大模型训练优化:FSDP、DeepSpeed ZeRO 与混合精度

26 阅读9分钟

AI 加速器系列 · 第 2 篇


7B 模型用 Adam 训练,光优化器状态就要 112GB 显存。A100 80GB?对不起,连门都进不去。

这还不是最离谱的——你还有模型参数、梯度、前向激活值没算。加起来直奔 130GB+,一张 H100 80GB 也只能干瞪眼。大模型训练的第一个关卡,从来不是速度,是"能不能装得下"。这篇文章拆解显存占用的每一块拼图,然后看 FSDP、DeepSpeed ZeRO 和混合精度三件套怎么把不可能变成可能。


一、显存都去哪了——训练时显存占用拆解

训练时 GPU 显存的占用,可以归纳为四大类。以 Llama-2 7B 为例,逐项算一遍就清楚了。

1.1 模型参数

参数量 7B,用 FP16 存储时为:

7 × 10^9 × 2 bytes = 14 GB

1.2 梯度

每个可训练参数在反向传播时都对应一个梯度,也是 FP16:

7 × 10^9 × 2 bytes = 14 GB

1.3 Adam 优化器状态

这是显存占用的"大户"。Adam 为每个参数维护三个 FP32 缓冲区:

  • Momentum(一阶矩):7B × 4 bytes = 28 GB
  • Variance(二阶矩):7B × 4 bytes = 28 GB
  • Master Parameter Copy(FP32 权重副本):7B × 4 bytes = 28 GB
28 + 28 + 28 = 84 GB

如果你用的是 FP32 存储参数和梯度(不使用混合精度),这个数字会直接翻倍——那就是另一篇哭诉显存不够的博客了。

1.4 前向激活值

前向传播中每一层的中间计算结果被保留下来,反向传播时用于计算梯度。激活值的大小取决于:

  • Batch size
  • 序列长度
  • 隐藏层维度
  • Transformer 层数
  • 是否开启激活检查点(Activation Checkpointing)

不做任何优化时,激活值轻松占到几十 GB。比如 batch size=4、序列长度=4096 的场景下,激活值约占 30-50 GB。

1.5 总账:一张 A100 80GB 够吗?

Model Weights (FP16)  : ████████████ 14 GB
Gradients  (FP16)     : ████████████ 14 GB
Optimizer States(FP32): ████████████████████████████████████████████████████████████████ 84 GB
Activations           : ████████████████████████████ 30-50 GB
─────────────────────────────────────────────────────────
Total                  : ≈ 142-162 GB

单卡 A100 80GB:直接 OOM(Out of Memory)。

这就是为什么大模型训练必须上分布式策略的根本原因——不是为了加速,是为了能用。


二、ZeRO 的三级火箭——把优化器状态"分片"到多张 GPU

DeepSpeed 的 ZeRO(Zero Redundancy Optimizer)不是一种优化,而是三级分片策略的组合:每提升一级,把更多东西切碎、分发出去,换来更大的可用容量。

2.1 ZeRO 三阶段对比

阶段分片内容每卡显存节省通信开销核心操作
ZeRO-1Optimizer States4× 倍速削减低(仅 AllReduce 梯度)各卡算出自己的优化器状态后分发
ZeRO-2+ Gradients8× 倍速削减中(ReduceScatter 替代 AllReduce)梯度算完即切分,不聚合完整副本
ZeRO-3+ Parameters线性缩放(N 卡=N 倍)高(参数需 AllGather)参数按需 AllGather,用完即弃

2.2 内存分布示意图

场景:4 GPUs,每个 GPU 在 ZeRO 各阶段下存储的内容。

        GPU 0           GPU 1           GPU 2           GPU 3

ZeRO-1: [FULL Params]   [FULL Params]   [FULL Params]   [FULL Params]
        [FULL Grads ]   [FULL Grads ]   [FULL Grads ]   [FULL Grads ]
        [1/4 OptSta ]   [1/4 OptSta ]   [1/4 OptSta ]   [1/4 OptSta ]

ZeRO-2: [FULL Params]   [FULL Params]   [FULL Params]   [FULL Params]
        [1/4 Grads  ]   [1/4 Grads  ]   [1/4 Grads  ]   [1/4 Grads  ]
        [1/4 OptSta ]   [1/4 OptSta ]   [1/4 OptSta ]   [1/4 OptSta ]

ZeRO-3: [1/4 Params]   [1/4 Params]    [1/4 Params]    [1/4 Params]
        [1/4 Grads  ]   [1/4 Grads  ]   [1/4 Grads  ]   [1/4 Grads  ]
        [1/4 OptSta ]   [1/4 OptSta ]   [1/4 OptSta ]   [1/4 OptSta ]

