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

基于PyG/GraphSage自编码器预测新节点与现有节点的链接

基于PyTorch Geometric的GraphSage自编码器链接预测:新增孤立节点与链接预测问题

问题背景

我是GNN领域新手,用PyTorch Geometric(PyG)实现了基于两层SAGEConv的GraphSage自编码器做链接预测,现在需要解决两个问题:

  1. 如何向现有图中添加一个无关联边的新节点(带特征张量)?
  2. 如何预测这个新节点与哪些现有节点存在高概率链接?

已定义的模型与训练函数如下:

import torch
import torch.nn.functional as F
from torch_geometric.nn import SAGEConv
from torch_geometric.utils import negative_sampling

class Net(torch.nn.Module):
    def __init__(self, in_channels, hidden_channels, out_channels, dropout=0.1):
        super().__init__()
        self.dropout = dropout
        self.conv1 = SAGEConv(in_channels, hidden_channels)
        self.conv2 = SAGEConv(hidden_channels, out_channels)
        

    def encode(self, x, edge_index):
        x = self.conv1(x, edge_index)
        x = F.relu(x)
        x = F.dropout(x, p=self.dropout)

        x = self.conv2(x, edge_index)
        x = F.relu(x)
        x = F.dropout(x, p=self.dropout)

        return x
    
    def decode(self, z, edge_label_index):
        return (z[edge_label_index[0]] * z[edge_label_index[1]]).sum(
            dim=-1
        )  # product of a pair of nodes on each edge

    def decode_all(self, z):
        prob_adj = z @ z.t()
        return (prob_adj > 0).nonzero(as_tuple=False).t()

def train_link_predictor(
    model, train_data, val_data, optimizer, criterion, n_epochs=100
):

    for epoch in range(1, n_epochs + 1):

        model.train()
        optimizer.zero_grad()
        
        # create node embeddings (aggregating neighbor nodes with GraphSage)
        z = model.encode(train_data.x, train_data.edge_index)
        
        # sampling training negatives for every training epoch
        neg_edge_index = negative_sampling(
            edge_index=train_data.edge_index, num_nodes=train_data.num_nodes,
            num_neg_samples=train_data.edge_label_index.size(1), method='sparse')

        edge_label_index = torch.cat(
            [train_data.edge_label_index, neg_edge_index],
            dim=-1,
        )
        
        # edge labels contain 1 for positive edges and 0 for negative edges
        edge_label = torch.cat([
            train_data.edge_label,
            train_data.edge_label.new_zeros(neg_edge_index.size(1))
        ], dim=0)
        
        # the decoder makes a prediction based on the node embeddings by calculating pairwise dot-product 
        out = model.decode(z, edge_label_index).view(-1)
        
        # the loss is calculated by minimizing the difference between predictions and labeled values for pos/neg edges
        loss = criterion(out, edge_label)        
        loss.backward()
        optimizer.step()

        val_auc = eval_link_predictor(model, val_data)
        writer.add_scalar("Loss/train", loss, epoch)
        writer.add_scalar("AUC/train", val_auc, epoch)
        if epoch % 10 == 0:
            print(f"Epoch: {epoch:03d}, Train Loss: {loss:.3f}, Val AUC: {val_auc:.3f}")
            

    return model

我曾考虑将新节点特征张量加入data.x,并在邻接矩阵data.edge_index中添加无边条目,但不确定这是否为最优可行方案。


解决方案

1. 添加无关联边的新节点

你的思路方向完全正确,直接扩展节点特征张量即可,不需要修改edge_index(因为新节点是孤立节点,没有边关联)。具体操作:

  • 准备新节点的特征张量new_node_x,需保证其维度与现有节点特征一致(形状为[1, in_channels])
  • 将其拼接在原data.x的末尾:
# 示例:随机生成新节点特征,实际替换为你的目标特征
new_node_x = torch.randn(1, data.x.size(1), device=data.x.device)
# 扩展节点特征张量
data.x = torch.cat([data.x, new_node_x], dim=0)
  • 注意:SAGEConv会自动处理孤立节点——当节点没有邻居时,默认会直接使用节点自身特征作为聚合结果(对应aggr='mean'的逻辑),因此无需对edge_index做任何额外修改。

2. 预测新节点与现有节点的高概率链接

利用训练好的模型生成节点嵌入,再通过解码逻辑计算新节点与所有现有节点的链接概率,最后筛选高概率节点即可:

步骤1:生成所有节点的嵌入

model.eval()
with torch.no_grad():
    # 生成包含新节点在内的所有节点嵌入
    z = model.encode(data.x, data.edge_index)
    new_node_z = z[-1]  # 新节点是最后一个,提取其嵌入
    existing_nodes_z = z[:-1]  # 提取所有原有节点的嵌入

步骤2:计算链接概率

复用模型的decode逻辑(点积求和),构造新节点与所有现有节点的配对边索引,计算概率:

# 构造边索引:新节点索引(data.num_nodes-1)与每个现有节点的配对
new_edge_index = torch.tensor([
    [data.num_nodes-1] * len(existing_nodes_z),  # 新节点索引重复N次(N为原有节点数)
    list(range(data.num_nodes-1))  # 原有节点的索引列表
], device=data.x.device)

with torch.no_grad():
    # 计算每个配对的链接概率分数
    prob_scores = model.decode(z, new_edge_index)

步骤3:筛选高概率节点

对概率分数降序排序,取Top-K的节点:

top_k = 10  # 按需设置要取的高概率节点数量
# 按概率从高到低排序,得到节点索引
sorted_indices = torch.argsort(prob_scores, descending=True)
top_nodes = sorted_indices[:top_k]
top_scores = prob_scores[top_nodes]

# 输出结果
print("Top高概率链接的现有节点索引:", top_nodes.cpu().numpy())
print("对应的概率分数:", top_scores.cpu().numpy())

额外优化提示

  • 也可以直接用decode_all里的逻辑生成概率邻接矩阵,再提取新节点对应的行:
with torch.no_grad():
    prob_adj = z @ z.t()  # 生成全量节点配对的概率矩阵
new_node_probs = prob_adj[-1][:-1]  # 新节点与所有原有节点的概率
  • 训练阶段无需包含新节点,因为GraphSage的归纳能力就是用来处理训练时未出现的节点,新增节点属于推理阶段的归纳任务,完全符合模型设计目标。

内容的提问来源于stack exchange,提问作者feinheitsbrei

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 04:07:01