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

AI 加速器系列 · 第 2 篇


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

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


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

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

1.1 模型参数

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

ini 复制代码
7 × 10^9 × 2 bytes = 14 GB

1.2 梯度

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

ini 复制代码
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
ini 复制代码
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 够吗?

java 复制代码
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-1 Optimizer States 4× 倍速削减 低(仅 AllReduce 梯度) 各卡算出自己的优化器状态后分发
ZeRO-2 + Gradients 8× 倍速削减 中(ReduceScatter 替代 AllReduce) 梯度算完即切分,不聚合完整副本
ZeRO-3 + Parameters 线性缩放(N 卡=N 倍) 高(参数需 AllGather) 参数按需 AllGather,用完即弃

2.2 内存分布示意图

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

ini 复制代码
        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 最精妙的设计在于"按需取用、用完即弃":

vbnet 复制代码
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 配置示例

json 复制代码
{
  "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 对比

python 复制代码
# ============ 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_SHARD ZeRO-3 模型很大、多卡均摊
HYBRID_SHARD 节点内 FULL_SHARD + 节点间复制 跨节点集群,减少跨机通信
_HYBRID_SHARD_ZERO2 节点内 ZeRO-2 + 节点间复制 ZeRO-2 的跨节点版本

3.3 从 DDP 迁移到 FSDP 的完整 diff

python 复制代码
# --- 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 架构和矩阵规模):

yaml 复制代码
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

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

scss 复制代码
┌──────────────────────────────────────────────────────────┐
│ 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 实现

python 复制代码
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 相同的指数位:

ini 复制代码
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
python 复制代码
# 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 为什么不是你想象中那样工作的。

相关推荐
樊小肆1 小时前
2568 万 token 才花 2 块 2:聊聊 DeepSeeker-Code 怎么吃满上下文缓存
前端·人工智能·后端
Zane19941 小时前
ClassName() 只是一步?拆开看 __new__ 和 __init__ 各自在干什么
后端·python
geovindu1 小时前
java: Memento Pattern
java·开发语言·后端·备忘录模式·行为模式
星哥的编程之路1 小时前
万字深度解析 Agent 学习路线
后端
程序边界1 小时前
SQL Server数据迁移这件事,远比你想的复杂——但也远比你想的简单(下)
后端
Zane19942 小时前
JRE 去哪了?一次讲清楚 JDK、JRE、JVM 到底是什么关系?
java·后端
JavaPub2 小时前
Ontology 本体论是什么?从哲学概念到 AI、知识图谱与软件工程
后端
liuxiaocheng2 小时前
聊聊 Vercel AI SDK 的流式协议:前后端到底是怎么"边想边说"的
前端·后端
胡萝卜术2 小时前
在浏览器中跑 DeepSeek-R1:WebGPU 推理全流程深度解析
前端·javascript·面试