关键认知:ZeRO-3 实现了完全的分片,显存占用与 GPU 数量几乎成反比。8 张 GPU 时,每卡参数+梯度+优化器状态的总开销从 112GB 降到约 14GB。

2.3 ZeRO-3 的前向传播过程

ZeRO-3 最精妙的设计在于"按需取用、用完即弃":

Step 1: AllGather —— 从所有 GPU 收集当前层的完整参数
        [GPU0:1/4] + [GPU1:1/4] + [GPU2:1/4] + [GPU3:1/4] → 完整参数

Step 2: Compute   —— 用完整参数执行当前层的前向计算

Step 3: Discard   —— 释放当前层完整参数,只保留本卡负责的分片

Step 4: Repeat    —— 进入下一层,回到 Step 1

反向传播同理,只是方向反过来叠加上梯度计算。整个过程只有正在计算的那一层持有完整参数,其余层的参数始终以分片形态存在。

2.4 DeepSpeed 配置示例

{
  "zero_optimization": {
    "stage": 2,
    "offload_optimizer": {
      "device": "cpu"
    },
    "overlap_comm": true,
    "reduce_bucket_size": 5e8
  }
}

ZeRO-1 适合起步(改动最小),ZeRO-2 是生产环境的常用配置(性价比最高),ZeRO-3 是跑超大规模模型的标准答案(配合 CPU Offload 可以跑 13B+ 模型在单卡上)。


三、FSDP:PyTorch 原生的 ZeRO-3

FSDP(FullyShardedDataParallel)是 PyTorch 官方在 1.11 中引入的分布式训练 API,本质上是 ZeRO-3 的 PyTorch 原生实现。它的设计哲学是:让 ZeRO 的使用体验和 DDP 一样简单。

3.1 DDP vs FSDP:API 对比

# ============ DDP:传统数据并行 ============
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP

model = MyModel()
model = model.to(device)
model = DDP(model, device_ids=[local_rank])
# 每张卡持有完整模型副本 → 显存吃紧

# ============ FSDP:完全分片数据并行 ============
import torch.distributed as dist
from torch.distributed.fsdp import (
    FullyShardedDataParallel as FSDP,
    ShardingStrategy
)

model = MyModel()
model = FSDP(
    model,
    sharding_strategy=ShardingStrategy.FULL_SHARD,  # ZeRO-3
    device_id=torch.cuda.current_device()
)
# 每张卡只持有 1/N 的模型参数 → 显存大幅削减

迁移成本极低:核心区别就是 DDP(model) 换成 FSDP(model),通信后端、进程组初始化、数据加载全部保持不变。

3.2 分片策略选择

FSDP 提供了三种策略,覆盖不同场景:

策略等价关系适用场景
FULL_SHARDZeRO-3模型很大、多卡均摊
HYBRID_SHARD节点内 FULL_SHARD + 节点间复制跨节点集群,减少跨机通信
_HYBRID_SHARD_ZERO2节点内 ZeRO-2 + 节点间复制ZeRO-2 的跨节点版本

3.3 从 DDP 迁移到 FSDP 的完整 diff

# --- DDP 版本 ---
from torch.nn.parallel import DistributedDataParallel as DDP

model = build_model()
model = model.to(device)
model = DDP(model, device_ids=[local_rank])

# --- FSDP 版本 ---
from torch.distributed.fsdp import (
    FullyShardedDataParallel as FSDP,
    MixedPrecision,
    ShardingStrategy,
    CPUOffload,
)

# FSDP 内置了混合精度和 CPU offload 的集成点
bf16_mp_policy = MixedPrecision(
    param_dtype=torch.bfloat16,   # 参数的前向传播类型
    reduce_dtype=torch.bfloat16,  # 梯度的通信类型
    buffer_dtype=torch.bfloat16,  # buffer 的类型
)

model = FSDP(
    model,
    sharding_strategy=ShardingStrategy.FULL_SHARD,
    mixed_precision=bf16_mp_policy,
    cpu_offload=CPUOffload(offload_params=False),  # 可选的 CPU 卸载
    auto_wrap_policy=partial(
        transformer_auto_wrap_policy,
        transformer_layer_cls={TransformerBlock}
    ),
)

auto_wrap_policy 决定了 FSDP 在哪一层"切分"模型——通常以每个 Transformer Block 为粒度。粒度太粗会降低通信效率(每次 AllGather 的数据量太大),粒度太细会增加通信次数(开销大于收益)。


四、混合精度训练(Automatic Mixed Precision)

即使拿 ZeRO-3 把显存问题解决了,还有一个瓶颈:计算速度

4.1 FP16 的理论加速

NVIDIA Tensor Core 的 FP16 吞吐量是 FP32 的 8-16 倍(具体取决于 GPU 架构和矩阵规模):

