Skip to content

分布式训练概述与通信原语

分布式训练是训练大语言模型的核心基础设施。理解通信原语是掌握所有分布式训练方案的基础。

分布式训练架构

分布式训练分类

类型切分对象通信需求适用规模
数据并行 (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 用于跨节点扩展。

相关资源

最近更新