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

PyTorch Geometric图神经网络训练后单缺失链接预测方法咨询

解决PyG异构图链接预测:单条边的概率计算问题

核心问题拆解

你拿到的模型输出是对数几率(Logits),这是二分类模型未经过激活的原始输出,需要通过sigmoid函数转换为0-1区间的概率值,才能直接解读为链接存在的可能性。

单条边预测的实现步骤

针对MovieLens的异构图场景,你需要构造包含目标用户-电影节点对的输入数据,再传入模型计算得分,最后转换为概率:

import torch

def is_there_a_link(user_node_id, movie_node_id, model, test_data):
    # 1. 构造单条异构图边的输入(匹配模型训练时的输入格式)
    # 复用原数据的节点特征和结构,仅替换目标边对
    edge_index = torch.tensor([[user_node_id], [movie_node_id]], device=test_data.device)
    input_data = test_data.clone()
    input_data['user', 'rates', 'movie'].edge_index = edge_index
    
    # 2. 切换到推理模式,关闭梯度计算
    model.eval()
    with torch.no_grad():
        # 3. 获取模型输出的logit,转换为概率
        logit = model(input_data)[0]  # 提取单条边的输出结果
        probability = torch.sigmoid(logit).item()
    
    # 4. 判断并返回结果
    return 'YES' if probability > 0.5 else 'NO'

# 调用示例
prediction = is_there_a_link(test_data['user'].node_id[1], test_data['movie'].node_id[3], model, test_data)
print(prediction)

关键细节说明

  • 输入构造逻辑:PyG异构图模型依赖完整节点特征计算嵌入,因此需要基于原测试数据克隆新对象,仅替换目标边对,确保模型能正确生成节点的嵌入表示。
  • 设备对齐:确保输入数据的设备(CPU/CUDA)与模型一致,避免运行时设备不匹配错误。
  • 推理模式优化:model.eval()会关闭dropout等训练专属的随机操作,torch.no_grad()禁用梯度计算,大幅提升推理效率。
  • Logit转概率:torch.sigmoid()将任意实数映射到0-1区间,值越接近1,代表用户与电影之间存在链接的概率越高。

额外优化:批量预测多条边

如果需要一次性预测多组用户-电影对,可以直接构造批量edge_index:

def predict_links(user_node_ids, movie_node_ids, model, test_data):
    edge_index = torch.tensor([user_node_ids, movie_node_ids], device=test_data.device)
    input_data = test_data.clone()
    input_data['user', 'rates', 'movie'].edge_index = edge_index
    
    model.eval()
    with torch.no_grad():
        logits = model(input_data)
        probabilities = torch.sigmoid(logits).tolist()
    
    return ['YES' if p > 0.5 else 'NO' for p in probabilities]

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.13 13:25:15