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 为什么不是你想象中那样工作的。