📑 本页目录(点开跳转)
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 里看可视化
信息关系
⚠️
allocated和reserved的差距就是碎片。 差距超过 20% 说明碎片严重,见文末。
🧰 二、七件武器(按性价比排序)
① 梯度检查点 —— 最常用
信息关系
# 🧩 骨架:`model` 来自你自己的代码,这一段只看写法
# 整个模型
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 —— 切分优化器状态
信息关系
③ 激活值量化 / 更低精度
信息关系
④ 换优化器
| 优化器 | 字节/参数 | 说明 |
|---|---|---|
| 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
关键信息
⑥ 参数高效微调(PEFT)
结果对照
🔑 这是显存问题最容易被忽略的解法: 很多人在硬啃全量微调的显存,而任务本身用 LoRA 就够了。
⑦ 减小 batch + 梯度累积
流程图
# 🧩 骨架:`loader` 来自你自己的代码,这一段只看写法
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差得多就是碎片。七件武器里最该记的四条:⭐ 梯度检查点把激活显存降一个量级、代价是多花两三成时间(选择性检查点只重算「算起来便宜、存起来贵」的那部分,性价比更高);⭐ 切优化器状态(ZeRO/FSDP)按 stage 递进;⭐ 换 8-bit 优化器几乎是白拿——一行的事、质量损失极小;⭐⭐ LoRA 是最容易被忽略的解法——很多人在硬啃全量微调的显存,而任务用 LoRA 就够,可训练参数掉到百分之一量级,梯度和优化器状态跟着掉。
⚠️ 两个「最后手段」:CPU offload 要走 PCIe,慢一个量级,实在装不下才用;减小 micro batch 会让 warp 不够、通信占比上升,MFU 反而变差。
⭐ 决策树按那个诊断实验分叉:batch=1 还 OOM 走 LoRA → 8-bit 优化器 → 切分 → 并行 → offload;不 OOM 走梯度检查点 → FlashAttention → 缩短序列 → 序列并行。
碎片的症状很有辨识度:「空着好几 GB 却分配不了小一半的显存」、跑几小时后突然 OOM。⭐ 解法是让形状固定(分桶)加上换分配器策略;⚠️
empty_cache()不是解药,反而更慢。⭐ 最后一句是顺序问题:FlashAttention 和 BF16 几乎无代价,它们本来就该是默认开的——还没开就先纠结 offload,是把顺序搞反了。
下一节 👉 10-数据并行与AllReduce.md