如何在PyTorch Geometric中利用边特征实现链接预测?
基于边特征的PyTorch Geometric链接预测修改方案
1. 数据预处理调整
首先确保自定义数据集正确加载边特征,data.edge_attr需为形状[num_edges, 22]的张量。原示例的RandomLinkSplit仅处理边结构,需手动扩展负采样逻辑,核心是为负样本补充边特征(负样本本身无天然边特征,需根据任务场景生成):
from torch_geometric.transforms import RandomLinkSplit import torch # 拆分数据集 transform = RandomLinkSplit( num_val=0.1, num_test=0.1, is_undirected=True, add_negative_train_samples=True, ) train_data, val_data, test_data = transform(data) # 为负样本生成边特征:示例用随机初始化,可替换为节点特征组合等自定义逻辑 train_data.neg_edge_attr = torch.randn(train_data.neg_edge_index.size(1), 22) val_data.neg_edge_attr = torch.randn(val_data.neg_edge_index.size(1), 22) test_data.neg_edge_attr = torch.randn(test_data.neg_edge_index.size(1), 22)
2. 模型结构修改
原示例仅用节点嵌入预测链接,需修改预测层融合边特征,也可自定义GNN层让节点嵌入生成过程利用边特征:
方案一:仅在预测层融合边特征
保留原GNN生成节点嵌入的逻辑,修改链接预测器,将节点对嵌入与边特征拼接后做分类:
from torch_geometric.nn import SAGEConv from torch.nn import Linear import torch.nn.functional as F class GNN(torch.nn.Module): def __init__(self, in_channels, hidden_channels, out_channels): super().__init__() self.conv1 = SAGEConv(in_channels, hidden_channels) self.conv2 = SAGEConv(hidden_channels, out_channels) def forward(self, x, edge_index): x = self.conv1(x, edge_index).relu() x = self.conv2(x, edge_index) return x class LinkPredictor(torch.nn.Module): def __init__(self, node_emb_dim, edge_feat_dim, hidden_channels): super().__init__() self.lin1 = Linear(2 * node_emb_dim + edge_feat_dim, hidden_channels) self.lin2 = Linear(hidden_channels, hidden_channels) self.lin3 = Linear(hidden_channels, 1) def forward(self, x_i, x_j, edge_attr): # 拼接源节点嵌入、目标节点嵌入、边特征 x = torch.cat([x_i, x_j, edge_attr], dim=-1) x = self.lin1(x).relu() x = self.lin2(x).relu() return self.lin3(x)
方案二:GNN传播过程利用边特征
自定义支持边特征的GraphSAGE层,让节点嵌入更新时融合邻居节点与对应边的特征:
from torch_geometric.nn import SAGEConv class EdgeAwareSAGEConv(SAGEConv): def forward(self, x, edge_index, edge_attr=None): if edge_attr is not None: row, col = edge_index # 拼接邻居节点特征与对应边特征 neighbor_combined = torch.cat([x[col], edge_attr], dim=-1) # 聚合邻居特征 out = self.propagate(edge_index, x=neighbor_combined, aggr=self.aggr) out = self.lin_l(out) if self.root_weight: out += self.lin_r(x) if self.normalize: out = F.normalize(out, p=2., dim=-1) return out return super().forward(x, edge_index) class GNN(torch.nn.Module): def __init__(self, node_in_dim, edge_feat_dim, hidden_channels, out_channels): super().__init__() self.conv1 = EdgeAwareSAGEConv(node_in_dim + edge_feat_dim, hidden_channels) self.conv2 = EdgeAwareSAGEConv(hidden_channels + edge_feat_dim, out_channels) def forward(self, x, edge_index, edge_attr): x = self.conv1(x, edge_index, edge_attr).relu() x = self.conv2(x, edge_index, edge_attr) return x
3. 训练循环修改
需在训练时传入边特征,分别计算正负样本的预测损失:
def train(): model.train() link_predictor.train() optimizer.zero_grad() # 生成节点嵌入(若用方案二,需传入edge_attr) z = model(train_data.x, train_data.edge_index, train_data.edge_attr) # 正样本预测:传入正边的节点对与对应边特征 pos_out = link_predictor(z[train_data.pos_edge_index[0]], z[train_data.pos_edge_index[1]], train_data.edge_attr) # 负样本预测:传入负边的节点对与生成的负边特征 neg_out = link_predictor(z[train_data.neg_edge_index[0]], z[train_data.neg_edge_index[1]], train_data.neg_edge_attr) # 二分类交叉熵损失 pos_loss = F.binary_cross_entropy_with_logits(pos_out, torch.ones_like(pos_out)) neg_loss = F.binary_cross_entropy_with_logits(neg_out, torch.zeros_like(neg_out)) loss = pos_loss + neg_loss loss.backward() optimizer.step() return loss.item()
验证、测试循环逻辑与训练一致,只需替换为val_data或test_data即可。
关键注意点
- 负样本边特征处理:负样本无天然边特征,需根据任务场景选择生成方式:异常边检测可使用标注的虚假边(带真实特征);潜在边预测可使用节点特征组合(如拼接、点积)或随机初始化。
- GNN层选择:标准GCN仅支持标量边权重,若需在传播中利用多维边特征,建议选择GAT、自定义GraphSAGE等支持多维边输入的层。
- 特征融合优化:除简单拼接外,可引入注意力层让模型自动学习节点嵌入与边特征的权重占比,提升模型效果。
内容的提问来源于stack exchange,提问作者wup017
相关产品推荐
相关产品推荐

