ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

基于 DGL 复现 BGRL 自监督图表示学习:架构解析、训练配置与实验复现指南

基于 DGL 复现 BGRL 自监督图表示学习:架构解析、训练配置与实验复现指南 基于 DGL 复现 BGRL 自监督图表示学习架构解析、训练配置与实验复现指南【免费下载链接】dglPython package built to ease deep learning on graph, on top of existing DL frameworks.项目地址: https://gitcode.com/gh_mirrors/dg/dglBGRLBootstrapped Graph Latent Representations是一种通过自举bootstrap实现的大规模图自监督表示学习方法其核心思路是使用在线网络与动量更新的目标网络在两个增强视图之间进行对比学习无需负样本与标签。本指南以 examples/pytorch/bgrl 中的 DGL 官方实现为主线完整讲解其模型结构、命令行参数、数据预处理、训练与评估协议并给出 WikiCS、Amazon、Coauthor、PPI 等标准数据集上可直接复现的实验命令与性能基线。读完本文你将能够在自己的环境中独立跑通 BGRL 的直推式与归纳式训练并理解每一行关键代码背后的原理。一、方法背景与本文定位BGRL 出自论文Large-Scale Representation Learning on Graphs via BootstrappingarXiv:2102.06514属于图对比学习 / 自监督学习的代表性工作之一。与依赖负采样或大 batch 的经典对比方法不同BGRL 仅依赖两个增强视图之间的余弦相似度最大化并借鉴 BYOL 的动量机制来稳定训练因而天然适用于大规模图。本文所讲解的代码是社区贡献者实现并合入 DGL 官方仓库的 examples/pytorch/bgrl 目录它完整复刻了论文的实验设定是理解 BGRL 如何在 DGL 生态中落地的最佳参考实现。二、代码文件结构examples/pytorch/bgrl目录共包含 5 个文件职责划分清晰文件职责main.py命令行入口参数解析、训练循环、周期评估、权重保存model.pyBGRL 框架、GCN / GraphSAGE_GCN 编码器、MLP 预测器、LayerNormutils.py数据集加载、图增强变换DropEdge FeatMask、余弦退火调度器eval_function.py下游评估逻辑回归 / 线性层微调PPI 使用 Micro-F1README.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与SAGEConvutils.py 依赖dgl.dataloading.GraphDataLoader与dgl.transforms的Compose / DropEdge / FeatMask / RowFeatNormalizer这些都是 DGL 长期稳定的公开 API因此较新的 DGL 版本通常也能兼容运行eval_function.py 依赖 scikit-learn 的LogisticRegression、GridSearchCV、OneVsRestClassifier、metrics等模块用于下游线性评估训练代码在 main.py 中会自动探测 CUDAtorch.cuda.is_available()为真则使用 GPU否则退回 CPU。由于 BGRL 默认训练 10000 个 epoch强烈建议在 GPU 上运行。四、数据集说明原文档给出六个数据集的统计信息覆盖两类任务DatasetTaskNodesEdgesFeaturesClassesWikiCSTransductive11,701216,12330010Amazon ComputersTransductive13,752245,86176710Amazon PhotosTransductive7,650119,0817458Coauthor CSTransductive18,33381,8946,80515Coauthor PhysicsTransductive34,493247,9628,4155PPI24 graphsInductive56,944818,71650121多标签数据集加载逻辑集中在 utils.py 的get_dataset中直推式数据集coauthor_cs / coauthor_physics / amazon_computers / amazon_photos / wiki_cs通过 DGL 内置数据集类CoauthorCSDataset、CoauthorPhysicsDataset、AmazonCoBuyComputerDataset、AmazonCoBuyPhotoDataset、WikiCSDataset加载并统一应用RowFeatNormalizer(subtract_minTrue)做行归一化其中 WikiCS 在 get_wiki_cs 中额外做了逐特征标准化(feat - mean) / std同时保留其官方的 20 组预划分 maskPPI 是归纳式任务get_ppi在 utils.py 中通过PPIDataset(modetrain/valid/test)分别取 train/val/test 三个集合并利用GraphDataLoader将 trainval 的 24 张图打包成 batch_size22 的图 batch同时为每张图写入batch节点属性供编码器中的 batch 级 LayerNorm 使用。五、核心模型结构解析BGRL 的整体架构在 model.py 的BGRL类中实现由四个组件构成在线编码器online encoder用于生成在线表示参与梯度更新预测器predictorMLP从在线表示预测目标投影目标编码器target encoder在线编码器的深拷贝不参与梯度更新requires_gradFalse权重通过动量滑动平均更新动量更新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(momentum0.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 是单隐层 MLPLinear(input_size, hidden_size) - PReLU(1) - Linear(hidden_size, output_size)默认hidden_size512对应--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_cs、coauthor_physics、amazon_photos、amazon_computers、wiki_cs、ppiamazon_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.pyDropEdgeL1588、FeatMaskL238、RowFeatNormalizerL111其中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.mm0.99起步并逐渐增大与论文设置一致。7.3 训练主循环main.py 的主循环逻辑每个 epoch 内更新学习率与动量 → 生成两个增强视图 → 对非 PPI 数据集调用dgl.add_self_loop加自环 → 前向计算对称余弦损失 →optimizer.step()更新在线网络 →model.update_target_network(mm)动量更新目标网络每--eval_epochs默认 250个 epoch 评估一次直推式任务打印Test AccuracyPPI 打印Best Val F1与Test F1训练结束后将在线编码器权重保存为{weights_dir}/bgrl-{dataset}.pt例如../weights/bgrl-ppi.pt。优化器采用AdamWmain.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普通直推式数据集除 WikiCSfit_logistic_regression使用 20% 数据训练、80% 测试的随机划分OneVsRestClassifier(LogisticRegression(solverliblinear))结合GridSearchCV在C ∈ {2^-10, ..., 2^10}上网格搜索重复--num_eval_splits20次取平均WikiCSfit_logistic_regression_preset_splits使用官方预置的 20 组 train/val mask每组内用验证集挑选最优C后在测试集上报 AccuracyPPIfit_ppi_linear先对表示做标准化再训练一个 100 步的线性分类头在weight_decay ∈ {2^-10, ..., 2^10}上按验证集 Micro-F1 选优最终报告测试集 Micro-F1。十、性能结果原文档给出的复现结果如下。Accuracy 报告为 20 次随机数据划分与模型初始化的均值 ± 标准差Micro-F1 报告为 20 次随机模型初始化的均值 ± 标准差官方代码与 DGL 实现仅各跑 1 次随机数据划分与初始化。直推式任务AccuracyDatasetWikiCSAm. Comp.Am. PhotosCo. CSCo. PhyAccuracy Reported79.98 ± 0.1090.34 ± 0.1993.17 ± 0.3093.31 ± 0.1395.73 ± 0.05Accuracy Official Code79.9490.6293.4593.4295.74Accuracy DGL80.0090.6493.3493.7695.79归纳式任务PPIDatasetPPIMicro-F1 Reported69.41 ± 0.15Accuracy Official Code68.83Micro-F1 DGL68.65可以看到DGL 实现与论文报告值及官方参考实现基本持平甚至略有超出如 WikiCS 的 80.00 与 Co. CS 的 93.76在六个数据集上均达到了同一量级的复现水平。十一、使用建议与注意事项训练成本默认 10000 个 epoch、每 250 epoch 评估一次全量训练耗时长建议先在 GPU 上运行并观察前 1000~2000 epoch 的评估曲线是否收敛权重保存模型只保存在线编码器的state_dictbgrl-{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 BootstrappingarXiv:2102.06514BGRL 方法与实验设定的原始出处DGL 官方实现examples/pytorch/bgrl本文所有代码引用均来自该目录下的 main.py、model.py、utils.py、eval_function.pyDGL 内置数据集与变换源码python/dgl/data/gnn_benchmark.pyAmazon / Coauthor 系列、python/dgl/data/ppi.py、python/dgl/data/wikics.py、python/dgl/transforms/module.pyDropEdge / 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),仅供参考
返回列表