PyTorch Geometric 示例库实战指南:从入门 GCN 到万亿边图规模扩展
2026/9/12 20:23:41 网站建设 项目流程

PyTorch Geometric 示例库实战指南:从入门 GCN 到万亿边图规模扩展

【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric

本指南以 PyTorch Geometric(PyG)官方仓库中的 examples 目录 为核心,系统梳理该示例库覆盖的各类 GNN 应用场景——从最基础的节点分类(GCN)、链接预测(含 Attract-Repel 与 LPFormer 等进阶方法),到大图基准(OGB)、关系深度学习(Relational Deep Learning)、异构图、LLM 与 GNN 协同、可解释性,以及基于 cuGraph 的万亿边图扩展方案。读者读完后,将掌握如何在当前仓库中快速定位并运行对应示例,理解每个示例背后的关键实现细节(数据加载、模型结构、训练/评估流程),并知道如何将单卡示例平滑迁移到torch.compile、多 GPU 与分布式场景。

示例库的整体布局与定位

examples/目录(examples/README.md)汇集了覆盖不同 GNN 使用场景的示例脚本,每个脚本都是"小而完整"的可运行程序,包含数据加载、模型定义、训练循环与评估逻辑。该 README 本身扮演"导航地图"角色,重点标出若干代表性示例,并划分为以下主题簇:

  • 入门与基础任务gcn.py(节点分类)、link_pred.py(链接预测)
  • 进阶链接预测ar_link_pred.py(Attract-Repel 嵌入)、lpformer.py(Graph Transformer 链接预测)
  • 大规模基准(OGB)ogbn_train.pyogbn_proteins_deepgcn.py
  • 关系深度学习rdl.py(RelBench 数据集)
  • 节点属性预测新数据集graphland.py
  • 工程化主题子目录:examples/compile(torch.compile)、examples/multi_gpu(多 GPU/分布式)、examples/hetero(异构图)、examples/llm(LLM 与 GNN 协同)、examples/explain(可解释性)
  • 极致扩展:通过 cuGraph 将 PyG 扩展到万亿级边的图数据

运行环境方面,README 特别建议 NVIDIA GPU 用户优先使用官方推荐的 NVIDIA PyG Container 中的设备选择逻辑)。

入门第一课:用 GCN 做节点分类

README 明确指出,gcn.py是最适合入门的示例——它演示了如何在小型同构图数据(Cora 等 Planetoid 数据集)上训练 GCN 模型做节点级预测。该脚本堪称 PyG 最小工作流的完整范本,核心链路如下。

数据加载与特征归一化

parser.add_argument('--dataset', type=str, default='Cora') parser.add_argument('--hidden_channels', type=int, default=16) parser.add_argument('--lr', type=float, default=0.01) parser.add_argument('--epochs', type=int, default=200) parser.add_argument('--use_gdc', action='store_true', help='Use GDC') parser.add_argument('--wandb', action='store_true', help='Track experiment')

数据部分(gcn.py)通过Planetoid数据集类下载 Cora,并套用T.NormalizeFeatures()变换做行归一化。设备选择使用torch_geometric.device('auto')自动探测可用硬件,训练指标可通过 torch_geometric.logging 的init_wandb/log接口输出——加--wandb即可接入 Weights & Biases 追踪实验。

模型与训练细节

class GCN(torch.nn.Module): def __init__(self, in_channels, hidden_channels, out_channels): super().__init__() self.conv1 = GCNConv(in_channels, hidden_channels, normalize=not args.use_gdc) self.conv2 = GCNConv(hidden_channels, out_channels, normalize=not args.use_gdc)

注意两个细节:其一,GCNConv默认在层内做归一化;当启用--use_gdc(Graph Diffusion Convolution)时,归一化被提前到数据预处理阶段完成,因此卷积层需显式关闭normalize。其二,优化器(gcn.py)对conv1施加weight_decay=5e-4、对conv2不施加,注释明确说明"只在第一个卷积层做权重衰减"——这是复现经典 GCN 论文超参的关键。训练循环(train/test函数)用data.train_mask/val_mask/test_mask分别计算交叉熵损失与三组准确率,并按验证集最优保存测试结果;--use_gdc对应的T.GDC变换配置为 PPR 扩散(alpha=0.05)+ TopK 稀疏化(k=128),这是将 GCN 升级为 GDC 版本的完整实操样例。

运行方式:

python examples/gcn.py --dataset Cora python examples/gcn.py --dataset CiteSeer --use_gdc --wandb

链接预测三连:从基础版到 Graph Transformer

基础版:link_pred.py

link_pred.py展示了一个最小化的 GNN 链接预测管线,核心创新点在于RandomLinkSplit变换把数据切分成 train/val/test 三份

