DGL GraphBolt 多 GPU 训练实战:用 DistributedItemSampler 与 DDP 训练 GraphSAGE 节点分类模型
2026/9/23 3:36:46 网站建设 项目流程
  • 人工智能
  • 机器学习
  • 深度学习
  • 图计算

【免费下载链接】dgl

Python package built to ease deep learning on graph, on top of existing DL frameworks.

项目地址:https://gitcode.com/gh_mirrors/dg/dgl
点击查看免费下载

本指南围绕 DGL 仓库中 examples/multigpu/graphbolt 目录下的多 GPU 训练示例展开,讲解如何在多张 GPU 上使用 GraphBolt 数据加载器(DataLoader)与 PyTorch 的分布式数据并行(Distributed Data Parallel,DDP)训练 GraphSAGE 节点分类模型。读者学完后将掌握DistributedItemSampler的分片采样原理、DDP 环境初始化、Join上下文管理器处理不均衡输入、以及跨 rank 加权聚合评估指标等完整的多卡训练技术方案。

快速运行

该示例的入口脚本位于 examples/multigpu/graphbolt/node_classification.py,运行方式如下:

python node_classification.py --gpu=0,1
  • --gpu参数接受逗号分隔的 GPU 编号列表,例如--gpu=0,1表示在 GPU 0 与 GPU 1 上并行训练;
  • 脚本会通过torch.multiprocessing.spawn为每个 GPU 派生一个子进程,每个子进程对应一个 DDP rank(world_size等于 GPU 数量);
  • 训练结束后会在 rank 0 上打印验证集准确率(每个 epoch)与测试集准确率。

前置知识

阅读本示例前,官方建议先熟悉两个基础内容:

  1. 单卡 GraphBolt 节点分类示例:即 examples/graphbolt/node_classification.py,它演示了使用gb.ItemSamplersample_neighborfetch_featuregb.DataLoader构建端到端 GraphBolt 训练流水线的方法,多卡版本正是将其中的ItemSampler替换为DistributedItemSampler后的分布式扩展。
  2. 经典 GraphSAGE 实现:即 examples/core/graphsage/node_classification.py,用于理解 GraphSAGE 模型的训练范式。

总体执行流程

脚本源码顶部的注释给出了完整的流程示意图,可概括为两个阶段:

main │ ├───> OnDiskDataset 预处理(gb.BuiltinDataset(args.dataset).load()) │ └───> run (multiprocessing) │ ├───> 初始化进程组并构建分布式 SAGE 模型(DDP) │ ├───> train │ │ │ ├───> 使用 DistributedItemSampler 构建 GraphBolt dataloader │ │ │ └───> 训练循环(SAGE.forward → 验证集评估 → 收集各 rank 的指标) │ └───> 测试集评估

主进程先加载数据集,然后通过mp.spawn启动world_size个子进程,每个子进程独立执行run(rank, world_size, args, devices, dataset)

多 GPU 环境初始化

run函数开头完成分布式环境设置(见 node_classification.py):

