分布式训练概述与通信原语
分布式训练是训练大语言模型的核心基础设施。理解通信原语是掌握所有分布式训练方案的基础。
分布式训练分类
| 类型 | 切分对象 | 通信需求 | 适用规模 |
|---|---|---|---|
| 数据并行 (DP) | 数据 | 梯度 AllReduce | 单节点-多节点 |
| 张量并行 (TP) | 权重 | 激活 AllReduce | 单节点 (NVLink) |
| 流水线并行 (PP) | 层 | 点对点通信 | 跨节点 |
| 序列并行 (SP) | 序列 | 激活 AllGather | 长上下文 |
集合通信原语
Broadcast
将一个进程的数据广播到所有进程:
python
import torch.distributed as dist
# Rank 0 的数据广播到所有 Rank
data = torch.tensor([1.0, 2.0, 3.0]) if dist.get_rank() == 0 else torch.zeros(3)
dist.broadcast(data, src=0)Reduce
将所有进程的数据归约到一个进程:
python
# 所有 Rank 的数据求和到 Rank 0
dist.reduce(data, dst=0, op=dist.ReduceOp.SUM)AllReduce
所有进程都获得归约结果(最常用的原语):
python
# 所有 Rank 同时获得全局梯度求和
dist.all_reduce(gradients, op=dist.ReduceOp.AVG)AllGather / ReduceScatter
python
# AllGather: 收集所有 Rank 的数据
output = [torch.zeros_like(data) for _ in range(dist.get_world_size())]
dist.all_gather(output, data)
# ReduceScatter: 归约后按 Rank 分发
output = torch.zeros(data.shape[0] // dist.get_world_size())
dist.reduce_scatter(output, [data], op=dist.ReduceOp.SUM)通信与计算重叠
python
# 梯度通信与计算重叠
for name, param in model.named_parameters():
if param.grad is not None:
# 异步通信
handle = dist.all_reduce(param.grad, async_op=True)
# 继续计算其他层
...
# 等待通信完成
handle.wait()通信量估算
数据并行的通信量 = 模型参数量 × 2(梯度 + AllReduce)。7B 模型 FP16 的单步通信量约 28GB。NVLink 600GB/s 带宽下约 47ms,InfiniBand 400Gbps 下约 560ms。
通信瓶颈
跨节点训练的通信延迟是主要瓶颈。建议 TP 限制在单节点内(NVLink),DP/PP 用于跨节点扩展。