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