基于 DGL 复现 BGRL 自监督图表示学习:架构解析、训练配置与实验复现指南
【免费下载链接】dglPython package built to ease deep learning on graph, on top of existing DL frameworks.项目地址: https://gitcode.com/gh_mirrors/dg/dgl
BGRL(Bootstrapped Graph Latent Representations)是一种通过"自举"(bootstrap)实现的大规模图自监督表示学习方法,其核心思路是使用在线网络与动量更新的目标网络,在两个增强视图之间进行对比学习,无需负样本与标签。本指南以 examples/pytorch/bgrl 中的 DGL 官方实现为主线,完整讲解其模型结构、命令行参数、数据预处理、训练与评估协议,并给出 WikiCS、Amazon、Coauthor、PPI 等标准数据集上可直接复现的实验命令与性能基线。读完本文,你将能够在自己的环境中独立跑通 BGRL 的直推式与归纳式训练,并理解每一行关键代码背后的原理。
一、方法背景与本文定位
BGRL 出自论文Large-Scale Representation Learning on Graphs via Bootstrapping(arXiv:2102.06514),属于图对比学习 / 自监督学习的代表性工作之一。与依赖负采样或大 batch 的经典对比方法不同,BGRL 仅依赖两个增强视图之间的余弦相似度最大化,并借鉴 BYOL 的动量机制来稳定训练,因而天然适用于大规模图。
本文所讲解的代码是社区贡献者实现并合入 DGL 官方仓库的 examples/pytorch/bgrl 目录,它完整复刻了论文的实验设定,是理解 BGRL 如何在 DGL 生态中落地的最佳参考实现。
二、代码文件结构
examples/pytorch/bgrl目录共包含 5 个文件,职责划分清晰:
| 文件 | 职责 |
|---|---|
| main.py | 命令行入口:参数解析、训练循环、周期评估、权重保存 |
| model.py | BGRL 框架、GCN / GraphSAGE_GCN 编码器、MLP 预测器、LayerNorm |
| utils.py | 数据集加载、图增强变换(DropEdge + FeatMask)、余弦退火调度器 |
| eval_function.py | 下游评估:逻辑回归 / 线性层微调,PPI 使用 Micro-F1 |
| README.md | 官方使用说明、实验命令与性能表 |
三、环境依赖与安装前提
原文档要求代码基于 Python 3.8 开发,并给出如下依赖版本组合:
dgl 0.8.3 numpy 1.21.2 torch 1.10.2 scikit-learn 1.0.2需要注意的是:
- 从源码调用链看,model.py 依赖
dgl.nn.pytorch.conv中的GraphConv与SAGEConv,utils.py 依赖dgl.dataloading.GraphDataLoader与dgl.transforms的Compose / DropEdge / FeatMask / RowFeatNormalizer,这些都是 DGL 长期稳定的公开 API,因此较新的 DGL 版本通常也能兼容运行; - eval_function.py 依赖 scikit-learn 的
LogisticRegression、GridSearchCV、OneVsRestClassifier、metrics等模块,用于下游线性评估; - 训练代码在 main.py 中会自动探测 CUDA:
torch.cuda.is_available()为真则使用 GPU,否则退回 CPU。由于 BGRL 默认训练 10000 个 epoch,强烈建议在 GPU 上运行。
四、数据集说明
原文档给出六个数据集的统计信息,覆盖两类任务:
| Dataset | Task | Nodes | Edges | Features | Classes |
|---|---|---|---|---|---|
| WikiCS | Transductive | 11,701 | 216,123 | 300 | 10 |
| Amazon Computers | Transductive | 13,752 | 245,861 | 767 | 10 |
| Amazon Photos | Transductive | 7,650 | 119,081 | 745 | 8 |
| Coauthor CS | Transductive | 18,333 | 81,894 | 6,805 | 15 |
| Coauthor Physics | Transductive | 34,493 | 247,962 | 8,415 | 5 |
| PPI(24 graphs) | Inductive | 56,944 | 818,716 | 50 | 121(多标签) |
数据集加载逻辑集中在 utils.py 的get_dataset中:
- 直推式数据集(coauthor_cs / coauthor_physics / amazon_computers / amazon_photos / wiki_cs)通过 DGL 内置数据集类
CoauthorCSDataset、CoauthorPhysicsDataset、AmazonCoBuyComputerDataset、AmazonCoBuyPhotoDataset、WikiCSDataset加载,并统一应用RowFeatNormalizer(subtract_min=True)做行归一化; - 其中 WikiCS 在 get_wiki_cs 中额外做了逐特征标准化
(feat - mean) / std,同时保留其官方的 20 组预划分 mask; - PPI 是归纳式任务,
get_ppi在 utils.py 中通过PPIDataset(mode="train"/"valid"/"test")分别取 train/val/test 三个集合,并利用GraphDataLoader将 train+val 的 24 张图打包成 batch_size=22 的图 batch,同时为每张图写入batch节点属性,供编码器中的 batch 级 LayerNorm 使用。
五、核心模型结构解析
BGRL 的整体架构在 model.py 的BGRL类中实现,由四个组件构成:
- 在线编码器(online encoder):用于生成在线表示,参与梯度更新;
- 预测器(predictor):MLP,从在线表示预测目标投影;
- 目标编码器(target encoder):在线编码器的深拷贝,不参与梯度更新(
requires_grad=False),权重通过动量滑动平均更新; - 动量更新:
update_target_network(mm)执行param_k = mm * param_k + (1 - mm) * param_q,其中mm为动量系数。
前向传播的逻辑(model.py)为:对两个增强视图分别过在线编码器得到online_y,经预测器得到online_q;目标视图在torch.no_grad()下过目标编码器得到target_y,训练目标即最大化预测结果与目标投影之间的余弦相似度。
5.1 GCN 编码器(直推式任务)
GCN 由GraphConv卷积层 +BatchNorm1d(momentum=0.99)+PReLU激活交替堆叠而成,层数由--graph_encoder_layer决定(例如512 256表示两层 GCN,输出维度为 256)。输入特征取自g.ndata["feat"]。
5.2 GraphSAGE_GCN 编码器(归纳式 PPI 任务)
GraphSAGE_GCN 是专为 PPI 多图归纳任务设计的 3 层网络:
- 使用
SAGEConv(..., "mean")均值聚合卷积; - 引入两条从原始输入到第 2、3 层的跳跃连接(
skip_lins); - 采用自定义的 LayerNorm,支持按
batch节点属性进行 batch 维度归一化(PPI 的图 batch 场景); - 激活函数为 PReLU。
5.3 MLP 预测器
MLP_Predictor 是单隐层 MLP:Linear(input_size, hidden_size) -> PReLU(1) -> Linear(hidden_size, output_size),默认hidden_size=512,对应--predictor_hidden_size参数。
5.4 训练损失
在 main.py 中,损失为对称余弦相似度损失的负均值:
loss = 2 - cos_sim(q1, y2.detach()).mean() - cos_sim(q2, y1.detach()).mean()其中q1, y2 = model(x1, x2)、q2, y1 = model(x2, x1),detach()确保目标分支不反传梯度,这与 BYOL/BGRL 的标准做法一致。
六、命令行参数详解
全部参数由 main.py 中的argparse定义,原文档将其分为四组:
数据集选项
| 参数 | 类型 | 说明 | 默认值 |
|---|---|---|---|
--dataset | str | 图数据集名称,可选coauthor_cs、coauthor_physics、amazon_photos、amazon_computers、wiki_cs、ppi | amazon_photos |
模型选项
| 参数 | 类型 | 说明 | 默认值 |
|---|---|---|---|
--graph_encoder_layer | list(int) | 卷积层隐藏维度,可传多个值表示层数与宽度 | [256, 128] |
--predictor_hidden_size | int | 预测器隐藏层大小 | 512 |
训练选项
| 参数 | 类型 | 说明 | 默认值 |
|---|---|---|---|
--epochs | int | 训练总 epoch 数 | 10000 |
--lr | float | 学习率 | 0.00001 |
--weight_decay | float | 权重衰减 | 0.00001 |
--mm | float | 目标网络动量系数 | 0.99 |
--lr_warmup_epochs | int | 学习率 warmup 周期 | 1000 |
--weights_dir | str | 权重保存目录 | ../weights |
增强选项(两个视图各一组,共两个值)
| 参数 | 类型 | 说明 | 默认值 |
|---|---|---|---|
--drop_edge_p | list(float) | 两个增强视图各自的边丢弃概率 | [0., 0.] |
--feat_mask_p | list(float) | 两个增强视图各自的节点特征掩码概率 | [0., 0.] |
评估选项
| 参数 | 类型 | 说明 | 默认值 |
|---|---|---|---|
--eval_epochs | int | 每隔多少 epoch 评估一次 | 250 |
--num_eval_splits | int | 评估时使用的数据划分 / 初始化次数 | 20 |
--data_seed | int | 数据划分随机种子 | 1 |
此外,源码中还有一个--num_experiments(默认 20),但主循环中未直接使用,属于论文报告 20 次随机初始化结果时的实验性参数。
七、训练流程与关键机制
7.1 两个视图的增强管线
增强变换由 get_graph_drop_transform 构造,顺序为:copy.deepcopy复制图 →(可选)DropEdge随机删边 →(可选)FeatMask按列随机掩码feat特征。--drop_edge_p与--feat_mask_p各自传入两个值,分别对应视图 1 与视图 2 的增强强度,从而实现"相同图、不同视角"。
对应 DGL 内置变换的实现位于 python/dgl/transforms/module.py:DropEdge(L1588)、FeatMask(L238)、RowFeatNormalizer(L111),其中FeatMask的p参数含义为"特征张量某一列被掩码的概率"。
7.2 学习率与动量的余弦退火调度
CosineDecayScheduler 实现带 warmup 的余弦退火:
- 学习率调度器:
CosineDecayScheduler(args.lr, args.lr_warmup_epochs, args.epochs),前 1000 个 epoch 线性升温至lr,随后余弦下降; - 动量调度器:
CosineDecayScheduler(1 - args.mm, 0, args.epochs),在训练中动量从 0.01 沿余弦曲线上升,对应代码中mm = 1 - mm_scheduler.get(step),即动量从args.mm(0.99)起步并逐渐增大,与论文设置一致。
7.3 训练主循环
main.py 的主循环逻辑:
- 每个 epoch 内:更新学习率与动量 → 生成两个增强视图 → 对非 PPI 数据集调用
dgl.add_self_loop加自环 → 前向计算对称余弦损失 →optimizer.step()更新在线网络 →model.update_target_network(mm)动量更新目标网络; - 每
--eval_epochs(默认 250)个 epoch 评估一次,直推式任务打印Test Accuracy,PPI 打印Best Val F1与Test F1; - 训练结束后,将在线编码器权重保存为
{weights_dir}/bgrl-{dataset}.pt,例如../weights/bgrl-ppi.pt。
优化器采用AdamW(main.py),仅优化model.trainable_parameters(),即在线编码器与预测器的参数,目标网络不参与优化。
八、实验复现命令
以下命令直接来自原文档,并给出注释说明:
直推式任务
# Coauthor CS python main.py --dataset coauthor_cs --graph_encoder_layer 512 256 --drop_edge_p 0.3 0.2 --feat_mask_p 0.3 0.4 # Coauthor Physics python main.py --dataset coauthor_physics --graph_encoder_layer 256 128 --drop_edge_p 0.4 0.1 --feat_mask_p 0.1 0.4 # WikiCS python main.py --dataset wiki_cs --graph_encoder_layer 512 256 --drop_edge_p 0.2 0.3 --feat_mask_p 0.2 0.1 --lr 5e-4 # Amazon Photos python main.py --dataset amazon_photos --graph_encoder_layer 256 128 --drop_edge_p 0.4 0.1 --feat_mask_p 0.1 0.2 --lr 1e-4 # Amazon Computers python main.py --dataset amazon_computers --graph_encoder_layer 256 128 --drop_edge_p 0.5 0.4 --feat_mask_p 0.2 0.1 --lr 5e-4归纳式任务
# PPI python main.py --dataset ppi --graph_encoder_layer 512 512 --drop_edge_p 0.3 0.25 --feat_mask_p 0.25 0. --lr 5e-3从上述命令可归纳出调参规律:不同数据集的最优增强强度差异明显,Amazon Computers 需要最强的边丢弃(0.5 / 0.4),而 PPI 的视图 2 不掩码特征(0.);学习率也需要按数据集单独调整,WikiCS / Amazon 系列通常需要1e-4 ~ 5e-4的量级。
九、下游评估协议详解
BGRL 训练完成后并不直接输出分类结果,而是冻结编码器、把学到的表示喂给简单的线性模型评估表示质量,逻辑集中在 eval_function.py:
- 普通直推式数据集(除 WikiCS):
fit_logistic_regression使用 20% 数据训练、80% 测试的随机划分,OneVsRestClassifier(LogisticRegression(solver="liblinear"))结合GridSearchCV在C ∈ {2^-10, ..., 2^10}上网格搜索,重复--num_eval_splits(20)次取平均; - WikiCS:
fit_logistic_regression_preset_splits使用官方预置的 20 组 train/val mask,每组内用验证集挑选最优C后在测试集上报 Accuracy; - PPI:
fit_ppi_linear先对表示做标准化,再训练一个 100 步的线性分类头,在weight_decay ∈ {2^-10, ..., 2^10}上按验证集 Micro-F1 选优,最终报告测试集 Micro-F1。
十、性能结果
原文档给出的复现结果如下。Accuracy 报告为 20 次随机数据划分与模型初始化的均值 ± 标准差,Micro-F1 报告为 20 次随机模型初始化的均值 ± 标准差;官方代码与 DGL 实现仅各跑 1 次随机数据划分与初始化。
直推式任务(Accuracy)
| Dataset | WikiCS | Am. Comp. | Am. Photos | Co. CS | Co. Phy |
|---|---|---|---|---|---|
| Accuracy Reported | 79.98 ± 0.10 | 90.34 ± 0.19 | 93.17 ± 0.30 | 93.31 ± 0.13 | 95.73 ± 0.05 |
| Accuracy Official Code | 79.94 | 90.62 | 93.45 | 93.42 | 95.74 |
| Accuracy DGL | 80.00 | 90.64 | 93.34 | 93.76 | 95.79 |
归纳式任务(PPI)
| Dataset | PPI |
|---|---|
| Micro-F1 Reported | 69.41 ± 0.15 |
| Accuracy Official Code | 68.83 |
| Micro-F1 DGL | 68.65 |
可以看到,DGL 实现与论文报告值及官方参考实现基本持平甚至略有超出(如 WikiCS 的 80.00 与 Co. CS 的 93.76),在六个数据集上均达到了同一量级的复现水平。
十一、使用建议与注意事项
- 训练成本:默认 10000 个 epoch、每 250 epoch 评估一次,全量训练耗时长,建议先在 GPU 上运行并观察前 1000~2000 epoch 的评估曲线是否收敛;
- 权重保存:模型只保存在线编码器的
state_dict(bgrl-{dataset}.pt),加载后可直接用于下游任务的特征提取; - 数据集首次加载:DGL 内置数据集首次使用会下载原始数据,需要网络连接;WikiCS 与 PPI 数据规模较大,注意磁盘空间;
- 随机性控制:
--data_seed控制评估时数据划分的随机状态,配合--num_eval_splits可得到稳定的平均指标;模型层面的随机性可通过 PyTorch 的全局种子自行控制; - 扩展性:BGRL 框架对编码器类型无硬性要求(目标编码器要求具备
reset_parameters方法),如需替换为 GAT、GraphSAGE 等其他 DGL 卷积层,只需参照GCN/GraphSAGE_GCN实现新的编码器并替换--graph_encoder_layer对应的网络即可。
十二、参考资料
- 论文:Large-Scale Representation Learning on Graphs via Bootstrapping(arXiv:2102.06514),BGRL 方法与实验设定的原始出处;
- DGL 官方实现:examples/pytorch/bgrl,本文所有代码引用均来自该目录下的 main.py、model.py、utils.py、eval_function.py;
- DGL 内置数据集与变换源码:python/dgl/data/gnn_benchmark.py(Amazon / Coauthor 系列)、python/dgl/data/ppi.py、python/dgl/data/wikics.py、python/dgl/transforms/module.py(DropEdge / FeatMask / RowFeatNormalizer);
- 若要参考仓库内其他自监督对比学习实现进行横向对比,可查看 examples/pytorch/grace、examples/pytorch/bgrl 等相邻示例目录。
【免费下载链接】dglPython package built to ease deep learning on graph, on top of existing DL frameworks.项目地址: https://gitcode.com/gh_mirrors/dg/dgl
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考