You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.19 07:34:52