ShardedDataParallel是Fairscale提供的模型分片并行方案,不能直接替代DDP;它将参数、梯度和优化器状态分片到各卡以节省显存,但需换用OSS优化器、适配混合精度、处理激活内存,并已逐步被PyTorch官方FSDP取代。

ShardedDataParallel 是什么,它真能替代 DDP 吗?
ShardedDataParallel 不是 PyTorch 官方 API,而是 Fairscale 库提供的一个封装(现已逐步迁入 torch.distributed.fsdp)。它和 torch.nn.parallel.DistributedDataParallel(DDP)目标不同:DDP 复制全部模型参数到每个 GPU,只做梯度同步;而 ShardedDataParallel 把模型参数、优化器状态、梯度按层或按参数分片(shard),让每个 GPU 只存一部分——这是为显存受限的大模型训练设计的。
- 它不能直接“替代” DDP,因为不兼容所有模型结构(比如带复杂跨设备 control flow 的自定义模块)
- 如果你用的是 Hugging Face
transformers模型 + 单机多卡(如 4×A10G),且显存总不够装下OPT-6.7B这类模型,ShardedDataParallel才值得考虑 - 当前主流已转向
FSDP(torch.distributed.fsdp),ShardedDataParallel在 Fairscale v0.4+ 中已被标记为 deprecated
如何把现有 DDP 训练脚本改造成 ShardedDataParallel?
改造不是“换一行 import”,核心在于三处解耦:参数分片时机、梯度归约方式、优化器状态管理。你得手动替换 DDP 包装逻辑,并确保模型 forward 不触发跨分片张量隐式通信。
- 替换前确认已安装兼容版本:
pip install fairscale==0.4.13(太高会报AttributeError: 'ShardedDataParallel' object has no attribute 'no_sync') - 不要再用
torch.nn.parallel.DistributedDataParallel包装模型,改用:from fairscale.nn.data_parallel import ShardedDataParallel as ShardedDDP from fairscale.optim.oss import OSS </li></ul><p>model = MyLargeModel() sharded_model = ShardedDDP(model, optimizer)</p>
- 优化器必须换成OSS(Optimizer State Sharding),否则优化器状态仍占满每张卡显存 -ShardedDDP不支持torch.cuda.amp.autocast的某些嵌套模式,需把autocast移到 forward 内部,或改用ShardedDDP(..., mixed_precision=True)(但仅限部分 dtype)常见报错:RuntimeError: expected scalar type Half but found Float
这个错误通常出现在启用混合精度但分片未对齐时。根本原因是:不同 GPU 上的参数分片 dtype 不一致,或
OSS优化器在 step 前没统一 cast。- 确保初始化模型后立刻调用
model.half(),而不是等包装成ShardedDDP后再转 -
OSS构造时显式传参:OSS(params=model.parameters(), optim=torch.optim.AdamW, foreach=False)(foreach=True在分片下易出 dtype 错误) - 检查是否手动调用了
loss.backward()后又用了optimizer.step()——ShardedDDP要求用sharded_optimizer.step(),且内部已包含 all-gather,重复调用会导致梯度乱序
为什么训着训着 OOM 了,明明参数分片了?
分片只管参数和优化器状态,不管中间激活(activations)。大 batch 或长序列仍可能撑爆单卡显存。
立即学习“Python免费学习笔记(深入)”;
- 必须配合梯度检查点(
torch.utils.checkpoint.checkpoint):对 Transformer 层级封装,避免保存全部 attention key/value - 关闭
ShardedDDP的bucket_cap_mb自动设置(默认 25),设为更小值如bucket_cap_mb=10,减少临时通信 buffer 占用 - 注意
ShardedDDP的reduce_buffer_size默认为 0,意味着禁用梯度 bucketing;若想节省显存,需手动设非零值,但会增加通信次数
Fairscale 的分片逻辑依赖参数名顺序和初始化一致性,哪怕 model.eval() 和 model.train() 切换时有轻微结构变化,都可能导致某卡分片数不匹配——这种问题不会立刻报错,而是在第 N 个 step 后突然
NCCL timeout。上线前务必跑至少 10 个 step 的 full-graph 验证。 - 确保初始化模型后立刻调用


















