执行摘要:很多人以为分布式训练调不动,是卡不够。真实情况往往是显存账没算清——你以为 70B 模型需要 1.12 TB 是在买显存,其实其中 75% 是三份完全可以切开的冗余状态。本文沿着一条真实的工程决策链展开:先把显存账本逐项拆开,再讲数据/张量/流水三种并行各自解决的是哪个维度的瓶颈,然后深入 ZeRO 三级分片与 PyTorch FSDP 的 FlatParameter 编排,接着算清激活重计算的算力-显存兑换率,最后落到最容易被忽略的一层——你的 GPU 有 30% 时间在等 NCCL。全程给出可直接套用的配置与代码。
一、显存账本:先算清楚再谈并行
混合精度(bf16)+ AdamW 是当前预训练的默认配置。此时每个参数的常驻显存并不是 2 字节,而是:
| 组件 | 精度 | 每参数字节 | 说明 |
|---|---|---|---|
| 前向/反向实际使用的参数 | bf16 | 2 | 真正参与 matmul 的副本 |
| 梯度 | bf16 | 2 | 反向产出,reduce-scatter 前常驻 |
| 主权重 master weights | fp32 | 4 | 优化器更新用,bf16 直接累加会掉精度 |
| 一阶动量 m | fp32 | 4 | Adam |
| 二阶动量 v | fp32 | 4 | Adam |
| 合计 | 16 |
一个 70B 模型的状态显存就是 70e9 × 16 ≈ 1.12 TB,而一张 A100/H100 只有 80 GB。14 张卡什么都不干,光放状态就满了——这就是"显存墙"的本体。
关键在于:这 16 字节里,只有 4 字节(bf16 参数 + bf16 梯度)是任何时刻都真正需要的,其余 12 字节是优化器的历史状态,而 Adam 对每个参数的更新是完全独立、逐元素的。这个逐元素性质,正是 ZeRO 能把它切开的数学前提。
二、三种并行各自解决什么问题
这是最常被混淆的地方。三种并行不是"三个选项",而是正交的三个维度:
- DP(数据并行):每卡一份完整模型,切的是 batch。解决吞吐,但显存零节省。反向后 all-reduce 梯度。
- TP(张量并行):切层内的权重矩阵(Megatron 的列切 + 行切配对)。解决"单层放不下"。代价是每层 4 次 all-reduce(前向 2 次、反向 2 次),通信量正比于
batch × seq × hidden,因此必须锁在节点内走 NVLink,跨节点用 TP 等于自杀。 - PP(流水并行):切层与层之间,只传边界激活,通信量最小。代价是流水线气泡,1F1B 调度下气泡占比约
(p-1)/m(p 为阶段数,m 为微批数),工程上取m ≥ 4p把它压到 25% 以内。
组合顺序是硬经验:TP 在节点内(8 卡 NVLink)、PP 跨节点(IB)、DP 放在最外层。原因就是通信量与带宽的匹配——把高频大流量的 all-reduce 留在 NVLink 上,把低频的激活点对点丢给网络。
另外别忘了第四个维度:长上下文场景下激活随序列长度线性膨胀,Sequence/Context Parallel + Ring Attention 把序列维切开,已经成为 128K 以上训练的标配。
三、ZeRO 三级分片:把冗余状态切碎
ZeRO 的思路极其朴素:既然优化器更新是逐元素的,那每张卡就只保存 1/N 份状态,用到别人的参数时再临时要。
| 策略 | 分片对象 | 单卡每参数字节 | 70B 在 64 卡上的单卡状态 | 相对 DP 通信量 |
|---|---|---|---|---|
| DDP | 无 | 16 | 1120 GB(放不下) | 2Φ |
| ZeRO-1 | 优化器状态 | 4 + 12/N | ≈ 293 GB(放不下) | 2Φ |
| ZeRO-2 | + 梯度 | 2 + 14/N | ≈ 155 GB(放不下) | 2Φ |
| ZeRO-3 | + 参数 | 16/N | ≈ 17.5 GB(可行) | 3Φ(1.5×) |
注意最后两列的取舍:ZeRO-3 用 1.5 倍通信量换来了从 1120 GB 到 17.5 GB 的降维打击。多出来的 1Φ 是在前向/反向时按需 all-gather 参数——每次用完立刻释放(param.release()),所以参数只在被用到的瞬间完整存在。
FSDP 的关键抽象:FlatParameter
PyTorch FSDP 里最重要的设计不是分片本身,而是 FlatParameter:把一个 wrap unit 内的所有参数 flatten 拼成一个 1D 大 buffer。这样一来:
- 一次 all-gather / reduce-scatter 就是一次大块集合通信,而不是几百个小 tensor 的碎片通信(每次都有 µs 级启动开销);
- 显存分配是整块的,避免反复 alloc/free 造成的碎片;
- bucket 切分与通信-计算重叠有了天然的调度单位。
所以 auto_wrap_policy 不是性能调优选项,而是正确性的一部分:按整个 TransformerBlock 包裹,一个 block 是一个 FSDP unit;如果偷懒只包顶层模型,你就得到一个巨大的 FlatParameter——前向第一层就得 all-gather 整个模型,显存与 ZeRO-3 的初衷彻底背离。
from functools import partial
from torch.distributed.fsdp import (
FullyShardedDataParallel as FSDP,
ShardingStrategy, MixedPrecision, BackwardPrefetch,
)
from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy
# reduce_dtype 用 fp32:梯度归约是累加操作,bf16 累加会引入显著误差
mp = MixedPrecision(
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
buffer_dtype=torch.bfloat16,
)
wrap = partial(
transformer_auto_wrap_policy,
transformer_layer_cls={TransformerBlock}, # 每个 block 一个 FSDP unit
)
model = FSDP(
model,
sharding_strategy=ShardingStrategy.FULL_SHARD, # ZeRO-3
mixed_precision=mp,
auto_wrap_policy=wrap,
backward_prefetch=BackwardPrefetch.BACKWARD_PRE, # 反向时预取下一层参数
forward_prefetch=True, # 前向时预取下一层
limit_all_gathers=True, # 抑制并发 all-gather 的峰值显存
use_orig_params=True, # 兼容非 FSDP 的算子/优化器
device_id=torch.cuda.current_device(),
)
四、激活重计算:用 33% 算力换 O(L)→O(1) 的显存
状态显存降下去之后,下一个瓶颈是激活。Transformer 每层需要为反向保留的中间结果大约是 10~20 × batch × seq × hidden 字节(attention 的 softmax 中间态、MLP 的激活输入、各种 norm 的输出……)。序列越长,这一项越主导。
完全重计算的做法是:只保存每层的输入,反向时重新跑一遍前向。前向 : 反向的算力比约为 1 : 2,所以多跑一次前向的代价是 +33% 算力,换来的激活显存从 O(L) 降到 O(1)(分段策略下是 O(√L))。
但这 33% 是可以讨价还价的。选择性重计算(selective activation checkpointing) 的洞察是:不同算子的"显存/算力"兑换率天差地别——
- attention 内部:中间激活巨大(softmax 的
seq²项),但重算几乎全是便宜的 element-wise 与小张量 GEMM; - MLP 内部:中间激活相对小,但重算要跑两个大 GEMM,代价高昂。
所以正确做法是只 checkpoint attention,保留 MLP 的中间激活:能拿到大部分显存收益,而重算代价从 33% 降到个位数百分比。
from torch.utils.checkpoint import checkpoint
def selective_block(layer, x, attn_mask):
def attn_part(x): # 放进重计算区域
return layer.attn(layer.norm1(x), attn_mask)
# use_reentrant=False:避免 RNG 状态陷阱与嵌套 checkpoint 问题
h = x + checkpoint(attn_part, x, use_reentrant=False)
h = h + layer.mlp(layer.norm2(h)) # MLP 不重算,中间激活保留
return h
这里有个真实的坑:use_reentrant=True(旧默认)会在反向时重跑前向并复用 RNG 状态,一旦前向里有 dropout 之类依赖 RNG 的算子,或者 checkpoint 嵌套,结果就会静默出错。生产环境一律显式写 use_reentrant=False。
五、通信-计算重叠:别让 GPU 等 NCCL
ZeRO-3 / FSDP 在真实训练里最常见的 profile 形态不是"矩阵乘占满",而是大段 gap 里 GPU 在等集合通信。这一层的优化往往比换更大的模型结构收益更高:
1)预取(prefetch)。 forward_prefetch=True 让 FSDP 在计算第 i 层时提前发起第 i+1 层的 all-gather;BACKWARD_PRE 在反向开始前预取下一层参数。代价很直接:峰值显存会多驻留约一层的完整参数——这是显存换时间的显式交易,显存紧张时关掉它常常是唯一正确的选择。
2)bucket 化 + 梯度累积的 no_sync。 通信是按 bucket(默认 25 MB)攒够再发的,避免小消息的延迟主导。更重要的是梯度累积场景:中间步骤根本不需要同步梯度,只在最后一步做一次 reduce-scatter,把 N 步通信压缩成 1 步。
import contextlib
for step, batch in enumerate(loader):
is_accum = (step + 1) % ACCUM != 0
ctx = model.no_sync() if is_accum else contextlib.nullcontext()
with ctx: # 累积步:只写本地梯度,不做 reduce-scatter
loss = model(batch).loss
(loss / ACCUM).backward()
if not is_accum:
optimizer.step() # 这一步才触发梯度同步
optimizer.zero_grad(set_to_none=True)
3)流与显存的隐性竞争。 NCCL 跑在独立 CUDA stream 上,但预取 buffer 与激活共享同一块显存池。预取越激进、并发 all-gather 越多,越容易触发 cudaMalloc 的同步式回收,把重叠收益全部吃回去。limit_all_gathers=True 就是为此存在的:宁可稍微串行,也不要在峰值点 OOM 后整轮重跑。
六、工程取舍对照
| 症状 | 真正的瓶颈 | 该动的旋钮 |
|---|---|---|
| 单卡连模型状态都放不下 | 优化器/梯度/参数冗余 | ZeRO-3 / FSDP FULL_SHARD |
| 单层权重或中间张量放不下 | 层内维度 | TP(限节点内,走 NVLink) |
| 层数太多、状态切完仍放不下 | 层间维度 | PP + 1F1B,微批数 ≥ 4×阶段数 |
| batch 开不大,吞吐上不去 | 激活显存 | 选择性重计算(只 checkpoint attention) |
| GPU 利用率 60%,profile 全是 gap | 通信串行 | prefetch + 梯度累积 no_sync + 调 bucket |
七、结论
把这条链路串起来看,分布式训练的工程本质是三次正交的降维,别用一个去解决另一个的问题:
- ZeRO/FSDP 解决的是"状态的冗余存储"——它不减少任何计算,只是让 12 字节/参数的优化器历史不再被复制 N 份;
- TP/PP 解决的是"单卡放不下一层/放不下所有层"——这是真正的模型切分,代价是通信;
- 激活重计算 解决的是"batch 开不大"——用可量化的算力(33%,优化后个位数)换显存,是全场唯一可以自由调汇率的交易。
真正负责任的调参顺序是:先把显存账本逐项算出来(参数、梯度、优化器、激活、临时缓冲),找到那一项是瓶颈,再决定动哪个旋钮。绝大多数"分布式训练跑不起来"的问题,追到底都是在这一张账本上算错了某一行——而不是卡不够。

发表评论 取消回复