基于 DGL 复现 BGRL 自监督图表示学习:架构解析、训练配置与实验复现指南
2026/9/23 1:53:04 网站建设 项目流程

基于 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.pyBGRL 框架、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中的GraphConvSAGEConv,utils.py 依赖dgl.dataloading.GraphDataLoaderdgl.transformsCompose / DropEdge / FeatMask / RowFeatNormalizer,这些都是 DGL 长期稳定的公开 API,因此较新的 DGL 版本通常也能兼容运行;
  • eval_function.py 依赖 scikit-learn 的LogisticRegressionGridSearchCVOneVsRestClassifiermetrics等模块,用于下游线性评估;
  • 训练代码在 main.py 中会自动探测 CUDA:torch.cuda.is_available()为真则使用 GPU,否则退回 CPU。由于 BGRL 默认训练 10000 个 epoch,强烈建议在 GPU 上运行。

四、数据集说明

原文档给出六个数据集的统计信息,覆盖两类任务:

DatasetTaskNodesEdgesFeaturesClasses
WikiCSTransductive11,701216,12330010
Amazon ComputersTransductive13,752245,86176710
Amazon PhotosTransductive7,650119,0817458
Coauthor CSTransductive18,33381,8946,80515
Coauthor PhysicsTransductive34,493247,9628,4155
PPI(24 graphs)Inductive56,944818,71650121(多标签)

数据集加载逻辑集中在 utils.py 的get_dataset中:

  • 直推式数据集(coauthor_cs / coauthor_physics / amazon_computers / amazon_photos / wiki_cs)通过 DGL 内置数据集类CoauthorCSDatasetCoauthorPhysicsDatasetAmazonCoBuyComputerDatasetAmazonCoBuyPhotoDatasetWikiCSDataset加载,并统一应用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类中实现,由四个组件构成:

  1. 在线编码器(online encoder):用于生成在线表示,参与梯度更新;
  2. 预测器(predictor):MLP,从在线表示预测目标投影;
  3. 目标编码器(target encoder):在线编码器的深拷贝,不参与梯度更新(requires_grad=False),权重通过动量滑动平均更新;
  4. 动量更新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定义,原文档将其分为四组:

数据集选项

参数类型说明默认值
--datasetstr图数据集名称,可选coauthor_cscoauthor_physicsamazon_photosamazon_computerswiki_csppiamazon_photos

模型选项

参数类型说明默认值
--graph_encoder_layerlist(int)卷积层隐藏维度,可传多个值表示层数与宽度[256, 128]
--predictor_hidden_sizeint预测器隐藏层大小512

训练选项

参数类型说明默认值
--epochsint训练总 epoch 数10000
--lrfloat学习率0.00001
--weight_decayfloat权重衰减0.00001
--mmfloat目标网络动量系数0.99
--lr_warmup_epochsint学习率 warmup 周期1000
--weights_dirstr权重保存目录../weights

增强选项(两个视图各一组,共两个值)

参数类型说明默认值
--drop_edge_plist(float)两个增强视图各自的边丢弃概率[0., 0.]
--feat_mask_plist(float)两个增强视图各自的节点特征掩码概率[0., 0.]

评估选项

参数类型说明默认值
--eval_epochsint每隔多少 epoch 评估一次250
--num_eval_splitsint评估时使用的数据划分 / 初始化次数20
--data_seedint数据划分随机种子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),其中FeatMaskp参数含义为"特征张量某一列被掩码的概率"。

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 的主循环逻辑:

  1. 每个 epoch 内:更新学习率与动量 → 生成两个增强视图 → 对非 PPI 数据集调用dgl.add_self_loop加自环 → 前向计算对称余弦损失 →optimizer.step()更新在线网络 →model.update_target_network(mm)动量更新目标网络;
  2. --eval_epochs(默认 250)个 epoch 评估一次,直推式任务打印Test Accuracy,PPI 打印Best Val F1Test F1
  3. 训练结束后,将在线编码器权重保存为{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"))结合GridSearchCVC ∈ {2^-10, ..., 2^10}上网格搜索,重复--num_eval_splits(20)次取平均;
  • WikiCSfit_logistic_regression_preset_splits使用官方预置的 20 组 train/val mask,每组内用验证集挑选最优C后在测试集上报 Accuracy;
  • PPIfit_ppi_linear先对表示做标准化,再训练一个 100 步的线性分类头,在weight_decay ∈ {2^-10, ..., 2^10}上按验证集 Micro-F1 选优,最终报告测试集 Micro-F1。

十、性能结果

原文档给出的复现结果如下。Accuracy 报告为 20 次随机数据划分与模型初始化的均值 ± 标准差,Micro-F1 报告为 20 次随机模型初始化的均值 ± 标准差;官方代码与 DGL 实现仅各跑 1 次随机数据划分与初始化。

直推式任务(Accuracy)

DatasetWikiCSAm. Comp.Am. PhotosCo. CSCo. Phy
Accuracy Reported79.98 ± 0.1090.34 ± 0.1993.17 ± 0.3093.31 ± 0.1395.73 ± 0.05
Accuracy Official Code79.9490.6293.4593.4295.74
Accuracy DGL80.0090.6493.3493.7695.79

归纳式任务(PPI)

DatasetPPI
Micro-F1 Reported69.41 ± 0.15
Accuracy Official Code68.83
Micro-F1 DGL68.65

可以看到,DGL 实现与论文报告值及官方参考实现基本持平甚至略有超出(如 WikiCS 的 80.00 与 Co. CS 的 93.76),在六个数据集上均达到了同一量级的复现水平。

十一、使用建议与注意事项

  1. 训练成本:默认 10000 个 epoch、每 250 epoch 评估一次,全量训练耗时长,建议先在 GPU 上运行并观察前 1000~2000 epoch 的评估曲线是否收敛;
  2. 权重保存:模型只保存在线编码器的state_dictbgrl-{dataset}.pt),加载后可直接用于下游任务的特征提取;
  3. 数据集首次加载:DGL 内置数据集首次使用会下载原始数据,需要网络连接;WikiCS 与 PPI 数据规模较大,注意磁盘空间;
  4. 随机性控制--data_seed控制评估时数据划分的随机状态,配合--num_eval_splits可得到稳定的平均指标;模型层面的随机性可通过 PyTorch 的全局种子自行控制;
  5. 扩展性: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),仅供参考

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

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

立即咨询