PyG 分布式训练完全指南:torch_geometric.distributed 架构、分区与采样原理
2026/9/13 14:13:06 网站建设 项目流程

PyG 分布式训练完全指南:torch_geometric.distributed 架构、分区与采样原理

【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric

导读:本文以 PyG(PyTorch Geometric)内置的torch_geometric.distributed模块为核心,系统讲解其在单机内存无法容纳的超大规模图数据上进行分布式 GNN 训练的完整技术方案——从 METIS 图分区、Local Graph/Feature Store 分布式存储,到基于 RPC 的分布式邻居采样与基于 DDP 的模型训练。读完本文,你将掌握分布式训练的全链路架构、核心类的源码级实现原理,以及如何在自己的集群中部署这套训练管线。

一、为什么需要分布式训练:从单机瓶颈到集群扩展

真实世界中的图往往包含数十亿个节点(如社交网络、推荐系统、知识图谱),远超单台机器内存的承载能力。此时,分布式 GNN 训练便成为必然选择:将大图切分为若干分区并分配到 CPU 集群的各个节点上,借助 PyTorch 的 Distributed Data Parallel(DDP)能力,对全量数据一次性部署同步模型训练。

PyG 从 2.5 版本起内置了第一个官方自研的分布式训练方案torch_geometric.distributed(由 Intel 与 Kumo AI 的工程师共同贡献),其核心设计是:

  • 使用RPC(Remote Procedure Calls)完成跨节点的邻居采样与远程特征获取;
  • 使用DDP完成数据并行的模型训练;
  • 该实现不依赖任何额外的第三方包,在默认 PyG 依赖栈之上即可运行。

⚠️ 注意:torch_geometric.distributed自 2.7.0 起已被标记为deprecated(弃用),官方不再维护。其后续替代方案为官方文档中的 分布式训练教程(multi_gpu_vanilla、multi_node_multi_gpu_vanilla)以及 NVIDIA cuGraph-GNN 方案。但理解这套经典架构,仍是掌握大规模图训练原理的最佳切入点。

二、六大关键优势:这套方案解决什么问题

官方教程 distributed_pyg.rst 总结了该模块的六项核心优势:

  1. 均衡图分区(Balanced Graph Partitioning):基于 METIS 算法,最小化跨计算节点采样子图时的通信开销;
  2. DDP + RPC 双通道:模型训练走 DDP,远程采样与特征获取走 RPC(基于 TCP/IP 协议与 gloo 通信后端),实现"各节点数据分区不同"的数据并行;
  3. 自定义 GraphStore / FeatureStore 接口:通过torch_geometric.data.GraphStoretorch_geometric.data.FeatureStore的灵活抽象,分别分发大图的结构信息与特征存储;
  4. 分布式邻居采样:可同时在本地区域与远程分区采样(经 RPC 通信通道),单机采样的全部高级功能均适用——异构采样、边级(link-level)采样、时序(temporal)采样等;
  5. 分布式数据加载器DistNeighborLoader等提供高层抽象,统一管理 sampler 进程,与标准 PyG 数据加载器无缝集成;
  6. 异步化:在 PyTorch RPC 之上引入 Pythonasyncio库做异步处理,进一步提升系统响应性与整体性能。

三、整体架构:六大组件的职责划分

torch_geometric.distributed对外导出的核心类见 distributed/init.py,共 8 个:

组件职责
Partitioner将图切分为多份,使每个节点只需在内存中加载本地数据
LocalGraphStore存储每个分区内的图拓扑结构(edge index),维护本地/全局 ID 映射
LocalFeatureStore存储节点级与边级特征,支持本地与远程特征的 put/get
DistNeighborSampler分布式采样算法:本地 + 远程采样,并基于 PyTorch RPC 合并结果
DistLoader分布式加载器的基类,封装 sampler 进程的初始化与清理
DistNeighborLoader管理分布式邻居采样与特征获取流程,将结果组装为标准 PyG Data 格式
DistLinkNeighborLoader边级(链接预测)场景的分布式采样加载器
DistContext保存当前进程的 rank、world_size 等分布式上下文信息

其工作流程可以概括为一条主线:Partitioner 切图 → LocalGraphStore/LocalFeatureStore 装载分区 → DistNeighborSampler 分布式采样 → DistLoader 系列加载器产出 mini-batch → 模型在 DDP 下训练

