基于 DGL 的 R-GCN 异构图节点分类:从 ogbn-mag 到 ogb-lsc-mag240m 的完整实战
【免费下载链接】dglPython package built to ease deep learning on graph, on top of existing DL frameworks.项目地址: https://gitcode.com/gh_mirrors/dg/dgl
R-GCN(Relational Graph Convolutional Network)是处理异构图(heterogeneous graph)节点分类任务的经典模型,它针对每条关系(relation)使用独立的权重矩阵,从而建模不同类型节点与边之间的复杂语义。本文以 DGL 官方示例 hetero_rgcn.py 为核心,完整讲解其在 OGB 两个真实大规模异构图数据集ogbn-mag与ogb-lsc-mag240m上的数据预处理、特征工程、模型结构与训练评估流程。读完本文,你将掌握:如何用 DGL 的图变换与邻居采样 API 搭建异构图上可扩展的 R-GCN 训练管线,如何为无原始特征的节点类型动态生成嵌入,以及如何在不同硬件配置(纯 CPU / 单 GPU)下启动训练并解读资源占用与精度结果。
一、示例概览与适用场景
该示例位于仓库 examples/core/rgcn/ 目录,目标是在异构图上完成节点分类:
- 同构图:所有节点与边类型相同,嵌入可以统一方式生成,无需类型级处理;
- 异构图:存在多种节点与边类型,需要为每种类型分别生成嵌入,才能精细捕获异构图的复杂结构与语义信息。
示例的完整功能流程(见 hetero_rgcn.py 中的流程图注释)为:
main ├── prepare_data 加载并预处理数据集 ├── rel_graph_embed 生成图嵌入(核心步骤) ├── 实例化 RGCN 模型 │ ├── RelGraphConvLayer(输入 → 隐藏层) │ └── RelGraphConvLayer(隐藏层 → 输出层) └── train ├── EntityClassify.forward(模型前向传播) └── evaluate(EntityClassify.evaluate 评估)作者在文档中明确说明:模型尚未针对最高精度进行调优,其价值在于展示一条完整、可运行、可扩展的异构图节点分类流水线。
二、支持的数据集与运行方式
2.1ogbn-mag:微软学术图谱子集
ogbn-mag是 OGB 官方节点属性预测数据集,包含四类节点(paper、author、institution、field_of_study)与四类关系,预测目标是paper节点的论文主题类别。预处理阶段:
- 使用
AddReverse()为每条边补充反向边(如writes之外新增rev_writes); - 使用
ToSimple()去除两点之间的重复边; author与institution两类节点没有原始特征,由嵌入层(embedding layer)动态生成。
在 CPU 上完成采样与训练/推理:
python3 hetero_rgcn.py --dataset ogbn-mag在 CPU 上采样、GPU 上训练/推理:
python3 hetero_rgcn.py --dataset ogbn-mag --num_gpus 12.2ogb-lsc-mag240m:大规模学术图谱
ogb-lsc-mag240m是 OGB 大规模挑战(LSC)中的超大规模数据集(约 2.44 亿节点、12.8 亿条边)。预处理阶段同样补充反向边、去除重复边,但特征处理策略不同:通过**消息传递(message passing)**预先为author与institution节点生成特征。由于该预处理通常耗时很长,README 提供了预处理的产物文件供直接下载使用:
paper-feat.npy:paper节点特征;author-feat.npy:author节点特征;inst-feat.npy:institution节点特征;hetero-graph.dgl:预处理后的异构图。
在 CPU 上完成采样与训练/推理:
python3 hetero_rgcn.py --dataset ogb-lsc-mag240m在 CPU 上采样、GPU 上训练/推理:
python3 hetero_rgcn.py --dataset ogb-lsc-mag240m --num_gpus 1三、完整命令行参数说明
除--dataset与--num_gpus外,脚本还提供了以下参数(见 hetero_rgcn.py 的参数解析部分):
| 参数 | 默认值 | 说明 |
|---|---|---|
--dataset | ogbn-mag | 训练数据集,可选ogbn-mag或ogb-lsc-mag240m |
--num_gpus | 0 | 使用的 GPU 数量,设为0表示纯 CPU 训练 |
--num_workers | 0 | 数据加载使用的 worker 进程数 |
--rootdir | ./dataset/ | OGB 数据集下载目录 |
--graph_path | ./graph.dgl | mag240m预处理图的加载路径 |
--paper_feature_path | ./paper-feat.npy | paper节点特征文件路径 |
--author_feature_path | ./author-feat.npy | author节点特征文件路径 |
--inst_feature_path | ./inst-feat.npy | institution节点特征文件路径 |
需要注意:--num_workers不应超过机器物理核心数。脚本通过psutil.cpu_count(logical=False)检测物理核心数量(逻辑核心数会因超线程等特性偏高),一旦num_workers >= expected_max会向 stderr 打印错误提示(hetero_rgcn.py)。
四、资源占用与训练耗时参考
README 给出的资源数据采集环境为 AWS EC2g4dn.metal(384GB RAM、96 vCPUs,Cascade Lake P-8259L、8 张 NVIDIA T4 16GB GPU)。其中 CPU 内存为free命令used字段的峰值(较粗略,RSS/USS/PSS更精确),GPU 内存为nvidia-smi记录的峰值。
4.1ogbn-mag(数据集约 1.1GB)
| 数据集大小 | CPU 内存占用 | GPU 数量 | GPU 内存占用 | 每 epoch 训练耗时 |
|---|---|---|---|---|
| ~1.1GB | ~7GB | 0 | 0GB | ~233s |
| ~1.1GB | ~5GB | 1 | 4.5GB | ~73.6s |
4.2ogb-lsc-mag240m(数据集约 404GB)
| 数据集大小 | CPU 内存占用 | GPU 数量 | GPU 内存占用 | 每 epoch 训练耗时 |
|---|---|---|---|---|
| ~404GB | ~72GB | 0 | 0GB | ~325s |
| ~404GB | ~61GB | 1 | 14GB | ~178s |
可以看出,使用 1 张 GPU 后训练耗时显著下降(ogbn-mag从 ~233s 降至 ~73.6s),同时 CPU 内存占用也随之降低(特征与中间结果更多驻留在 GPU 上)。这些数字为特定环境下的经验值,实际耗时随机器配置与数据加载 worker 数而变化。
五、数据预处理:图变换与邻居采样
5.1 图变换管线
在prepare_data中,ogbn-mag通过 DGL 的图变换 API 完成预处理(hetero_rgcn.py):
transform = Compose([ToSimple(), AddReverse()]) g = transform(g)三个变换的含义如下(实现位于 python/dgl/transforms/module.py):
ToSimple(module.py#L866-L925):将图转为无平行边的简单图。可选参数return_counts(保存原始边数的边特征名,默认count)与aggregator(重复边特征合并方式:arbitrary/sum/mean,默认arbitrary);AddReverse(module.py#L712-L791):为每条边(i,j)添加反向边(j,i)。对异构图,会为每个边类型新增rev_前缀的反向边类型(如('paper','cites','paper')之外新增('paper','rev_cites','paper')),可选copy_edata控制是否复制边特征;Compose(module.py#L1084):将多个变换按顺序组合为一个可调用对象。
mag240m路径则直接从预处理好的hetero-graph.dgl文件加载图,并显式指定g.formats(["csc"])以 CSC 格式存储(利于按目标节点聚合邻居的采样与消息传递)。
5.2 邻居采样与数据加载
模型采用两层图卷积,配合 DGL 的多层邻居采样器(hetero_rgcn.py):
sampler = dgl.dataloading.MultiLayerNeighborSampler([25, 10], fused=False) train_loader = dgl.dataloading.DataLoader( g, split_idx["train"], sampler, batch_size=1024, shuffle=True, num_workers=num_workers, device=device, )即第一层每节点采样 25 个邻居、第二层采样 10 个邻居,训练 batch size 为 1024,采样在 CPU 上进行。这种"CPU 采样 + 训练/推理设备可选"的架构正是 README 中"Sample on CPU and train/infer on CPU/GPU"两种模式的实现基础。评估阶段使用 batch size 4096、num_workers=0的 DataLoader(hetero_rgcn.py)。
六、特征工程:嵌入层 vs 预计算特征
这是同构与异构图分类的核心差异所在,脚本中由rel_graph_embed函数(hetero_rgcn.py)实现:
def rel_graph_embed(graph, embed_size): node_num = {} for ntype in graph.ntypes: if ntype == "paper": continue node_num[ntype] = graph.num_nodes(ntype) return HeteroEmbedding(node_num, embed_size)它遍历图中所有节点类型,为除paper之外的每个类型维护一张独立的(node_num[ntype], embed_size)嵌入表,返回dgl.nn.HeteroEmbedding实例。HeteroEmbedding(python/dgl/nn/pytorch/hetero.py#L345-L428)内部是多个torch.nn.Embedding组成的ModuleDict,每个节点类型独立训练;reset_parameters()采用 Xavier 均匀初始化。
两个数据集的差异体现在特征来源与维度:
ogbn-mag:输入特征维度feat_size = 128,paper节点使用数据集自带的原始特征(g.ndata["feat"]["paper"]),其余类型由HeteroEmbedding动态学习(hetero_rgcn.py);ogb-lsc-mag240m:输入特征维度feat_size = 768,三类节点特征均从磁盘上的.npy文件以mmap_mode="r+"内存映射方式读取(特征过大无法整体载入内存),且由于原始特征为 float16 而模型权重为 float32,前向前需.float()转换(hetero_rgcn.py)。源码注释同时指出,当前尚未启用 GPU 上的混合精度训练([TODO] 标记)。
extract_embed(hetero_rgcn.py)则负责在采样得到的input_nodes上索引嵌入层,只对非paper类型调用HeteroEmbedding。
七、模型结构:RelGraphConvLayer 与 EntityClassify
7.1 单层图卷积RelGraphConvLayer
RelGraphConvLayer(hetero_rgcn.py)是模型的基本构件,其结构包含三部分:
HeteroGraphConv:为每条关系实例化一个GraphConv(in_size, out_size, norm="right", weight=False, bias=False)。norm="right"表示按目标节点入度归一化聚合消息(等价于对接收消息取平均);weight=False, bias=False是因为脚本改用自定义权重矩阵。HeteroGraphConv(python/dgl/nn/pytorch/hetero.py#L12-L120)会为每个目标节点类型聚合来自不同关系子模块的输出,默认聚合方式为求和(aggregate='sum');self.weight(ModuleDict):为每条关系创建一个无偏置的nn.Linear(in_size, out_size),即关系专属的权重矩阵,前向时通过mod_kwargs={"weight": weight.T}注入卷积模块;self.loop_weights(ModuleDict):为每个节点类型创建带偏置的nn.Linear,作用相当于残差连接——用目标节点自身特征更新输出。源码注释特别强调:这不代表图中存在自环边,只是类似于残差连接的操作。
前向过程(hetero_rgcn.py)为:
g = g.local_var() # 防止修改原始图(副作用隔离) weight_dict = {rel: {"weight": self.weight[rel].weight.T} for rel in relation_names} inputs_dst = {k: v[: g.number_of_dst_nodes(k)] for k, v in inputs.items()} hs = self.conv(g, inputs, mod_kwargs=weight_dict) # 对每个节点类型:h = conv 结果 + loop_weight(inputs_dst),再经激活与 dropout7.2 整体模型EntityClassify
EntityClassify(hetero_rgcn.py)堆叠两层RelGraphConvLayer:
- 第一层:输入特征 → 64 维隐藏层,激活函数 ReLU,dropout 0.5;
- 第二层:64 维隐藏层 → 输出类别数,无激活。
关系列表由list(set(g.etypes))去重后排序得到。前向时逐层对采样的blocks应用卷积。
八、训练与评估流程
训练主循环(hetero_rgcn.py)要点:
- 预测目标类型固定为
paper(category = "paper"),每 epoch 仅遍历训练集 batch; - 训练 3 个 epoch。源码注释解释:该数据集上通常第 1~2 个 epoch 即达到最佳验证性能,因此 max epoch 设为 3;
- 损失为
log_softmax后的负对数似然F.nll_loss; - 优化器为
torch.optim.Adam,学习率0.01;itertools.chain将模型与嵌入层参数合并交给优化器统一更新(hetero_rgcn.py); - 每个 epoch 结束后在验证集与测试集上评估,
mag240m评估时会保存 test-dev 提交文件(evaluator.save_test_submission)。
设备选择逻辑为:cuda:0当且仅当torch.cuda.is_available()且--num_gpus > 0,否则回退 CPU(hetero_rgcn.py)。
脚本还在初始化后调用reset_parameters()(hetero_rgcn.py),源码注释解释了原因:若不重置参数,模型会沿用上一次运行的参数,可能因陷入较差的局部最优而导致结果偏差或次优性能。
九、运行结果示例
9.1ogbn-mag准确率
README 提供的训练日志(3 个 epoch):
Epoch: 01, Loss: 2.3386, Valid: 47.67%, Test: 46.96% Epoch: 02, Loss: 1.5563, Valid: 47.66%, Test: 47.02% Epoch: 03, Loss: 1.1557, Valid: 46.58%, Test: 45.42% Test accuracy 45.38509.2ogb-lsc-mag240m准确率
README 提供的验证集日志(3 个 epoch,文档未给出测试精度):
Epoch: 01, Loss: 2.0798, Valid: 52.04% Epoch: 02, Loss: 1.8652, Valid: 54.51% Epoch: 03, Loss: 1.8175, Valid: 53.71%需要再次强调:以上精度来自未调优的模型,作为"端到端流程可运行"的基线参考,而非性能上限。
十、源码级要点小结
- 图变换三件套:
ToSimple+AddReverse+Compose是实现"去除重复边 + 补充反向关系"的标准管线,均可在 python/dgl/transforms/module.py 中查阅实现; - 类型级特征策略:
HeteroEmbedding为无原始特征的节点类型提供可学习嵌入(python/dgl/nn/pytorch/hetero.py#L345-L428),有原始特征的paper节点直接使用图上的feat数据; - 关系专属卷积:
HeteroGraphConv按关系分配独立GraphConv子模块,配合自定义ModuleDict权重实现标准 R-GCN 的"每关系一矩阵"范式,并以loop_weights充当残差连接; - 可扩展采样管线:
MultiLayerNeighborSampler([25, 10])+ 1024 batch size 使模型能扩展到 4 亿节点的超大规模异构图,且采样与训练设备可分离,适配纯 CPU 与 CPU+GPU 两种部署形态; - 大规模特征读取:
mmap_mode="r+"内存映射加载超大.npy特征,避免一次性载入内存导致 OOM。
若希望复现上述结果,可在安装 DGL 与 OGB 相关依赖后,按第二节给出的命令直接运行 hetero_rgcn.py,并根据自身机器配置调整--num_workers、--rootdir与特征文件路径。
【免费下载链接】dglPython package built to ease deep learning on graph, on top of existing DL frameworks.项目地址: https://gitcode.com/gh_mirrors/dg/dgl
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考