PyTorch 2.x分布式训练是深度学习框架落地大规模模型的基本功。单卡放不下模型时,DDP和FSDP是两条主流路径:DDP做数据并行,每张卡持有一份完整模型副本;FSDP把参数、梯度、优化器状态分片,用通信换显存。两者解决的问题不同,选型依据也不同。
DDP数据并行原理与训练代码
DDP(DistributedDataParallel)在每次反向传播时通过AllReduce同步各卡梯度,保证所有副本参数一致。实现要点只有三步:初始化进程组、包装模型、改造训练循环。
import torch
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
def train(local_rank, world_size):
dist.init_process_group("nccl", rank=local_rank, world_size=world_size)
torch.cuda.set_device(local_rank)
model = DDP(model, device_ids=[local_rank])
# 训练循环与单卡一致,梯度同步由DDP自动完成
for x, y in dataloader:
loss = criterion(model(x), y)
loss.backward()
optimizer.step()
启动命令用torchrun:torchrun –nproc_per_node=8 train.py。数据切片用DistributedSampler,保证每个进程看到互不重叠的样本。
FSDP参数分片原理与配置
FSDP(FullyShardedDataParallel)借鉴ZeRO-3思路,把参数、梯度、优化器状态按卡分片。前向时用all-gather收集完整参数,反向时reduce-scatter聚合梯度,每卡显存占用从O(N)降到O(N/卡数)。
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy
policy = transformer_auto_wrap_policy(transformer_layer_cls={LlamaDecoderLayer})
model = FSDP(model, auto_wrap_policy=policy, sharding_strategy=ShardingStrategy.FULL_SHARD)
FSDP按transformer层自动切分并逐层展开,通信粒度为单层,显存峰值明显下降。显存仍不够时再叠加激活检查点(activation checkpointing)。
DDP与FSDP显存占用和通信开销对比
实测数据(7B模型、8卡A100 80G):DDP每卡需要模型权重、梯度、优化器状态完整副本,约16GB显存;FSDP FULL_SHARD每卡降到约4-5GB。代价是前向的AllGather和反向的ReduceScatter通信量增加,小batch下吞吐会低于DDP。
选型依据一句话:单卡能装下模型副本时DDP吞吐更高;模型超过单卡显存时FSDP是唯一可选路径。训练吞吐优先选DDP,显存受限场景选FSDP。
分布式训练调参要点与常见问题
梯度累积(gradient_accumulation_steps)等效放大batch size,不额外占显存;学习率随batch size线性缩放;多机场景确认NCCL_IB_TC与NCCL_SOCKET_IFNAME指向正确网卡,否则报NCCL timeout。
常见问题排查:NCCL超时检查网络与IB配置;OOM优先开FSDP激活检查点;梯度不一致确认所有进程的backward调用次数一致,跳过batch的场景必须用no_sync()包裹。
原创文章,作者:小编,如若转载,请注明出处:https://www.yunthe.com/pytorch2x-fen-bu-shi-xun-lian-shi-zhan-ddp-yu-fsdp-bing/