🏠 总目录📚 本教程 显存优化
📑 本页目录(点开跳转)

09 · 显存优化全家桶

28 分钟 | ⭐ OOM 时按这个顺序试


🎯 一句话

OOM 不是一个问题,是四个问题 —— 参数、梯度、优化器状态、激活, 每一块的解法完全不同。 第 3 章教你算它们各占多少,这一章教你怎么把每一块压下去。

不用检查点全部激活都存着L1L2L3L4L5L6L7L8用检查点只存这几个L1L2L3L4L5L6L7L8重算重算重算重算重算反向传播需要前向的激活值 —— 存不下就重新算一遍显存省 60~70%换来训练慢 20~30%什么时候用显存不够 / 想把 batch 开大时⭐ 这是显存优化里最直白的一个权衡:拿时间买空间,没有免费的部分
只存少数几层的激活,其余的反向时重新算一遍。⭐ 显存省 60~70%,代价是训练慢 20~30% —— 这是显存优化里最直白的一个权衡:拿时间买空间,没有免费的部分。

🩺 一、先诊断:你的 OOM 属于哪一类

import torch
torch.cuda.reset_peak_memory_stats()
train_step()
print(f"峰值 {torch.cuda.max_memory_allocated()/1e9:.1f} GB  "
      f"已保留 {torch.cuda.max_memory_reserved()/1e9:.1f} GB")

# ⭐ 更有用的:完整的显存快照,能看出每一块是谁占的
torch.cuda.memory._record_memory_history()
train_step()
torch.cuda.memory._dump_snapshot("mem.pickle")
# 拖到 https://pytorch.org/memory_viz 里看可视化
   ⭐ 一个关键的判断:把 batch 调到 1,还 OOM 吗?

   还 OOM  → 是【参数/梯度/优化器状态】的问题
             → 调 batch 没用,要用 ZeRO / 量化 / 换优化器
   不 OOM  → 是【激活】的问题
             → 梯度检查点、序列并行都能救

⚠️ allocatedreserved 的差距就是碎片。 差距超过 20% 说明碎片严重,见文末。


🧰 二、七件武器(按性价比排序)

① 梯度检查点 —— 最常用 ⭐

   正常:前向时保存【每一层】的激活,反向时用

   检查点:只保存【少数几个】节点的激活
          反向时从最近的节点【重新前向计算】出中间激活

   ⭐ 显存:O(L) → O(√L)
   ⚠️ 时间:+20~35%(多一次前向)
# 整个模型
from torch.utils.checkpoint import checkpoint_sequential
out = checkpoint_sequential(model.layers, segments=4, input=x)

# ⭐ HuggingFace 一行
model.gradient_checkpointing_enable()

# ⭐ 更精细:选择性检查点(PyTorch 2.x)
#    只重算便宜的(LayerNorm/激活函数),保留贵的(矩阵乘)
from torch.utils.checkpoint import create_selective_checkpoint_contexts

💡 选择性检查点值得一提:全量重算掉 30% 速度, 但只重算那些"算起来便宜、存起来贵"的操作, 常常能用 5–10% 的时间换到大部分的显存收益。

② ZeRO / FSDP —— 切分优化器状态 ⭐⭐

   [第 3 章]算过:优化器状态是最大的一块(12 字节/参数)

   ZeRO 把它切开分给 N 张卡:
   Stage 1:切优化器状态      → 省 4x
   Stage 2:+ 切梯度          → 省 8x
   Stage 3:+ 切参数          → 省 Nx  ⭐

   🔗 第 11 章详讲

③ 激活值量化 / 更低精度

   BF16 激活 → FP8 激活:显存再减半
   ⚠️ 需要 H100 且要处理 scaling(第 6 章)

④ 换优化器

优化器 字节/参数 说明
Adam(FP32 状态) 12 标准
8-bit Adam ~4 bitsandbytes质量损失很小
Adafactor ~1–2 分解二阶矩,大模型常用
SGD + momentum 4 省但大模型效果差
LOMO / 融合更新 ~0 边算梯度边更新,不存梯度
import bitsandbytes as bnb
opt = bnb.optim.AdamW8bit(model.parameters(), lr=3e-4)   # ⭐ 一行换掉

⑤ CPU Offload

   把优化器状态 / 参数放到 CPU 内存,需要时搬过来

   ⚠️ 走 PCIe(第 3 章:比 HBM 慢约 40 倍)
   → 只在【实在装不下】时用,会显著拖慢

   💡 例外:推理时 offload 不常用的专家(MoE)是划算的