四、图分区:METIS 切分与 Halo 节点

4.1 分区原理

分布式训练的第一步,是把图拆成若干子图,让集群各节点只加载自己那一份。分区建立在pyg-libMETIS 算法的实现之上(pyg_lib.partition.metis),即便对大规模图也能高效完成切分。

关键设计要点:

  • 输入要求:METIS 要求输入为无向、同构图,因此Partitioner会先执行必要预处理(异构数据先to_homogeneous()),再对异构数据对象做正确的分布与索引重建;
  • 均衡目标:默认情况下,METIS 在最小化分区之间边数的同时,尽量平衡每个分区中各类节点的数量,保证采样器可"本地计算"、无需跨节点通信;
  • Halo 节点:每个节点获得唯一归属分区,而落入其他分区的 1 跳邻居(halo nodes)会被复制到本分区。Halo 节点保证单节点在单层中的邻居采样可以完全本地完成;
  • 非确定性警告:METIS 分区具有非确定性,不同迭代结果可能不同。但所有计算节点必须访问同一份分区数据,因此官方建议:在一个节点上生成分区后,将数据复制到集群所有成员,或把分区目录放到共享存储中。

4.2 分区目录结构

以 ogbn-products 同构图切分为两份为例,Partitioner 产出的目录结构如下:

partitions └─ obgn-products ├─ ogbn-products-partitions │ ├─ part_0 │ ├─ part_1 │ ├─ META.json │ ├─ node_map.pt │ └─ edge_map.pt ├─ ogbn-products-label │ └─ label.pt ├─ ogbn-products-test-partitions │ ├─ partition0.pt │ └─ partition1.pt └─ ogbn-products-train-partitions ├─ partition0.pt └─ partition1.pt

从 partition.py 的源码注释可以看到更细致的内部布局:每个part_{pid}/目录下包含graph.pt(含edge_idrowcolsize)、node_feats.pt(含global_ididfeats)、edge_feats.pt;异构图则将node_map/edge_map拆成按节点类型、边类型组织的子目录。顶层META.json记录num_partsnode_typesedge_typesis_heterois_sorted等元信息,供load_partition_info()读取(见 partition.py)。

4.3 Partitioner 的构造与能力

Partitioner的构造函数签名(partition.py):

Partitioner(data, num_parts, root, recursive=False)
  • dataDataHeteroData对象;
  • num_parts:分区数量(须大于 1);
  • root:分区数据集保存的根目录;
  • recursive:若为True,使用多级递归二分(multilevel recursive bisection)替代多级 k-way 分区(默认False)。

从源码看(partition.py),generate_partition()内部会:先对异构数据取time属性做时间信息保留 →to_homogeneous()→ 交给ClusterDatakeep_inter_cluster_edges=Truesparse_format='csc')执行 METIS → 依据node_perm/partptr/edge_perm重建全局↔本地 ID 映射 → 按列(目标节点)排序为 CSC 格式 → 逐分区保存graph.ptnode_feats.ptedge_feats.pt及映射文件。值得强调的是,Partitioner完整保留节点特征、边特征以及节点级/边级的时间属性is_node_level_time/is_edge_level_time),这为后续时序采样提供了数据基础。

五、分布式数据存储:LocalGraphStore 与 LocalFeatureStore

分区完成后,每个集群节点需要能高效地访问"本地分区 + 远程分区"的数据。方案是对 PyG 的GraphStoreFeatureStore远程接口实例化,并结合内建的 RPC 请求收发 API,构成互联的分布式数据存储。

5.1 LocalGraphStore:图拓扑容器

LocalGraphStore(local_graph_store.py)实现GraphStore接口,核心能力包括:

  • 只存储本分区内的本地图连接及其 halo 节点信息;
  • 远程连通性:通过节点/边的 "partition book"(分区 ID → 全局节点/边 ID 的映射)查询任意节点/边属于哪个分区(本地或全局);
  • 全局标识符:为节点与边维护全局 ID,保证跨分区映射一致。

从源码看,它内部维护_edge_index_edge_attr_edge_id三个字典,并提供三类构造方式:

  • from_data(...):从同构图数据构造(local_graph_store.py);
  • from_hetero_data(...):从异构HeteroData构造,按边类型分别存储(local_graph_store.py);
  • from_partition(root, pid):直接从分区文件加载(local_graph_store.py),内部调用load_partition_info()读取META.json与映射文件。