A100 Tensor Core:
  FP32 throughput : 19.5 TFLOPS
  FP16 throughput : 312  TFLOPS  (16×)
  BF16 throughput : 312  TFLOPS

H100 Tensor Core (with FP8):
  FP32 throughput : 67   TFLOPS
  FP16/FP16 throughput : 990  TFLOPS (约 15×)

把模型参数、前向计算、梯度计算全部换成 FP16,理论上能提速一个数量级。

4.2 两个致命问题

但直接全 FP16 训练,模型会当场发散。原因有二:

问题 1:梯度下溢(Gradient Underflow)

  • FP16 的表示范围是 6e-8 到 65504
  • 梯度中很多值小于 6e-8 时,直接变成 0
  • 零梯度 = 参数不更新 = 训练停滞

问题 2:权重更新的精度丢失

  • 权重更新量(lr × gradient)的量级远小于权重本身
  • FP16 有效精度约 3-4 位十进制 → 加一个极小的 delta 等于没加

4.3 解决方案:FP32 Master Weights + Loss Scaling

业界标准做法是"三步走":

┌──────────────────────────────────────────────────────────┐
│ Forward Pass (FP16)                                      │
│  ┌─────────┐    ┌─────────┐    ┌──────────────────────┐  │
│  │ FP16    │ →  │ FP16    │ →  │ Loss × Scale Factor  │  │
│  │ Weights │    │ Compute │    │ (prevent underflow)  │  │
│  └─────────┘    └─────────┘    └──────────┬───────────┘  │
│                                           │               │
│ Backward Pass (FP16)         ◄────────────┘               │
│  ┌─────────┐    ┌──────────────────────┐                  │
│  │ FP16    │ ←  │ Unscale Gradients    │                  │
│  │ Grads   │    │ (reverse scaling)    │                  │
│  └────┬────┘    └──────────────────────┘                  │
│       │                                                   │
│ Weight Update (FP32)                                      │
│  ┌────▼─────────────────────────────────┐                │
│  │ FP32 Master Weights                  │                │
│  │  += lr × FP32(unscaled_grad)         │                │
│  │  → Copy back to FP16 for next step   │                │
│  └──────────────────────────────────────┘                │
└──────────────────────────────────────────────────────────┘
  • FP16 前向+反向:享受 Tensor Core 的加速
  • FP32 主权重:确保权重更新的精度,攒够足够的信息
  • Loss Scaling:前向时将 Loss 乘以一个大数(如 2^16),反向后再除以同样倍数,把"太小了会下溢"的梯度拉到 FP16 的表示范围内

4.4 PyTorch 实现

from torch.cuda.amp import autocast, GradScaler

model = build_model()
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)
scaler = GradScaler()  # 默认 init_scale=2^16

for batch in dataloader:
    optimizer.zero_grad()

    # 前向传播:自动选择 FP16 算子
    with autocast():
        loss = model(batch)

    # 反向传播:scaler 管理梯度缩放
    scaler.scale(loss).backward()
    scaler.step(optimizer)
    scaler.update()

scaler.update() 会根据梯度是否出现 inf/nan 动态调整缩放因子:

  • 本轮无溢出 → 增大 scale(追求更大动态范围)
  • 本轮有溢出 → 跳过本次更新,缩小 scale

4.5 BF16:更优雅的方案

BF16(Brain Floating Point 16)用更少的小数位换来了与 FP32 相同的指数位:

FP32: [1 sign][8 exponent][23 mantissa]    范围: 1e-38
FP16: [1 sign][5 exponent][10 mantissa]    范围: 6e-8   ← 容易下溢
BF16: [1 sign][8 exponent][7  mantissa]    范围: 1e-38  ← 与 FP32 相同

BF16 的优势:

  • 不需要 Loss Scaling:指数位与 FP32 一致,不存在下溢问题
  • 代码更简洁,少了一个需要调参的 GradScaler
  • 精度更低(7 位尾数 vs 10 位尾数),但大模型训练基本不受影响

BF16 的代价:

  • 需要 A100 或更新的 GPU(V100 不支持)
  • 某些对精度敏感的操作(如 Embedding、LayerNorm)仍需 FP32
# BF16 混合精度:不需要 GradScaler
with autocast(dtype=torch.bfloat16):
    loss = model(batch)
loss.backward()  # 直接 backward,无需 scaler
optimizer.step()

一句话总结

ZeRO 把显存分成 N 份存 N 张卡上(空间换空间),混合精度把计算从 FP32 压成 FP16(精度换速度),两者组合是大模型训练的标配。FSDP 是 PyTorch 官方对这套思想的一键封装。

下一篇:GPU 集群通信解剖——NCCL 的 Ring、Tree 与 CollNet 拓扑,以及 AllReduce 为什么不是你想象中那样工作的。