执行摘要:很多人以为分布式训练调不动,是卡不够。真实情况往往是显存账没算清——你以为 70B 模型需要 1.12 TB 是在买显存,其实其中 75% 是三份完全可以切开的冗余状态。本文沿着一条真实的工程决策链展开:先把显存账本逐项拆开,再讲数据/张量/流水三种并行各自解决的是哪个维度的瓶颈,然后深入 ZeRO 三级分片与 PyTorch FSDP 的 FlatParameter 编排,接着算清激活重计算的算力-显存兑换率,最后落到最容易被忽略的一层——你的 GPU 有 30% 时间在等 NCCL。全程给出可直接套用的配置与代码。


一、显存账本:先算清楚再谈并行

混合精度(bf16)+ AdamW 是当前预训练的默认配置。此时每个参数的常驻显存并不是 2 字节,而是:

组件精度每参数字节说明
前向/反向实际使用的参数bf162真正参与 matmul 的副本
梯度bf162反向产出,reduce-scatter 前常驻
主权重 master weightsfp324优化器更新用,bf16 直接累加会掉精度
一阶动量 mfp324Adam
二阶动量 vfp324Adam
合计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无161120 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%,优化后个位数)换显存,是全场唯一可以自由调汇率的交易。

真正负责任的调参顺序是:先把显存账本逐项算出来(参数、梯度、优化器、激活、临时缓冲),找到那一项是瓶颈,再决定动哪个旋钮。绝大多数"分布式训练跑不起来"的问题,追到底都是在这一张账本上算错了某一行——而不是卡不够。

点赞(0) 打赏

评论列表 共有 0 条评论

暂无评论
立即
投稿

微信公众账号

微信扫一扫加关注

发表
评论
返回
顶部