关键查询方法get_partition_ids_from_nids(ids, node_type)get_partition_ids_from_eids(eids, edge_type)(local_graph_store.py)返回节点/边所属的分区 ID,这是采样器判断"本地采样还是远程 RPC 采样"的依据。

5.2 LocalFeatureStore:节点/边特征存储

LocalFeatureStore同时承担节点级与边级特征存储,提供高效的put/get例程,负责训练过程中跨分区、跨机器的特征检索与更新:

  • 在本机管理的分区内,节点与边特征本地存储
  • 通过 RPC 请求实现远程特征查找,无缝应对采样结果跨分区的情形;
  • 维护节点与边的全局标识符,保证跨分区映射一致。

官方教程给出了一段使用LocalFeatureStore异步获取节点特征的内部示例:

import torch from torch_geometric.distributed import LocalFeatureStore from torch_geometric.distributed.event_loop import to_asyncio_future feature_store = LocalFeatureStore(...) async def get_node_features(): # Retrieve node features for specific node IDs: node_id = torch.tensor([1]) future = feature_store.lookup_features(node_id) return await to_asyncio_future(future)

可以看到,lookup_features()返回一个 future 对象,通过to_asyncio_future()转成 Python 异步对象后await即可拿到特征——这正是"RPC 异步化 + asyncio"设计在特征获取环节的直接体现。

六、分布式邻居采样:DistNeighborSampler

DistNeighborSampler专为分区存储在多个机器上的大图设计,解决分布式环境下邻居采样的挑战,保证 GNN 训练的扩展性与性能。

6.1 异步采样与异步特征收集

分布式邻居采样基于异步的torch.distributed.rpc调用实现:

  • 各机器独立、异步地从本地图分区选择邻居,无需等待其他机器完成采样;
  • 除采样外,特征收集同样是异步的;
  • 这种异步设计最大化并行度,显著加速训练。

6.2 可定制的采样策略

DistNeighborSampler对采样策略提供完全灵活的定制能力,包括:

  • 节点采样 vs. 边采样(node sampling / edge sampling);
  • 同构图 vs. 异构采样(homogeneous / heterogeneous sampling);
  • 时序采样 vs. 静态采样(temporal / static sampling)。

6.3 三步采样工作流

一批种子节点在被数据加载器交给模型forward之前,需要经历以下三个主要步骤:

  1. 分布式节点采样:分布式场景下,同一 batch 的种子节点可能分属不同分区,多个机器会同时采样,因此需要跨机器同步采样结果以获得下一层种子节点——这与单机采样有本质差异。本地分区的节点在本地采样;远程分区的节点由存储该分区的机器负责采样。采样逐层进行,采样出的节点又作为下一层的种子节点;
  2. 分布式特征查找:每个分区存储其内部节点/边的特征数组。若某台机器采样结果中包含不属于本分区的节点或边,它便向这些节点/边所属的远程服务器发起 RPC 请求获取特征;
  3. 数据转换:基于采样器输出与获取到的节点(或边)特征,构建 PyG 的DataHeteroData对象,该对象构成后续模型计算所用的 batch。

七、分布式数据加载:DistNeighborLoader 与 DistLinkNeighborLoader

DistNeighborLoaderDistLinkNeighborLoader提供了采样引擎之上的简单 API——它们在内部完整封装了 sampler 进程的初始化与清理。值得注意的是,这两个分布式加载器分别继承自标准单机加载器torch_geometric.loader.NodeLoadertorch_geometric.loader.LinkLoader,因此训练脚本中的用法与单机几乎一致(见 dist_neighbor_loader.py)。

Batch 生成与单机略有差异:(本地 + 远程)特征获取被内联进 sampler 中,而不是拆成"采样 → 特征获取"两步,以此限制 RPC 数量。由于所有 sampler 子进程之间异步处理,sampler 最终把输出放入一个torch.multiprocessing.Queue,由主进程消费。

DistNeighborLoader的关键参数(dist_neighbor_loader.py):

参数说明默认值
data(LocalFeatureStore, LocalGraphStore)元组
num_neighbors每层每节点采样邻居数;-1表示取全部;异构图可为按边类型区分的字典
master_addr/master_port分布式加载器 RPC 通信的主节点地址与端口
current_ctx当前进程的DistContext上下文(rank、world_size 等)
concurrencyRPC 并发度,即异步处理队列的最大尺寸1
num_rpc_threadsRPC 线程数16
async_sampling是否启用异步采样(启用时使用multiprocessing.Queue作为通道)True
replace/subgraph_type/disjoint/temporal_strategy/time_attrNeighborLoader一致的采样语义