transform = T.Compose([ T.NormalizeFeatures(), T.ToDevice(device), T.RandomLinkSplit(num_val=0.05, num_test=0.1, is_undirected=True, add_negative_train_samples=False), ])

应用该变换后,数据对象会转变成(train_data, val_data, test_data)三元组。模型采用"编码器-解码器"结构:encode用两层 GCN 生成节点嵌入,decode对候选边两端嵌入做点积得到边分数;训练时每个 epoch 都重新做一次负采样(link_pred.py),把negative_sampling采样出的负边与正边拼接后交给BCEWithLogitsLoss监督;评估用roc_auc_score。脚本末尾的decode_all还能对整个图预测所有可能存在边的概率邻接矩阵。

进阶版:ar_link_pred.py 与 Attract-Repel 嵌入

ar_link_pred.py实现了基于论文《Pseudo-Euclidean Attract-Repel Embeddings for Undirected Graphs》的改进链接预测方法,README 指出该方案可显著提升 AUC(在对应论文报告中最多提升 23%)。其思想是把嵌入维度显式拆成"吸引"与"排斥"两部分

class ARLinkPredictor(torch.nn.Module): def __init__(self, in_channels): super().__init__() self.attract_dim = in_channels // 2 self.repel_dim = in_channels - self.attract_dim def forward(self, z_i, z_j): z_i_attr, z_i_repel = z_i[:, :self.attract_dim], z_i[:, self.attract_dim:] z_j_attr, z_j_repel = z_j[:, :self.attract_dim], z_j[:, self.attract_dim:] attract_score = (z_i_attr * z_j_attr).sum(dim=1) repel_score = (z_i_repel * z_j_repel).sum(dim=1) return attract_score - repel_score

脚本通过--use_ar开关在传统 MLP 打分器(LinkPredictor)与 AR 打分器之间切换,并支持 Cora/CiteSeer/PubMed 三个数据集;训练采用train_test_split_edges切分正负边,损失为正负样本二分类损失之和。更实用的是,脚本在训练结束后会输出R-fraction(排斥部分能量占比),用于量化吸引/排斥空间的分离程度——这是理解 AR 方法可解释性的关键指标:

python examples/ar_link_pred.py --use_ar --dataset Cora

进阶版:lpformer.py 与 LPFormer

lpformer.py演示了用 Graph Transformer 家族成员LPFormer(论文《Empowering GNNs with Edge-based Flexible Graph Transformer》)在ogbl-ppa数据集上做链接预测。它展示了 PyG 与 OGB 生态协作的完整流程:

  • PygLinkPropPredDataset加载数据、dataset.get_edge_split()取得正负边划分;
  • 模型初始化时传入ppr_thresholds(控制 CN/1-hop/多跳 PPR 的截断阈值列表)与gcn_cache=True等关键参数;
  • 通过model.calc_sparse_ppr(...)预计算稀疏 PPR 矩阵作为 Transformer 的结构先验;
  • 训练时对当前 batch 的正边做 mask 处理,避免信息泄漏;
  • 支持--runs多随机种子运行并汇总均值 ± 标准差,评估用 OGB 官方的Evaluator计算hits@K

README 同时提醒:ogbl-citation2的评估方式与其他 OGB 链接预测数据集不同(见脚本内注释及 LPFormer 原仓库说明),直接套用本脚本需要额外适配。

大图基准:在 OGB 数据集上训练 GNN

ogbn_train.py:多模型 + 邻居采样

ogbn_train.py是面向大规模 OGB 节点分类的通用训练脚本,支持ogbn-arxiv(默认)、ogbn-products(约 6200 万条边)与ogbn-papers100M(约 16 亿条边)三个数据集,README 以此为例说明 PyG 如何在大图上训练。它的工程化要点极具参考价值:

  • 模型可切换:通过--modelsagegatsgformerpolynormer之间选择,其中 SGFormer(默认,Graph Transformer 类)与 Polynormer 是 README 特别提及的代表性方法;
  • 邻居采样:用NeighborLoader做多跳子图采样,--fan_out控制每层邻居数、--batch_size控制 batch 大小、--num_workers并行加载;对 Transformer 类模型启用disjoint=True(子图不共享节点);
  • 图预处理开关--use_directed_graph控制是否保留有向图(默认转为无向图并reduce='mean'),--add_self_loop控制是否添加自环;
  • 内存预警:当选择ogbn-papers100M且系统内存不足约 390GB 时,脚本会打印显式警告——这直观说明了该数据集对硬件的要求;
  • Polynormer 特殊处理:论文推荐 7 层结构,脚本在层数不符时会提示;训练时先做--local_epochs轮局部(local)训练作为 warmup,再通过设置model._global = True切换为全局注意力;
  • 训练基建seed_everything(123)固定随机种子,ReduceLROnPlateau依据验证集准确率调度学习率,最终输出训练/推理/总体的平均与中位 epoch 耗时、最佳验证准确率与测试准确率。
