DGL + MXNet 实现 GraphSAGE 归纳式节点分类:从论文复现到参数调优
【免费下载链接】dglPython package built to ease deep learning on graph, on top of existing DL frameworks.项目地址: https://gitcode.com/gh_mirrors/dg/dgl
GraphSAGE(Graph Sample and Aggregation)是图表示学习领域极具代表性的归纳式(inductive)方法,其核心思想是学习一个聚合函数(aggregator),将节点自身特征与其邻居特征结合,从而为训练阶段从未见过的节点生成嵌入。本文基于 DGL 仓库中的 MXNet 参考实现(examples/mxnet/graphsage/main.py),完整讲解如何在 Cora、Citeseer、Pubmed 三个经典引文网络上复现 GraphSAGE 节点分类,深入剖析模型结构与 SAGEConv 层的源码实现,并给出可复现的完整命令行参数说明与结果对照。
GraphSAGE 与示例程序概览
该示例对应论文Inductive Representation Learning on Large Graphs(NeurIPS 2017),文中模型通过采样并聚合邻居特征来生成节点嵌入,与直推式(transductive)方法(如 GCN)不同,GraphSAGE 可以泛化到未见过的图结构,因此适合大规模、动态变化的图场景。示例代码是对论文作者开源参考实现(williamleif/graphsage-simple)的简化复刻,聚焦于节点分类这一核心任务。
示例程序的结构非常清晰:
| 文件 | 作用 |
|---|---|
| examples/mxnet/graphsage/main.py | 完整的训练与评估入口:数据集加载、模型构建、训练循环、验证与测试 |
| examples/mxnet/graphsage/README.md | 运行说明与基准结果 |
模型层SAGEConv并不在本示例目录内,而是复用 DGL 统一提供的图卷积层模块 python/dgl/nn/mxnet/conv/sageconv.py,这也是 DGL 封装消息传递范式、屏蔽后端差异的典型体现:同一份 GraphSAGE 逻辑在 PyTorch、MXNet、TensorFlow 三个后端都有对应实现。
环境准备与依赖安装
示例运行只依赖一个额外的 Python 包requests(用于数据集下载)。在已安装 DGL 与 MXNet 的环境中执行:
pip install requests之后即可直接运行。若尚未安装 DGL 的 MXNet 后端,可参考仓库根目录 README.md 中的安装说明,按对应 MXNet 版本安装 DGL 包。
数据集:DGL 内置引文网络
示例通过 python/dgl/data/citation_graph.py 中定义的CoraGraphDataset、CiteseerGraphDataset、PubmedGraphDataset加载数据。三个数据集均为论文引用网络:节点表示论文,边表示引用关系,特征是论文的词袋向量(已做行归一化),任务是预测论文所属类别。
各数据集统计信息(来自源码 docstring)如下:
| 数据集 | 节点数 | 边数 | 类别数 | 特征维度 | 训练/验证/测试划分 |
|---|---|---|---|---|---|
| Cora | 2708 | 10556 | 7 | 1433 | 140 / 500 / 1000 |
| Citeseer | 3327 | 9228 | 6 | 3703 | 120 / 500 / 1000 |
| Pubmed | 19717 | 88651 | 3 | 500 | 60 / 500 / 1000 |
数据加载后,图对象以dgl.DGLGraph形式返回,节点特征、标签和三个 mask(train_mask、val_mask、test_mask)分别存放在g.ndata["feat"]、g.ndata["label"]、g.ndata["train_mask"]等字段中。程序入口处调用register_data_args(parser)(定义于 python/dgl/data/init.py)注册--dataset参数,运行时在main()中根据数据集名选择对应的 Dataset 类,并通过data.num_classes获得类别数、data.graph.number_of_edges()获得边数。
需要说明的是:Cora、Citeseer、Pubmed 均属于同质图(homogeneous graph),且CitationGraphDataset默认添加反向边(reverse_edge=True),以保证信息在无向语义下的传播。
模型结构:GraphSAGE 的 MXNet 实现
模型类 GraphSAGE
main.py中定义的GraphSAGE(nn.Block)是一个标准的多层堆叠结构,由三层SAGEConv组成:
- 输入层:
SAGEConv(in_feats, n_hidden, aggregator_type, feat_drop=dropout, activation=activation),将原始特征映射到隐藏维度; - 隐藏层:循环添加
n_layers - 1个SAGEConv(n_hidden, n_hidden, ...); - 输出层:
SAGEConv(n_hidden, n_classes, aggregator_type, feat_drop=dropout, activation=None),输出层不使用激活函数,直接接交叉熵损失。
前向传播forward(features)非常简单:将特征逐层喂给所有卷积层即可,因为聚合逻辑完全封装在SAGEConv内部:
def forward(self, features): h = features for layer in self.layers: h = layer(self.g, h) return hSAGEConv 层的源码级剖析
核心层实现在 python/dgl/nn/mxnet/conv/sageconv.py。其数学形式为:
h_N(i)^(l+1) = aggregate({h_j^l, ∀j ∈ N(i)}) h_i^(l+1) = σ(W · concat(h_i^l, h_N(i)^(l+1))) h_i^(l+1) = norm(h_i^(l+1))即先聚合邻居特征得到h_neigh,再与自身特征拼接后经线性变换(可选激活与归一化)。构造参数如下:
| 参数 | 类型 | 默认值 | 说明 |
|---|---|---|---|
in_feats | int 或 (int, int) | 必填 | 输入特征维度;支持同质图与单向二分图(源/目标节点维度不同时传元组)。gcn聚合器要求源目标特征维度一致 |
out_feats | int | 必填 | 输出特征维度 |
aggregator_type | str | mean | 聚合器类型,取值mean/gcn/pool/lstm,非法取值抛出DGLError |
feat_drop | float | 0.0 | 特征 dropout 概率 |
bias | bool | True | 是否添加可学习偏置 |
norm | callable | None | 输出归一化函数 |
activation | callable | None | 输出激活函数 |
forward中根据聚合器类型走不同的消息传递路径(通过 DGL 的update_all完成,详见 python/dgl/heterograph.py):
- mean:
fn.copy_u("h", "m")复制源节点特征,fn.mean("m", "neigh")对邻居取均值得到h_neigh; - gcn:先对邻居特征求和,再与自身特征相加后除以
(入度 + 1)做归一化,等价于 GCN 的规范化传播,因此要求源与目标特征维度一致(check_eq_shape校验); - pool:先用一个全连接层
fc_pool对每个源节点特征做 ReLU 变换,再对邻居取逐元素最大值(fn.max); - lstm:在 MXNet 实现中目前直接
raise NotImplementedError,仅保留接口。
对于非gcn聚合器,最终输出为fc_self(h_self) + fc_neigh(h_neigh),即自身变换与邻居聚合变换之和;gcn则只对h_neigh做变换。所有全连接层均使用 Xavier 初始化(mx.init.Xavier(magnitude=sqrt(2.0)))。此外,forward使用graph.local_scope()包裹,避免在原始图上残留临时特征,并显式处理了无向边图(graph.num_edges() == 0)的边界情况。
训练流程详解
main()中的训练流程可分为以下步骤:
- 数据加载与设备迁移:
--gpu为负时使用mx.cpu(0),否则调用g = g.int().to(ctx)将图迁移到 GPU,并将特征、标签同步到对应上下文; - 自环处理:
g = dgl.remove_self_loop(g)后g = dgl.add_self_loop(g),确保每个节点在聚合时包含自身信息(gcn 聚合器下此操作与deg + 1归一化配合); - 模型与优化器:
gluon.Trainer(model.collect_params(), "adam", {"learning_rate": args.lr, "wd": args.weight_decay}),使用 Adam 优化器; - 训练循环:
mx.autograd.record()下前向计算得到pred,损失为带train_mask掩码的SoftmaxCELoss,按训练样本数取平均;loss.backward()后trainer.step(batch_size=1)更新参数——由于是全图训练(full-batch),batch_size 取 1 表示一次更新覆盖全部节点; - 日志输出:从第 3 个 epoch 起统计每轮耗时,并打印 Loss、验证集 Accuracy 以及吞吐量
ETputs(KTEPS)(每秒千条边,n_edges / mean_dur / 1000); - 测试评估:训练结束后在
test_mask上计算最终测试准确率。
评估函数evaluate直接对全图做前向,取argmax预测类别,与标签比对后按 mask 加权计算准确率。
完整命令行参数
在main.py末尾通过argparse注册了全部超参数,register_data_args额外补充--dataset。汇总如下:
| 参数 | 默认值 | 说明 |
|---|---|---|
--dataset | 必填 | 可选cora/citeseer/pubmed |
--dropout | 0.5 | 特征 dropout 概率 |
--gpu | -1 | GPU 编号,负值表示使用 CPU |
--lr | 1e-2 | 学习率 |
--n-epochs | 200 | 训练轮数 |
--n-hidden | 16 | 隐藏层单元数 |
--n-layers | 1 | 隐藏层数量(不含输入输出层) |
--weight-decay | 5e-4 | L2 正则权重 |
--aggregator-type | gcn | 聚合器类型:mean/gcn/pool/lstm |
运行与基准结果
在仓库目录下执行(示例命令将完整命令写为):
python3 main.py --dataset cora --gpu 0三个数据集在默认超参数下的参考准确率(来自 examples/mxnet/graphsage/README.md,在对应 GPU 上复现):
- cora:约0.817
- citeseer:约0.699
- pubmed:约0.790
这些结果与论文报告的水平相当,可当作复现正确性的判定基准。无 GPU 时可用--gpu -1跑 CPU 版本,训练时间会相应增长。
调优建议与扩展方向
- 聚合器选择:默认
gcn聚合器在三个数据集上表现稳定;mean聚合器对特征分布更鲁棒,pool聚合器表达能力更强但参数量更大,适合在特征维度较高(如 Citeseer 的 3703 维)时尝试。注意lstm在 MXNet 后端尚未实现,选择后会直接报NotImplementedError; - 层数与隐藏维度:
--n-layers增加可捕获更高阶邻居信息,但全图训练下易出现过平滑(oversmoothing),建议从 1~2 层起步;--n-hidden增大通常能提升拟合能力,但也需同步提高--weight-decay防止过拟合; - 正则化:默认
dropout=0.5、weight_decay=5e-4是引文网络上的常见配置,数据集较小时可适当增大 dropout; - 迁移到大规模图:本示例采用全图训练,Pubmed(约 2 万节点)已是上限附近。若要扩展到 Reddit 等更大规模的图,DGL 提供了基于邻居采样的分布式/小批量训练方案,相关实现可见 examples/pytorch/graphsage 与 python/dgl/dataloading,这也正是 GraphSAGE 归纳式设计的用武之地。
小结
本文以 DGL 仓库的 MXNet 参考实现为主线,完整梳理了 GraphSAGE 归纳式节点分类的复现路径:从数据集加载、模型堆叠到 SAGEConv 四种聚合器的底层实现,再到训练循环与超参调优。示例代码虽然精简,却涵盖了 DGL 消息传递范式(update_all+copy_u/mean/sum/max)与后端统一模块(dgl.nn.mxnet.conv.SAGEConv)的核心用法,是理解图神经网络工程化落地的一个高质量范本。读者可在此基础上替换数据集、调整聚合器,或进一步探索采样式训练以支撑更大规模的图数据。
【免费下载链接】dglPython package built to ease deep learning on graph, on top of existing DL frameworks.项目地址: https://gitcode.com/gh_mirrors/dg/dgl
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考