从构造函数源码可见,若不显式传入dist_sampler,加载器会自动构建DistNeighborSampler并同时初始化DistLoader(RPC 通信管理)与NodeLoader(batch 生成与转换)两条继承链,并通过transform_sampler_output=channel_get从队列取回采样结果。

八、通信层:DDP 与 RPC 的分工协作

该方案同时使用两种torch.distributed通信技术:

  • torch.distributed.rpc:负责远程采样调用与分布式特征检索;
  • torch.distributed.ddp:负责数据并行的模型训练。

官方选择torch.distributed.rpc而非 gRPC 等替代方案,关键原因是:PyTorch RPC 原生理解张量类型数据,无需像其他 RPC 方案那样先把 JSON 或用户数据序列化/数字化为张量,从而避免了额外的序列化与数字化开销。

8.1 DDP 组初始化

DDP 组在主训练脚本中以标准方式初始化:

torch.distributed.init_process_group( backend='gloo', rank=current_ctx.rank, world_size=current_ctx.world_size, init_method=f'tcp://{master_addr}:{ddp_port}', )

提示:基于 CPU 的采样推荐使用gloo通信后端。

8.2 RPC 组初始化

RPC 组初始化更为复杂,因为它发生在每个 sampler 子进程中,通过数据加载器的worker_init_fn(由 PyTorch 在 worker 进程初始化阶段直接调用)完成。该函数依次:

  1. 为每个 worker 定义分布式上下文,并分配 group 与 rank;
  2. 初始化自己的分布式邻居采样器;
  3. 在 RPC 组中注册新成员——这条 RPC 连接在子进程存续期间保持打开。

此外,实现还借助 Python 标准库的atexit模块注册进程终止时的额外清理行为,保证 RPC 连接与队列资源的正确释放。

九、性能基准:ogbn-products 上的扩展性表现

官方教程在 PyTorch 2.1 上给出了 GraphSAGE 模型在 ogbn-products 数据集上的扩展性基准(下表为不同分区数与 batch size 下单个 epoch 的训练耗时):

#Partitionsbatch_size=1024batch_size=4096batch_size=8192
198s47s38s
245s30s24s
438s21s16s
829s14s10s
1622s13s9s

基准软硬件环境(来自官方教程原文):

  • 硬件:2x Intel(R) Xeon(R) Platinum 8360Y CPU @ 2.40GHz,36 核,HT/Turbo 开启,NUMA 2,总内存 256GB(16x16GB DDR4 3200 MT/s),2x Ethernet Controller X710(10GbE SFP+),1x ConnectX-6,Rocky Linux 8.8;
  • 软件:Python 3.9、PyTorch 2.1、PyG 2.5、pyg-lib 0.4.0。

可以看到:在固定 batch size 下,分区数从 1 增加到 16 时训练耗时近似线性下降(如batch_size=1024下从 98s 降至 22s,约 4.5 倍加速);更大的 batch size 在同等分区数下始终更快,体现了该方案良好的横向扩展能力。

十、现状与迁移建议

需要再次强调:torch_geometric.distributed自 PyG 2.7.0 起已弃用(模块文档见 modules/distributed.rst,弃用声明同样体现在 examples/distributed/pyg/README.md 中),代码中也会在构造时打印弃用警告(见 dist_neighbor_loader.py)。如果你正在启动新的分布式训练项目,官方推荐的路径是:

  1. 阅读 分布式训练教程索引,其中包含单机多卡(multi_gpu_vanilla)与多机多卡(multi_node_multi_gpu_vanilla)两篇基于原生 PyTorch DDP 的入门教程;
  2. 对 NVIDIA GPU 用户,官方推荐接入 cuGraph-GNN 以获得可扩展的分布式 GNN 训练能力;
  3. 即便如此,本文所述的"分区 → 分布式存储 → 分布式采样 → DDP 训练"四层架构思想与 RPC 通信细节,仍然是大规模图训练领域最具参考价值的工程实践之一,理解它对你评估任何分布式 GNN 框架都大有裨益。

【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric

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

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

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

立即咨询