python examples/ogbn_train.py --dataset ogbn-arxiv --model sgformer python examples/ogbn_train.py --dataset ogbn-products --model sage -b 1024 --fan_out 15

ogbn_proteins_deepgcn.py:训练深度 GNN

ogbn_proteins_deepgcn.py演示如何在ogbn-proteins数据集上训练深层 GNN,README 将其定位为"深度 GNN 训练范本"。几个实现亮点:

  • 特征工程:用scatter(..., reduce='sum')把边特征聚合到节点上作为初始节点特征(该数据集原始节点无特征);
  • 模型结构:28 层DeepGCNLayer,每层由GENConv(softmax 聚合、可学习温度t)、LayerNormReLU构成,使用block='res+'残差连接,并通过ckpt_grad=i % 3每三层做一次梯度检查点以节省显存;
  • 数据划分RandomNodeLoader将图随机切分为 40 份用于训练、5 份用于测试;
  • 多标签评估:任务是 112 类多标签二元分类,用BCEWithLogitsLoss训练,并用 OGBEvaluator计算 ROC-AUC。

关系深度学习:rdl.py 与 RelBench

rdl.py展示了如何基于RelBench数据集(论文见 README 引用)做关系深度学习(Relational Deep Learning,RDL)。RDL 的核心思路是把关系型数据库中的多张表建模为异构时序图,然后在图上做端到端预测。该示例是 PyG 与 torch_frame 之外两个生态(RelBench、PyTorch Frame)协同的完整样板,架构层次分明:

  1. 文本嵌入GloveTextEmbedding基于 SentenceTransformer 的 GloVe 平均词向量模型,把文本列编码为张量;
  2. 异构特征编码器HeteroEncoder为每种节点类型维护一个 PyTorch FrameResNet,按列语义类型(categorical/numerical/multicategorical/embedding/timestamp)选择对应的stype_encoder(如EmbeddingEncoderLinearEncoderTimestampEncoder);
  3. 时序编码器HeteroTemporalEncoderPositionalEncoding编码"种子时间与节点时间戳的相对差(换算为天数)";
  4. 消息传递HeteroGraphSAGEHeteroConv对每种边类型实例化SAGEConv,层间配LayerNorm(mode="node")
  5. 预测头:目标节点类型经过MLP输出预测。

数据侧,make_pkey_fkey_graph把数据库按主外键关系物化为HeteroData图对象;NeighborLoader支持input_time/time_attr/temporal_strategy做时序邻居采样,AttachTargetTransform负责在 batch 生成后把标签挂到对应节点上(因为时序采样中同一节点可能多次出现、对应不同标签)。任务类型通过get_task_type_params自动适配:回归用L1Loss+MAE,二分类用BCEWithLogitsLoss+ROC-AUC。训练结束后按验证指标保存best_model.pt并加载测试:

python examples/rdl.py --dataset rel-f1 --task <task_name> --epochs 10

--task的可选值以 RelBench 官网 发布的任务为准,脚本启动时会校验。)

节点属性预测新基准:graphland.py

graphland.py对应 README 提到的GraphLand 数据集(论文《GraphLand: A Scientific Discovery Dataset for Graph Neural Networks》),用于节点属性预测。它反映了 PyG 示例库持续跟进学术界新基准的节奏,用法与其他数据集脚本一致:加载数据集 → 构建模型 → 训练评估。

工程化进阶:编译、多 GPU、异构图、LLM 与可解释性

torch.compile:examples/compile

examples/compile 目录提供了使用torch.compile加速 PyG 模型的示例。以 compile/gcn.py 为例,与入门版相比有两处关键工程改动:

  • 数据预处理阶段用T.GCNNorm()提前完成归一化,卷积层设置normalize=False,避免编译过程中出现图中断与 CPU 通信;
  • 模型定义后直接model = torch.compile(model, dynamic=False)即可获得优化后的执行图,训练循环保持不变。这是把 PyG 模型接入 PyTorch 2.x 编译器的标准姿势。

多 GPU 与分布式:examples/multi_gpu