⑥ 参数高效微调(PEFT)

   ⭐ 如果你只是【微调】而不是预训练,这是最省的路:

   LoRA:冻结原权重,只训练低秩增量 A·B
   → 可训练参数降到 0.1~1%
   → 梯度和优化器状态跟着降到 1%  ⭐⭐

   QLoRA:base 模型用 4-bit 量化 + LoRA
   → 单张 24GB 卡能微调 33B 模型

🔑 这是显存问题最容易被忽略的解法很多人在硬啃全量微调的显存,而任务本身用 LoRA 就够了。

⑦ 减小 batch + 梯度累积

   micro_batch=1,累积 32 步 → 等效 batch=32

   ⚠️ 但要知道代价:
   · micro batch 太小 → GPU 利用率下降(第 2 章:warp 不够)
   · 通信次数不变但计算变少 → 通信占比上升

   ⭐ 所以这是【最后手段】,不是第一手段
for i, batch in enumerate(loader):
    loss = model(batch) / accum_steps        # ⭐ 别忘了除
    loss.backward()
    if (i + 1) % accum_steps == 0:
        opt.step(); opt.zero_grad(set_to_none=True)

🧭 三、决策树(照着走)

   OOM 了
   │
   ├─ batch=1 还 OOM?
   │   ├─ 是 → 参数/优化器的问题
   │   │      ① 是微调吗?→ 【LoRA / QLoRA】⭐ 最省事
   │   │      ② 换 8-bit Adam / Adafactor
   │   │      ③ 上 ZeRO-2 → ZeRO-3(第 11 章)
   │   │      ④ 还不行 → 张量并行 / 流水线并行(第 12-13 章)
   │   │      ⑤ 最后 → CPU offload(慢)
   │   │
   │   └─ 否 → 激活的问题
   │          ① 【梯度检查点】⭐ 先试这个
   │          ② FlashAttention(第 8 章,把 O(n²) 降到 O(n))
   │          ③ 序列长度能不能减?能不能 packing?
   │          ④ 序列并行(超长序列时)
   │          ⑤ 减小 micro batch + 梯度累积
   │
   └─ 时好时坏 / 显存够但报 OOM? → 【碎片】,见下

🧩 四、显存碎片:那个"明明还有显存却 OOM"

   症状:
   · 报错说"tried to allocate 2GB,free 8GB"但还是失败
   · 训练跑了几小时后突然 OOM
   · reserved 比 allocated 大很多(>20%)

   原因:PyTorch 的缓存分配器把显存切成了很多小块,
        没有【连续】的大块可用
# ⭐ 三个常用手段
import torch
import os
os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True"
#                                        ↑ PyTorch 2.x,对碎片非常有效

torch.cuda.empty_cache()          # 释放缓存(慢,别放在循环里)

# ⭐ 最有效的预防:让形状固定
#    变长序列 → 分桶(bucketing)到固定几个长度

⚠️ empty_cache() 不是解药:它会让分配器重新向驱动申请,反而变慢真正的解法是让分配模式变规律 —— 固定形状、分桶、expandable_segments


📋 五、一张速查表

手段 省哪块 省多少 代价
梯度检查点 激活 O(L)→O(√L) +20~35% 时间
FlashAttention 激活 O(n²)→O(n) 几乎无代价
ZeRO-3 / FSDP 参数+梯度+优化器 ÷N 通信增加
8-bit Adam 优化器 12→4 字节 极小的质量损失
LoRA 梯度+优化器 ÷100 只适用微调
BF16 参数+激活 ÷2 ✅ 基本无代价
CPU offload 任意 ⚠️ 慢很多
减小 batch 激活 线性 MFU 下降

注意 FlashAttention 和 BF16 这两行它们几乎没有代价 —— 应该是默认开启的,不算"优化手段"。 如果你还没开这两个就在纠结 offload,顺序错了


🔗 和站内其他章的关系

相关的地方 这里的位置
第 3 章 显存四大块 本章是每一块的解法
第 3 章 只有激活和 batch 有关 诊断树的第一个分叉
第 6 章 混合精度 省激活的基础手段
第 8 章 FlashAttention 省激活最划算的一个
第 11 章 ZeRO 省优化器状态
ML 基础第 10 章 有意思的对照:那里也是"按性价比排序的武器库"
《Kaggle竞赛方法论》02 训练策略与正则化 显存优化综合策略 同一张武器表的竞赛版(单卡视角);本章补的是"为什么"和多卡那几件 ⭐
《Kaggle竞赛方法论》05 预训练模型详解 LoRA / Adapter / P-Tuning 第 ⑥ 件武器 PEFT 的展开 —— 怎么选、rank 设多少在那边

