显存优化——让 70B 模型在有限显存中跑起来 / Memory Optimization for Running 70B Models
📅 创建时间:2026-06-02 🏷️ 标签:#显存 #ZeRO #Gradient-Checkpointing #CPU-Offload #Activation #Activation-Recomputation 📚 前置知识:[[01-gpu-hardware]](GPU 硬件基础) [[02-distributed-training]](分布式训练) 📚 相关知识:[[04-mixed-precision]](混合精度) [[08-efficient-finetuning]](高效微调)
场景:DeepSpeed ZeRO 为什么能让 70B 模型在 8 张 80GB 卡上跑起来
┌─────────────────────────────────────────────────────────────┐
│ │
│ 你的 8 张 H100 80GB 集群。 │
│ 想跑 Llama-3-70B。 │
│ │
│ 显存需求估算(BF16,seq_len=4096, batch=1): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 模型参数:70B × 2 bytes = 140 GB │ │
│ │ 梯度:70B × 2 bytes = 140 GB │ │
│ │ 优化器状态:70B × 4 bytes = 280 GB │ │
│ │ Activation:~100 GB(32 层 × seq² × batch) │ │
│ │ │ │
│ │ 总计:660 GB │ │
│ │ 可用:8 × 80GB = 640 GB → 仍然不够! │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 但 DeepSpeed ZeRO-3 + Activation Checkpointing 后: │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 参数分片:每卡 140GB / 8 = 17.5 GB │ │
│ │ 梯度分片:每卡 140GB / 8 = 17.5 GB │ │
│ │ 优化器分片:每卡 280GB / 8 = 35 GB │ │
│ │ Activation Recomputation:~20 GB │ │
│ │ 模型层分片(PP):每卡 17.5 GB │ │
│ │ │ │
│ │ 总计:约 50 GB/卡 → 8 × 80GB 够用了! │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 这些技术是怎么工作的?这就是本章要回答的问题。 │
│ │
└─────────────────────────────────────────────────────────────┘第1节:ZeRO——用通信换显存
ZeRO 的三阶段演进
┌─────────────────────────────────────────────────────────────┐
│ ZeRO 三个 Stage 的显存节省效果 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 基线(DDP): │
│ 每张卡存:参数 + 梯度 + 优化器状态(全量副本) │
│ → 每卡显存 = P + G + O = 3×模型参数 │
│ → 70B 模型:每卡存 140 + 140 + 280 = 560 GB │
│ → 不可能 │
│ │
│ ZeRO-1(优化器状态分片): │
│ 每张卡只存 1/N 的优化器状态 │
│ → 每卡显存 = P + G + O/N = 140 + 140 + 280/8 = 215 GB │
│ → 节省约 60% │
│ │
│ ZeRO-2(梯度分片): │
│ 梯度也分片,加上优化器分片 │
│ → 每卡显存 = P + G/N + O/N = 140 + 17.5 + 35 = 192.5 GB │
│ → 节省约 66% │
│ │
│ ZeRO-3(参数也分片): │
│ 参数 + 梯度 + 优化器状态全部按 rank 分片 │
│ → 每卡显存 = P/N + G/N + O/N = 140/8 + 17.5 + 35 ≈ 70 GB│
│ → 节省约 87.5%,8 卡可以跑 70B 模型 │
│ │
│ 关键:ZeRO-3 的代价是每次 Forward/Backward 前需要 │
│ AllGather 参数,用通信量换显存 │
│ │
└─────────────────────────────────────────────────────────────┘ZeRO-3 的 Forward 和 Backward 过程
┌─────────────────────────────────────────────────────────────┐
│ ZeRO-3 Forward 过程详解 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 问题:参数被分片到 8 张卡,每张卡只有 1/8 的参数 │
│ 前向计算时,每层需要完整的权重 │
│ │
│ 解决:AllGather 获取需要的参数部分,计算后释放 │
│ │
│ 以 Transformer Layer 为例(8 卡分片): │
│ │
│ Step 1:Attention Layer │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ Forward Pass │ │
│ │ │ │
│ │ Layer 0: │ │
│ │ GPU 0~7: AllGather(W_qkv_0) → 各自获得完整 W_qkv_0│ │
│ │ 计算 Q/K/V │ │
│ │ AllGather(W_o_0) → 各自获得完整 W_o_0 │ │
│ │ 计算 Attention Output │ │
│ │ 释放 W_qkv_0, W_o_0(不再占用显存) │ │
│ │ │ │
│ │ Layer 1~31: 同上 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ Step 2:MLP Layer │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ GPU 0~7: AllGather(W_up_0) → 各自获得完整 W_up_0 │ │
│ │ 计算 Up Projection │ │
│ │ AllGather(W_down_0) → 各自获得完整 W_down_0│ │
│ │ 计算 Down Projection │ │
│ │ 释放 W_up_0, W_down_0 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 通信开销: │
│ → 每层 2 次 AllGather(QKV + O),2 次(Up + Down) │
│ → 32 层 × 4 次 AllGather = 128 次通信/步 │
│ → 通信量 = 参数大小(140GB),但被分摊到 128 次 │
│ │
│ 优化:通信和计算 overlap(下一个 Layer 的通信 + 当前计算) │
│ │
└─────────────────────────────────────────────────────────────┘ZeRO 在 DeepSpeed 中的配置
python
# DeepSpeed ZeRO-3 配置示例
ds_config = {
"zero_optimization": {
"stage": 3, # ZeRO-3
"offload_optimizer": {
"device": "cpu", # 优化器状态 offload 到 CPU
"pin_memory": True,
},
"offload_param": {
"device": "nvme", # 参数 offload 到 NVMe(更慢但更便宜)
},
"allgather_partitions": True, # AllGather 获取参数分区
"allgather_bucket_size": 5e8, # bucket 大小,控制通信粒度
"overlap_comm": True, # 通信计算 overlap
"contiguous_gradients": True, # 梯度连续存储,减少碎片
},
"fp16": {
"enabled": True,
"loss_scale": 0, # 动态 loss scale
"loss_scale_window": 1000,
},
"gradient_clipping": 1.0,
}第2节:Gradient Checkpointing——用计算换显存
问题:Activation 显存随 Batch 和 Sequence 线性增长
┌─────────────────────────────────────────────────────────────┐
│ Activation 显存爆炸的原因 │
├─────────────────────────────────────────────────────────────┤
│ │
│ Forward 时,每一层的每个算子都要保存输入(用于 Backward) │
│ │
│ Transformer Layer 的 Activation 显存: │
│ │
│ Layer 0: │
│ Q = X @ W_q → 保存 X, W_q, Q(用于 dW 计算) │
│ K = X @ W_k → 保存 X, W_k, K │
│ V = X @ W_v → 保存 X, W_v, V │
│ S = Q @ K^T → 保存 Q, K, S(Attention Score) │
│ A = softmax(S/d) @ V → 保存 S, V, A(用于 dV 计算) │
│ O = A @ W_o → 保存 A, W_o, O │
│ │
│ 关键问题:Attention Score 矩阵 S = (seq_len, seq_len) │
│ seq_len=4096 → S = 4096 × 4096 × 4 bytes = 64 MB │
│ 32 层 → 64 MB × 32 = 2 GB │
│ batch_size=16 → 2 GB × 16 = 32 GB │
│ │
│ 更糟的是:Backward 时需要 S 的每一行来计算 dV, │
│ 需要 A 的每一列来计算 dQ/K │
│ → 不能简单删除,必须全部保存 │
│ │
└─────────────────────────────────────────────────────────────┘解法:Forward 时不存,Backward 时重算
┌─────────────────────────────────────────────────────────────┐
│ Gradient Checkpointing(也叫 Activation Recomputation)│
├─────────────────────────────────────────────────────────────┤
│ │
│ 核心思想: │
│ → 不存储中间 activation,用时重新算 │
│ → 牺牲 20-30% 的计算时间,换取 50%+ 的显存节省 │
│ │
│ 朴素做法(存所有 activation): │
│ Forward: X → A → B → C → D → loss │
│ Backward: loss → dD → dC → dB → dA → dX │
│ → 存 A, B, C, D,Backward 直接用 │
│ │
│ Checkpointing 做法(只存部分): │
│ Forward: X → A → B → C → D → loss │
│ Backward: │
│ 从 D 重算 C → dD → dC │
│ 从 C 重算 B → dC → dB(用保存的 C) │
│ 从 B 重算 A → dB → dA │
│ 从 A 重算 X → dA → dX │
│ │
│ 在 Transformer 中的实现: │
│ → 每隔 1 层(或若干层)保存 checkpoint │
│ → 32 层:保存 16 个 checkpoint(每 2 层一个) │
│ → Forward 时:每个 checkpoint 后的层重新计算 │
│ │
│ PyTorch 实现: │
│ torch.utils.checkpoint.checkpoint(model, input) │
│ # 手动指定哪些层用 checkpoint,哪些层正常保存 │
│ │
└─────────────────────────────────────────────────────────────┘Checkpointing 的 trade-off 分析
┌─────────────────────────────────────────────────────────────┐
│ Checkpointing 的计算 vs 显存 trade-off │
├─────────────────────────────────────────────────────────────┤
│ │
│ 以 Llama-3-8B 为例,seq_len=4096, batch=1: │
│ │
│ 不做 Checkpointing: │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ Activation 显存 ≈ 64 GB │ │
│ │ Forward 计算:1 次 │ │
│ │ Backward 计算:1 次 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 每层 Checkpoint(每层一个 checkpoint): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ Activation 显存 ≈ 2 GB(每层只存输入) │ │
│ │ Forward 计算:1 次 │ │
│ │ Backward 计算:1 次(每层重算一遍 forward) │ │
│ │ → 计算量增加 30%(2 次完整 forward/层) │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 每 2 层 Checkpoint: │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ Activation 显存 ≈ 4 GB │ │
│ │ Forward 计算:1 次 │ │
│ │ Backward 计算:1.15 次(更少的重算) │ │
│ │ → 计算量增加 15% │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 实际推荐: │
│ → 显存紧张时:每层 checkpoint(省显存,多 30% 计算) │
│ → 显存充裕时:每 2 层 checkpoint(省 15% 计算) │
│ → 显存极充裕:不做 checkpoint(最快) │
│ │
│ ⚠️ FlashAttention 天然支持选择性 Checkpointing: │
│ → FlashAttention 只存储 Attention 的压缩状态(MHA/GA) │
│ → 比 naive attention 节省大量显存 │
│ │
└─────────────────────────────────────────────────────────────┘第3节:CPU 和 NVMe Offload——把显存卸载到内存和磁盘
为什么需要 Offload
┌─────────────────────────────────────────────────────────────┐
│ CPU / NVMe Offload 的场景 │
├─────────────────────────────────────────────────────────────┤
│ │
│ ZeRO-3 已经把显存省到了极致,但仍有场景不够: │
│ │
│ 场景 1:单卡跑 7B 模型,显存不够 │
│ → ZeRO-3 可以在 1 张 80GB 卡上跑 7B(需要 CPU Offload) │
│ │
│ 场景 2:消费级显卡(RTX 3090 24GB)跑 LoRA 微调 │
│ → 量化(INT4)+ CPU Offload → 勉强能跑 │
│ │
│ 场景 3:极致显存优化,把显存全给 Activation │
│ → ZeRO-3 + 优化器 CPU Offload → 参数全在显存 │
│ │
│ Offload 的速度关系: │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ GPU HBM: ~3.5 TB/s │ │
│ │ CPU RAM: ~100 GB/s(PCIe) │ │
│ │ NVMe SSD: ~5 GB/s(PCIe 4.0) │ │
│ │ 差距: 700x(HBM vs NVMe) │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 原则: │
│ → CPU RAM Offload:可以接受(有 overlap) │
│ → NVMe Offload:只有优化器状态可以(太慢) │
│ │
└─────────────────────────────────────────────────────────────┘DeepSpeed CPU Offload 配置
python
# ZeRO-3 + 优化器 CPU Offload
ds_config = {
"zero_optimization": {
"stage": 3,
"offload_optimizer": {
"device": "cpu", # 优化器状态 offload 到 CPU
"pin_memory": True, # 固定内存,传输更快
"buffer_count": 5, # CPU 缓冲数量
},
"offload_param": {
"device": "nvme", # 参数 offload 到 NVMe
"nvme_path": "/local/nvme", # NVMe 路径
"buffer_size": 1e8, # buffer 大小
},
},
}Offload 对训练效率的影响
┌─────────────────────────────────────────────────────────────┐
│ CPU Offload 的性能代价分析 │
├─────────────────────────────────────────────────────────────┤
│ │
│ ZeRO-3 无 Offload: │
│ → 所有参数在 GPU,Forward/Backward 全部 GPU 执行 │
│ → 效率高,但需要大量显存 │
│ │
│ ZeRO-3 + 优化器 CPU Offload: │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ Forward/Backward: GPU 执行(参数从 CPU 拉到 GPU) │ │
│ │ Optimizer Step: CPU 执行(Adam 状态在 CPU) │ │
│ │ │ │
│ │ 时间分解: │ │
│ │ Forward + Backward: 100ms(GPU 计算) │ │
│ │ CPU 参数拉回 GPU: 20ms(PCIe 传输) │ │
│ │ Optimizer Step: 30ms(CPU 计算) │ │
│ │ GPU 参数拉回 CPU: 5ms │ │
│ │ ───────────────────────────── │ │
│ │ 总计: 155ms(vs 100ms) │ │
│ │ 效率损失: ~35% │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 优化策略: │
│ 1. Overlap:下一个 micro-batch Forward 时,执行上一个的 │
│ Optimizer Step(GPU 计算 + CPU 计算并行) │
│ 2. Pin Memory:固定 CPU 内存,减少传输开销 │
│ 3. NVMe 异步 IO:参数卸载到 NVMe 时用异步 IO │
│ │
│ 结论:CPU Offload 适合显存严重不足的场景, │
│ 但会牺牲 20-50% 的训练效率 │
│ │
└─────────────────────────────────────────────────────────────┘第4节:显存优化技术的组合策略
不同模型规模的最优配置
┌─────────────────────────────────────────────────────────────┐
│ 不同场景下的显存优化策略 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 7B 模型(单卡 A100 80GB): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 方案 A:纯混合精度(FP16/BF16) │ │
│ │ → 模型 14GB + 梯度 14GB + 优化器 28GB = 56GB │ │
│ │ → Activation(batch=4, seq=2048)≈ 8GB │ │
│ │ → 总计 64GB,勉强可跑,但 batch size 很小 │ │
│ │ │ │
│ │ 方案 B:ZeRO-2 + Activation Checkpointing │ │
│ │ → 梯度+优化器分片:42GB/8 = 5.25GB/卡 │ │
│ │ → Activation Checkpointing:省 50% │ │
│ │ → 可以增大 batch size 到 16 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 70B 模型(8 卡 A100 80GB): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 方案:ZeRO-3 + Activation Checkpointing + PP │ │
│ │ → 参数分片:140GB/8 = 17.5GB/卡 │ │
│ │ → 梯度分片:140GB/8 = 17.5GB/卡 │ │
│ │ → 优化器分片:280GB/8 = 35GB/卡 │ │
│ │ → Activation Checkpointing(每 2 层):~15GB │ │
│ │ → 总计:~85GB/卡 → 8 卡 640GB 可跑 │ │
│ │ → 但效率较低(ZeRO-3 通信量大) │ │
│ │ │ │
│ │ 更优方案:ZeRO-3 + PP + Activation Checkpointing │ │
│ │ → PP 分担每卡层数,减少 ZeRO-3 通信 │ │
│ │ → PP=4, ZeRO-3 + Checkpointing → ~55GB/卡 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 400B+ 模型(千卡集群): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 方案:TP + PP + DP + ZeRO-3 + Activation CP │ │
│ │ → 3D 并行是必须的 │ │
│ │ → 每个维度都优化到极致 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
└─────────────────────────────────────────────────────────────┘显存优化技术总结
┌─────────────────────────────────────────────────────────────┐
│ 显存优化技术全景图 │
├─────────────────────────────────────────────────────────────┤
│ │
│ │ 技术 │ 显存节省 │ 计算代价 │ 通信代价 │
│ ├───────────────────────┼──────────┼──────────┼──────────┤
│ │ 混合精度 BF16 │ 50% │ 无 │ 无 │
│ ├───────────────────────┼──────────┼──────────┼──────────┤
│ │ ZeRO-1(优化器分片) │ ~60% │ 无 │ 中 │
│ ├───────────────────────┼──────────┼──────────┼──────────┤
│ │ ZeRO-2(梯度分片) │ ~66% │ 无 │ 中 │
│ ├───────────────────────┼──────────┼──────────┼──────────┤
│ │ ZeRO-3(参数分片) │ ~87% │ 无 │ 高 │
│ ├───────────────────────┼──────────┼──────────┼──────────┤
│ │ Activation CP(全层) │ ~70% │ +30% │ 无 │
│ ├───────────────────────┼──────────┼──────────┼──────────┤
│ │ Activation CP(间隔) │ ~50% │ +15% │ 无 │
│ ├───────────────────────┼──────────┼──────────┼──────────┤
│ │ CPU 优化器 Offload │ ~50% │ +20% │ 高 │
│ ├───────────────────────┼──────────┼──────────┼──────────┤
│ │ NVMe 参数 Offload │ ~20% │ +40% │ 极高 │
│ ├───────────────────────┼──────────┼──────────┼──────────┤
│ │ FlashAttention │ ~40% │ ~0 │ 无 │
│ │
│ 组合策略(70B/8卡): │
│ BF16 + ZeRO-3 + Activation CP + PP → ~55GB/卡 │
│ │
└─────────────────────────────────────────────────────────────┘升华:显存优化的工程哲学
┌─────────────────────────────────────────────────────────────┐
│ 显存优化的核心 trade-off │
├─────────────────────────────────────────────────────────────┤
│ │
│ 显存优化的本质:用其他资源换显存 │
│ │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 换法 │ 代价 │ │
│ ├─────────────────────────────────────────────────────┤ │
│ │ 通信换显存(ZeRO) │ NCCL 带宽 │ │
│ │ 计算换显存(CP) │ Forward 额外计算 │ │
│ │ 时间换显存(Offload) │ CPU/NVMe 传输延迟 │ │
│ │ 精度换显存(量化) │ 模型精度可能下降 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 选择原则: │
│ → 通信带宽充裕(IB HDR)→ 优先 ZeRO-3 │
│ → 通信带宽受限(以太网)→ 优先 Activation CP │
│ → 单卡场景 → 量化 + CPU Offload │
│ │
│ 一句总结: │
│ 显存永远不够,但永远有办法。 │
│ 代价是什么,才是需要权衡的。 │
│ │
└─────────────────────────────────────────────────────────────┘"AI 可查 vs 必须理解"清单
AI 可查:
✅ DeepSpeed ZeRO 的具体配置参数
✅ torch.utils.checkpoint 的详细 API
✅ 不同模型的 Activation 显存计算公式
必须理解:
🔴 ZeRO Stage 1/2/3 各自节省多少显存,以及通信代价如何
🔴 ZeRO-3 Forward 时为什么要 AllGather,以及什么时候可以 overlap
🔴 Gradient Checkpointing 的核心思想(不存重算),以及计算代价是多少
🔴 FlashAttention 为什么能省显存(压缩状态 vs 完整 S 矩阵)
🔴 CPU Offload vs NVMe Offload 的速度差异,以及什么可以 offload
🔴 不同模型规模(7B/70B/400B)应该用什么显存优化组合学习状态:🟡 开始学习