摘要
当模型大到一张 GPU 装不下、或者数据多到一张卡跑不动的时候,PyTorch 提供了两套分布式训练方案来解决这个问题:DDP(DistributedDataParallel)和 FSDP(FullyShardedDataParallel)。它们都不是独立的框架,而是 PyTorch 内置的分布式训练技术,分别代表了"数据并行"和"参数分片"两种完全不同的思路,也是目前几乎所有大模型训练岗位都会要求熟悉的基本功。
背景与问题
单卡训练有两个绕不开的瓶颈:一是数据量太大,一张卡训练太慢;二是模型本身太大,参数、梯度、优化器状态加起来根本装不进一张卡的显存。前者靠"多卡同时算、结果汇总"就能解决,后者则必须想办法把模型本身也拆开分布到多张卡上。
PyTorch 针对这两类问题给出的答案分别是 DDP 和 FSDP。它们都属于 torch.distributed 生态下的并行策略 API,底层通信默认走 NCCL——如果说 NCCL 解决的是"卡与卡之间怎么说话",DDP 和 FSDP 解决的就是"该让每张卡负责什么、什么时候需要说话"。
核心思路与优势
DDP:每张卡一份完整模型,只拆数据
DDP 的思路很直接:每个进程(通常对应一张 GPU)都持有一份完整的模型副本,各自吃一部分数据、独立完成前向和反向传播,算出各自的梯度。初始化时,DDP 会把 0 号进程的 state_dict()(参数和 buffer)广播给所有其他进程,保证所有副本从同一个起点出发;之后每一步反向传播结束,各进程的梯度会通过 allreduce 求和平均,再各自更新参数,从而保持所有模型副本始终一致。
DDP 有个值得了解的实现细节:梯度分桶(bucketing)。它不会等所有梯度都算完才去同步,而是把参数梯度组织成若干个 bucket,大致按模型参数的逆序分配(因为反向传播时梯度也大致按这个逆序变为 ready),一个 bucket 里的梯度全部就绪就立刻异步发起一次 allreduce,这样通信和计算可以尽量重叠、少等待。所有进程必须严格按相同顺序执行这些 allreduce,顺序不一致会直接导致结果错误甚至训练卡死,这也是 DDP 内部会用固定的 bucket 顺序而不是"谁先好谁先发"的原因。
DDP 的优势是实现简单、开销小、工程上非常成熟;代价是每张卡都要装下完整的模型、梯度和优化器状态,模型越大,这份冗余就越吃显存——一旦模型本身超出单卡显存容量,DDP 就无能为力了。
FSDP:把模型本身也拆开分片
FSDP 解决的正是 DDP 解决不了的那个问题。它把模型参数、梯度、优化器状态都切分成片,分别存放在不同 GPU 上,平时每张卡只保有属于自己的那一小份;只有真正要计算某一层时,才通过 all-gather 临时把这一层的完整参数收集齐,算完立刻释放(reshard),显存不会被长期占用。
PyTorch 目前推荐的是 FSDP2,核心 API 是 fully_shard,用法是先对模型的每个子层分别包一层,再对整个根模型包一次:
from torch.distributed.fsdp import fully_shard
model = Transformer()
for layer in model.layers:
fully_shard(layer)
fully_shard(model)
和 FSDP1 相比,FSDP2 的分片方式也不一样:FSDP1 是把一组 tensor 拉平、拼接后再整体切分;FSDP2 则是按 dim-0 对每一个参数单独切分(torch.chunk(dim=0)),分片粒度更细,也更直观——冻结参数的限制更少,还支持不需要额外通信的分片 state dict。
FSDP 的显存效率明显好于 DDP,能让原本单卡装不下的大模型(十亿参数级别往上)训练成为可能,或者在同样显存下跑更大的 batch size;代价是频繁的 all-gather / reduce-scatter 带来了更高的通信开销,分片粒度越细,省的显存越多,但通信次数更多、单次传输更小,通信延迟和调度开销也更高,需要根据实际硬件和网络带宽做权衡。
面向人群
- 模型训练/算法工程师:日常做多卡训练,需要清楚什么时候用 DDP 就够、什么时候必须上 FSDP,否则要么白白浪费显存,要么直接 OOM。
- AI Infra / MLSys 方向的工程师:设计训练集群的并行策略、排查显存和通信瓶颈,DDP 和 FSDP 的内部机制是绕不开的基础知识。
- 准备大模型相关岗位面试的求职者:DDP、FSDP 和 NCCL 通常被放在同一个知识层次上考察——从底层通信到上层并行策略的完整技术栈,是越来越多岗位描述里的加分项甚至硬性要求。
实践步骤
第一步:用 DDP 训练一个能塞进单卡显存的模型
初始化通信组、绑定设备、用 DistributedDataParallel 包裹模型:
import os
import torch
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
local_rank = int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(local_rank)
dist.init_process_group(backend="nccl")
model = MyModel().to(local_rank)
model = DDP(model, device_ids=[local_rank])
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
配合 DistributedSampler 让每张卡只读取数据集的一部分,训练循环和单卡几乎没有区别:
for x, y in dataloader:
optimizer.zero_grad()
loss = loss_fn(model(x), y)
loss.backward()
optimizer.step()
反向传播里的梯度同步是 DDP 自动完成的,业务代码基本感知不到通信过程。用 torchrun --nnodes=1 --nproc_per_node=4 train.py 这样的命令启动即可。
第二步:模型太大装不下时,换成 FSDP2
把 DDP 换成 fully_shard,其余训练逻辑基本不用大改:
import os
import torch
import torch.distributed as dist
from torch.distributed.fsdp import fully_shard
local_rank = int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(local_rank)
dist.init_process_group(backend="nccl")
model = Transformer(vocab_size=32000, n_layers=32)
for layer in model.layers:
fully_shard(layer)
fully_shard(model)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
注意优化器必须在 fully_shard 之后创建,因为这时模型参数已经变成了分片后的 DTensor。训练循环写法和普通 PyTorch 训练几乎一样:
for x, y in dataloader:
optimizer.zero_grad()
loss = loss_fn(model(x), y)
loss.backward()
optimizer.step()
梯度裁剪也可以直接作用在这些 DTensor 参数上,不需要特殊处理,位置和单卡训练一样,放在 loss.backward() 之后、optimizer.step() 之前:
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
第三步:开启混合精度,进一步压缩显存和通信量
FSDP2 通过 MixedPrecisionPolicy 在包装时指定精度策略。param_dtype 决定 all-gather 出来的完整参数用什么精度做前向和反向计算,reduce_dtype 决定梯度规约(reduce-scatter)用什么精度。常见配置是计算走 bfloat16、梯度规约转回 float32 保证数值稳定;注意分片存放的那份参数仍然保持原始精度,优化器也是在原始精度的分片参数上做更新:
from torch.distributed.fsdp import fully_shard, MixedPrecisionPolicy
fsdp_kwargs = {
"mp_policy": MixedPrecisionPolicy(
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
)
}
for layer in model.layers:
fully_shard(layer, **fsdp_kwargs)
fully_shard(model, **fsdp_kwargs)
第四步:性能调优的几个关键点
- DDP 侧:
bucket_cap_mb控制梯度分桶大小,桶太小同步次数多、开销大,桶太大又会让计算和通信的重叠变差;find_unused_parameters=True能支持部分参数不参与本次反向传播的场景,但会引入额外开销,非必要不要开启。配合torch.compile()使用时,要先用 DDP 包裹模型再调用torch.compile()(即torch.compile(DDP(model, ...))),这样才能触发 DDPOptimizer,在 allreduce 的 bucket 边界处切分前向计算图,把通信和计算的重叠恢复回来;顺序写反就得不到这个优化。 - FSDP 侧:FSDP2 默认提供隐式 prefetching,让 all-gather 和计算尽量重叠,建议先用默认配置跑起来看性能再决定是否手动调整;分片粒度(对每一层单独
fully_shardvs 更粗粒度地包装)直接影响显存和通信的权衡,层数多、单层不大的模型(比如 Transformer)通常按层分片效果最好。
第五步:判断该用哪一个
简单判断标准:模型能完整放进单卡显存,就优先用 DDP,实现简单、开销低、工程成熟;模型显存需求超出单卡容量(比如十亿参数级别往上的大模型预训练/微调),才需要上 FSDP 用分片换空间。像 TorchTitan 这样的大模型预训练平台已经把 FSDP2 作为默认的并行方案,并在此基础上叠加张量并行、流水线并行等更复杂的策略来支撑更大规模的训练——但这些都建立在先理解 DDP 和 FSDP 各自的适用边界之上。