✅ 检查点

  1. 判断 OOM 属于哪一类的关键实验是什么?
  2. allocatedreserved 的差距说明什么?
  3. 梯度检查点的显存和时间各是什么量级?
  4. 什么是选择性检查点?它好在哪?
  5. 8-bit Adam 能省多少?代价是什么?
  6. 为什么说 LoRA 是"最容易被忽略的解法"?
  7. 为什么"减小 batch"是最后手段而不是第一手段?
  8. 显存碎片的三个症状?最有效的预防是什么?
  9. 速查表里哪两个手段"几乎无代价"?这意味着什么?
👀 答案
  1. 把 batch 调到 1,看还 OOM 吗。还 OOM → 参数/梯度/优化器的问题(调 batch 没用);不 OOM → 激活的问题(梯度检查点能救)。
  2. 差距就是碎片。超过 20% 说明碎片严重。
  3. 显存 O(L) → O(√L),时间 +20~35%(多一次前向)。
  4. 只重算那些"算起来便宜、存起来贵"的操作(LayerNorm、激活函数),保留贵的(矩阵乘)。好在常常用 5-10% 的时间就换到大部分显存收益,而全量重算要掉 30%。
  5. 优化器状态从 12 字节/参数降到约 4。代价是极小的质量损失bitsandbytes 一行换掉。
  6. 因为很多人在硬啃全量微调的显存,而任务本身用 LoRA 就够了。LoRA 把可训练参数降到 0.1-1%,梯度和优化器状态跟着降到 1%。
  7. 因为 micro batch 太小会让 GPU 利用率下降(warp 不够,第 2 章),而且通信次数不变但计算变少,通信占比上升 —— MFU 会明显变差。
  8. ①报错说"tried to allocate 2GB, free 8GB"却失败 ②跑了几小时后突然 OOM ③reserved 比 allocated 大 20% 以上。最有效的预防:让形状固定(变长序列分桶)+ expandable_segments:True。⚠️empty_cache() 反而会变慢。
  9. FlashAttention 和 BF16。意味着它们应该是默认开启的,不算"优化手段"——如果还没开这两个就在纠结 offload,顺序错了。

🛑 可以停在这里

走神救援

OOM 不是一个问题是四个(参数/梯度/优化器/激活),每块解法不同。⭐诊断的关键实验:把 batch 调到 1 还 OOM 吗——还 OOM = 参数/优化器问题(调 batch 没用),不 OOM = 激活问题。allocatedreserved 差 >20% = 碎片。用 torch.cuda.memory._dump_snapshot + pytorch.org/memory_viz 看是谁占的。七件武器:①⭐梯度检查点(显存 O(L)→O(√L),时间 +20~35%;⭐选择性检查点只重算"算起来便宜存起来贵"的,用 5-10% 时间换大部分收益)②⭐⭐ZeRO/FSDP(切优化器状态,Stage1省4x/Stage2省8x/Stage3省Nx)③激活 FP8 ④换优化器8-bit Adam:12→4 字节,质量损失极小bnb.optim.AdamW8bit 一行;Adafactor ~1-2 字节)⑤CPU offload(⚠️走 PCIe 慢 40 倍,实在装不下才用)⑥⭐LoRA/QLoRA最容易被忽略的解法——很多人硬啃全量微调的显存而任务用 LoRA 就够;可训练参数降到 0.1-1%,梯度和优化器跟着降到 1%;QLoRA 让 24GB 卡微调 33B)⑦减小 batch + 梯度累积(⚠️最后手段:micro batch 太小则 warp 不够、通信占比上升,MFU 变差)。决策树:batch=1 还 OOM → LoRA → 8-bit Adam → ZeRO-2/3 → 张量/流水线并行 → offload;不 OOM → 梯度检查点 → FlashAttention → 缩短序列/packing → 序列并行 → 减 micro batch。碎片:症状是"free 8GB 却分配不了 2GB"、跑几小时后突然 OOM;⭐解法是让形状固定(分桶)+ PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,⚠️empty_cache() 不是解药反而更慢。⭐速查表里 FlashAttention 和 BF16 几乎无代价——它们该是默认开的,还没开就纠结 offload 是顺序错了

下一节 👉 10-数据并行与AllReduce.md

打卡记录保存在你的浏览器里,首页能看到总进度