device = devices[rank] torch.cuda.set_device(device) dist.init_process_group( backend="nccl", # 分布式 GPU 训练使用 NCCL 后端 init_method="tcp://127.0.0.1:12345", world_size=world_size, rank=rank, )

要点说明:

  • backend="nccl":多 GPU 训练推荐使用 NCCL 后端,它针对 NVIDIA GPU 通信做了深度优化;
  • init_method:示例使用tcp://127.0.0.1:12345作为默认的进程组初始化方式,仅适用于单机多卡场景;若在跨机集群中使用,需要替换为实际可用的通信地址;
  • main中还会设置os.environ["OMP_NUM_THREADS"] = str(mp.cpu_count() // 2 // world_size)限制线程数以避免资源竞争,并通过mp.set_sharing_strategy("file_system")设置多进程共享策略。

初始化进程组后,脚本会将数据集中的图结构与特征搬运到对应设备:

if args.storage_device == "pinned": graph = dataset.graph.pin_memory_() feature = dataset.feature.pin_memory_() else: graph = dataset.graph.to(args.storage_device) feature = dataset.feature.to(args.storage_device)
  • storage_device == "pinned"时调用pin_memory_()将图与特征固定在 CPU 的页锁定内存(pinned memory)中,支持异步传输与重叠取数;
  • 否则通过.to()将数据放到 CPU 或 CUDA 设备上,具体取决于--mode参数。

DistributedItemSampler:多卡数据分片的核心

这是多 GPU 版本区别于单卡版本的核心组件,其完整实现位于 python/dgl/graphbolt/item_sampler.py。它与单卡ItemSampler的最大区别在于:原始 item 集合会先按 replica(进程)切分成互不重叠的子集,每个 rank 只在自己的子集上执行(可选的)shuffle 与 batch 划分,因此每个 replica 始终获得确定且互斥的一份数据。

构造参数

参数类型默认值说明
item_setItemSet/HeteroItemSet必填待采样的数据,例如train_setvalidation_settest_set
batch_sizeint必填mini-batch 的大小,即一批处理的样本数量
drop_lastboolFalse是否丢弃最后一个不完整的 batch
shuffleboolFalse是否在采样前打乱数据
drop_uneven_inputsboolFalse是否让所有 rank 的 batch 数量保持一致,丢弃多出的部分
seedintNone可复现的随机种子;为None时自动生成

构造时(item_sampler.py),采样器会通过dist.get_world_size()dist.get_rank()获取当前分布式环境信息,并在world_size > 1时调用_align_seeds(src=0)dist.broadcast将 rank 0 的种子同步给所有 rank(item_sampler.py),从而保证各进程的随机行为一致、便于复现实验。

实际分片行为

内部通过calculate_range计算每个 rank 与每个 worker 的起止范围(实现在 python/dgl/graphbolt/internal/item_sampler_utils.py)。以torch.arange(15)batch_size=2、4 个 replica 为例,源码 docstring 给出了各类参数组合下的输出:

  • 全部为False时:Replica#0 得到[0,1],[2,3],Replica#1 得到[4,5],[6,7],Replica#2 得到[8,9],[10,11],Replica#3 得到[12,13],[14],各 rank 的 batch 数不均等;
  • drop_last=True, drop_uneven_inputs=False:Replica#3 只剩[12,13],其余不变;
  • drop_last=True, drop_uneven_inputs=True:所有 rank 都只保留 1 个 batch(Replica#0[0,1]、Replica#1[4,5]、Replica#2[8,9]、Replica#3[12,13]),各 rank batch 数量完全一致;
  • shuffle=True时,各 rank 在自己的子集内以seed + epoch为随机种子打乱数据,且每个 epoch 顺序都会变化。

注意DistributedItemSampler特意没有使用torch.utils.data.functional_datapipe装饰,即不支持函数式调用,但可以继续追加其他可迭代 datapipe(如copy_tosample_neighborfetch_feature)。

构建分布式 Dataloader 流水线

示例中的create_dataloader函数(node_classification.py)完整演示了多卡数据流水线的组装方式:

datapipe = gb.DistributedItemSampler( item_set=itemset, batch_size=args.batch_size, drop_last=is_train, shuffle=is_train, drop_uneven_inputs=is_train, ) if args.storage_device != "cpu": datapipe = datapipe.copy_to(device) datapipe = datapipe.sample_neighbor( graph, args.fanout, overlap_fetch=args.storage_device == "pinned", asynchronous=args.storage_device != "cpu", ) datapipe = datapipe.fetch_feature(features, node_feature_keys=["feat"]) if args.storage_device == "cpu": datapipe = datapipe.copy_to(device) dataloader = gb.DataLoader(datapipe, args.num_workers)

各步骤的作用:

  1. gb.DistributedItemSampler:按 rank 切分 item 子集并产出 mini-batch,训练阶段(is_train=True)同时开启shuffledrop_lastdrop_uneven_inputs,验证/测试阶段三者均关闭;
  2. copy_to(device)(非 CPU 存储时提前执行):将数据先拷贝到目标设备,使后续采样操作直接在 GPU 上运行;
  3. sample_neighbor(graph, fanout, ...):为每个 batch 的种子节点采样邻居,fanout长度必须与模型层数一致(默认10,10,10对应三层 GraphSAGE);overlap_fetch在 pinned 内存模式下开启以重叠取数,asynchronous在非 CPU 存储时开启异步采样;
  4. fetch_feature(features, node_feature_keys=["feat"]):为采样得到的子图拉取节点特征;
  5. gb.DataLoader(datapipe, args.num_workers):用num_workers个进程并行加载数据。

训练循环中,从每个data中解出三部分(node_classification.py):

x = data.node_features["feat"] # 第一层计算图的源节点特征 y = data.labels # 最后一层计算图的目标节点标签 blocks = data.blocks # 逐层的消息传递块 y_hat = model(blocks, x)

其中blocks是邻居采样产生的多层计算图块(Block),逐层经过SAGE.forward中的SAGEConv(聚合方式为"mean",隐藏层 256 维),中间层经过ReLUDropout(0.5)(node_classification.py)。

模型与训练循环:DDP + Join + 加权聚合

构建分布式模型

model = SAGE(in_size, hidden_size, out_size).to(device) model = DDP(model)

每个 rank 拥有一份完整的模型副本(replica),DDP负责在反向传播时通过allreduce同步梯度。特征输入维度in_size通过feature.size("node", None, "feat")[0]动态获取,out_sizenum_classes

Join 上下文管理器处理不均衡输入

DDP 要求所有 rank 的输入数量一致,否则程序可能报错或挂起。示例提供了两种解决方案:

  1. PyTorch 的Join上下文管理器(示例采用的方式):
with Join([model]): for data in (tqdm.tqdm(train_dataloader) if rank == 0 else train_dataloader): ...
  1. drop_uneven_inputs=True(在DistributedItemSampler中设置),通过丢弃多出的 batch 使各 rank 的 batch 数量一致。

示例在训练阶段同时开启drop_uneven_inputsJoin,双重保障训练不会因输入不均衡而中断。进度条tqdm只在 rank 0 上显示,避免多进程输出刷屏。

加权聚合损失与准确率

由于各 GPU 处理的样本数量可能不同,简单取平均会得到有偏的指标。示例实现了weighted_reduce(node_classification.py):

def weighted_reduce(tensor, weight, dst=0): dist.reduce(tensor=tensor, dst=dst) weight = torch.tensor(weight, device=tensor.device) dist.reduce(tensor=weight, dst=dst) return tensor / weight
  • dist.reduce将各 rank 的张量归约到dst指定的进程(默认 rank 0),默认使用ReduceOp.SUM求和;
  • 损失乘以各自样本数求和后再除以总样本数,得到精确的加权平均损失;验证准确率同样以acc * num_val_items加权求和后除以总样本数(node_classification.py)。

训练每个 epoch 后还会调用torch.cuda.synchronize()同步,保证计时准确,并在 rank 0 打印Epoch / Average Loss / Accuracy / Time

验证与测试评估

evaluate函数(node_classification.py)在torch.no_grad()下遍历 dataloader:

  • 每个 rank 在自己的验证/测试子集上推理,收集预测y_hats与标签y
  • 使用torchmetrics.functional.accuracytask="multiclass")计算准确率;
  • 返回(准确率, 样本数)weighted_reduce加权平均。

测试阶段同样只在 rank 0 打印Test Accuracy,最后调用dist.destroy_process_group()清理进程组(node_classification.py)。

命令行参数详解

脚本通过argparse提供以下参数(node_classification.py):

参数默认值可选值/说明
--gpu"0"逗号分隔的 GPU 编号,如0,1,2,3;GPU 数量即world_size
--epochs10训练轮数
--lr0.001学习率(Adam 优化器)
--batch-size1024mini-batch 大小
--fanout"10,10,10"邻居采样扇出,逗号分隔;长度必须与模型层数一致
--num-workers0数据加载进程数
--gpu-cache-size0GPU 特征缓存容量(字节)
--dataset"ogbn-products"支持ogbn-arxivogbn-productsogbn-papers100M
--mode"pinned-cuda"数据存储位置与训练设备组合:cpu-cuda(图/特征在 CPU 内存)、pinned-cuda(图/特征在页锁定内存)、cuda-cuda(图/特征在 GPU 显存)

其中--mode会被拆分为storage_device与训练设备两部分(args.storage_device, _ = args.mode.split("-")),并决定数据流水线中copy_tooverlap_fetchasynchronous的行为。当--gpu-cache-size > 0且存储设备不是cuda时,示例会用gb.gpu_cached_feature(实现见 python/dgl/graphbolt/impl/gpu_cached_feature.py)为节点特征挂上 GPU 缓存,将热点特征缓存在显存中以减少 PCIe 传输。

数据集通过gb.BuiltinDataset(args.dataset).load()加载为OnDiskDataset,训练/验证/测试子集分别取自dataset.tasks[0]train_setvalidation_settest_set,类别数来自dataset.tasks[0].metadata["num_classes"]

使用建议

  • 单机多卡场景下默认的tcp://127.0.0.1:12345初始化即可工作;跨机训练需修改init_method并配合 DGL 的分布式工具链使用;
  • --fanout长度必须与模型层数严格一致,否则采样阶段会报错;
  • 验证/测试阶段不要开启shuffledrop_last,以保证评估覆盖全部样本;
  • 若发现 DDP 训练因输入不均衡而卡死,优先确认训练 dataloader 是否同时启用了Join上下文管理器或drop_uneven_inputs
  • 进一步了解 GraphBolt 数据加载器各 datapipe 的详细用法,可参考 notebooks/graphbolt/walkthrough.ipynb。
  • 人工智能
  • 机器学习
  • 深度学习
  • 图计算

【免费下载链接】dgl

Python package built to ease deep learning on graph, on top of existing DL frameworks.

项目地址:https://gitcode.com/gh_mirrors/dg/dgl
点击查看免费下载

相关推荐

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

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

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

立即咨询