examples/multi_gpu/README.md 对分布式示例做了系统分类(详见该文档的示例表格),可归纳为三条主线:

  • NVIDIA GPU + cuGraph:官方推荐方案,单节点、多节点、链接预测等负载均有现成脚本(位于 cuGraph-PyG 示例仓库,安装方式见官方文档的 cuGraph 加速章节);
  • 纯 PyTorch 分布式distributed_batching.py(单节点,多小图图级预测,DataLoader+DistributedSampler)、distributed_sampling.py(单节点,Reddit 大图节点分类,NeighborLoader子图采样)、distributed_sampling_multinode.py与配套的.sbatch(多节点 + Slurm 提交)、papers100m_gcn.py/papers100m_gcn_multinode.py(约 16 亿边图)、pcqm4m_ogb.py(图级回归)、mag240m_graphsage.py(大规模异构图)、taobao.py(异构图链接预测)、model_parallel.py(手动把不同层放到不同 GPU 的模型并行);
  • Intel GPU(XPU)distributed_sampling_xpu.py支持单节点多卡的同构图邻居采样训练。

异构图:examples/hetero

examples/hetero 覆盖了异构图的典型任务:二部图 GraphSAGE(bipartite_sage.py)、带标签传播的 HAN(han_imdb.py)、异构卷积(hetero_conv_dblp.py)、异构链接预测(hetero_link_pred.py)、HGT(hgt_dblp.py)、CSV 数据加载(load_csv.py)、元路径随机游走(metapath2vec.py)、推荐系统(recommender_system.py)、时序链接预测(temporal_link_pred.py)以及to_hetero_mag.py等,适合作为异构建模的起点合集。

LLM 与 GNN 协同:examples/llm

examples/llm 集中了将大语言模型与 GNN 联合训练的示例,覆盖文本到知识图谱的 RAG(txt2kg_rag.py)、文本问答(txt2qa.py)、分子与大分子场景(git_mol.pyglem.pymolecule_gpt.pyprotein_mpnn.py)、图检索(g_retriever.pyrelbench_gretriever.py)等方向。README 引用的 GNN+LLM 研讨与演讲资料可帮助理解这一交叉方向的研究动机。

可解释性:examples/explain

examples/explain 提供 GNN 可解释性示例,包括基于 Captum 的解释器(captum_explainer.py及异构链接预测版)、GNNExplainer 系列(图分类、链接预测)、GraphMask 与 MGNAN 图分类等,对应 PyG 官方 explain 模块的实践入口。

终极扩展:cuGraph 把 PyG 带到万亿边

README 最后专门用一节介绍如何借助 cuGraph 把 PyG 扩展到万亿级边的图。背景信息如下:

  • cuGraph是 RAPIDS 框架下专注于 GPU 加速图分析的包集合,支持属性图与数千 GPU 规模扩展;
  • cuGraph GNNcugraph-gnn)通过cuGraph-PyGWholeGraph两个子项目为 PyTorch/PyG 提供原生 GPU 加速插件,底层基于pylibcugraph/libcugraph的高性能 C++ 采样原语,并配套libwholegraph/pylibwholegraph实现分布式边列表与嵌入存储;
  • 用户既可以直接使用这些底层库,也可以走 cuGraph-PyG 的高层 API——它直接实现了 PyG 的GraphStoreFeatureStoreNodeLoaderLinkLoader接口,因此上层训练代码几乎无需改动,即可无缝获得 GPU 加速的图存储与采样能力。

README 给出的落地路径是:参照官方安装指南完成 cuGraph 加速配置后,直接使用 cuGraph-PyG 示例仓库中现成的单节点/多节点/链接预测训练脚本(位于cugraph-pyg/cugraph_pyg/examples目录,覆盖完整可扩展工作流)。这与 examples/multi_gpu 中的建议一脉相承:追求 NVIDIA GPU 上极致性能时,优先采用 cuGraph 生态。

如何快速上手:推荐路径总结

结合 README 的导览顺序,建议按如下路径探索示例库:

  1. 入门:先跑 examples/gcn.py,理解数据加载 → 模型 → 训练/评估的最小闭环;
  2. 任务扩展:依次阅读 examples/link_pred.py(链接预测)、examples/ar_link_pred.py(AR 改进版)、examples/lpformer.py(Graph Transformer 版);
  3. 规模升级:用 examples/ogbn_train.py 在多模型/多数据集上做邻居采样训练,对照 examples/ogbn_proteins_deepgcn.py 学习深层 GNN 技巧;
  4. 场景拓展:按需进入 examples/hetero、examples/llm、examples/explain 等主题目录;
  5. 性能优化:参考 examples/compile 接入torch.compile,参考 examples/multi_gpu 扩展到多 GPU,最后按官方 cuGraph 指南迈向万亿边规模。

所有示例脚本均位于仓库examples/目录下,可直接以python examples/<name>.py运行(OGB 相关脚本需先安装ogb包,RDL 示例需安装relbenchtorch_framesentence-transformers等依赖);每个脚本都尽量保持自包含,是理解 PyG 各模块 API 用法的最直接教材。

【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询