报错现象
报错原文(PyTorch)
```
torch.cuda.OutOfMemoryError: CUDA out of memory. Tried to allocate 2.00 GiB
(GPU 0; 23.69 GiB total capacity; 21.43 GiB already allocated; 1.02 GiB free;
22.10 GiB reserved in total by PyTorch)
If reserved but unallocated memory is large try setting
PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True to avoid fragmentation.
```
报错原文(TensorFlow / JAX)
```
tensorflow.python.framework.errors_impl.ResourceExhaustedError:
OOM when allocating tensor with shape[32,12,512,512] and type float
```
在什么操作下出现
- 训练刚起步,第一个 iteration 就炸;
- 训练跑了几百上千步之后突然炸;
- 从训练切到验证/评估阶段时炸;
- 推理时长文本、大图、长视频一次性喂进去炸;
- 保存 checkpoint、写日志、做指标计算时炸;
- 多卡启动后,只有某一张卡炸。
影响范围
单个进程直接退出,训练中断。在 Jupyter 里会让 kernel 挂掉,但显存未必马上释放,重启后再跑还是 OOM。多卡训练里只要一张卡 OOM,整个 DDP 作业通常一起挂掉,前面的训练进度全丢。这类问题不会自己变好,只会随着序列变长、并发变高更容易复现。
可能原因
按出现概率从高到低排:
1. 单次迭代的峰值显存超了容量。batch size 太大、序列太长、图像分辨率太高,或者模型本身太大。
2. 显存被别的进程占着,或者上一次的训练进程没退干净,残留僵尸进程一直挂着显存。
3. 缓存池碎片化。报错里 reserved in total by PyTorch 很大,free 很小 —— 显存其实被 PyTorch 缓存着,但形状对不上,分配不出来。
4. 代码里的隐式显存泄漏。loss 累加时写 total_loss += loss 而不是 loss.item(),验证阶段没加 torch.no_grad(),或者变量一直持有计算图。
5. 精度问题。全程 FP32 训练,没有用混合精度,显存占用比 BF16/FP16 高出一截。
6. 优化器状态和梯度占大头。全量微调大模型时,Adam 的状态量往往比模型权重还大。
7. 多卡配置不当。DDP 每张卡都存一份完整模型副本,显存不会因为卡多而下降;device_map="auto" 分配不均时,某张卡会先炸。
8. 数据与环境相关。pin_memory=True、num_workers 过多、动态 shape 导致缓存多份 kernel 与显存。
逐条排查与解决
1. 先看卡上到底谁在占显存
```bash
nvidia-smi
nvidia-smi --query-compute-apps=pid,process_name,used_memory --format=csv
```
判断方法:如果 nvidia-smi 里显示的显存占用远大于模型理论大小,而且进程 PID 不是你当前这个训练任务,基本就是残留进程。
清理残留进程(确认无用再执行):
```bash
kill -9 <PID>
或者按卡查占用进程
fuser -v /dev/nvidia*
```
清完之后再跑一次,如果显存立刻恢复正常,说明是原因 2,属于最好解决的一类。
2. 看 PyTorch 自己的显存账本
在训练脚本里插几行:
```python
import torch
print("allocated:", torch.cuda.memory_allocated() / 1024**3, "GiB")
print("reserved :", torch.cuda.memory_reserved() / 1024**3, "GiB")
print("max alloc:", torch.cuda.max_memory_allocated() / 1024**3, "GiB")
print(torch.cuda.memory_summary())
```
判断方法:
allocated接近卡容量 → 真的是模型/激活太大,往下面的显存优化方向走;allocated不大但reserved接近卡容量 → 碎片化或缓存未释放;max_memory_allocated比当前allocated大很多 → 说明某个瞬间有峰值,通常是验证阶段或某次异常长的输入。
针对碎片化,设置分配器参数:
```bash
export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True
```
这条环境变量要在启动 Python 之前设置,进程内再改无效。不同 PyTorch 版本支持的参数项略有差异,以官方文档当前版本为准。
在代码里手动释放缓存(排查阶段用,别放进热循环):
```python
torch.cuda.empty_cache()
```
empty_cache() 只是把 PyTorch 缓存池还给驱动,不会释放还在被引用的张量,所以它救不了真正的显存泄漏。
3. 定位到底是哪一行炸的
PyTorch 的报错栈经常指向不到真正的 kernel。加一个环境变量让执行同步化:
```bash
CUDA_LAUNCH_BLOCKING=1 python train.py
```
跑得会慢很多,但 traceback 会准确指到出错的那次前向/反向。定位完记得去掉。这个方法只用于定位,别带进正式训练。
4. 降 batch + 梯度累积
先把 batch size 减半试试,如果显存立刻够用,就确认是原因 1。
```python
accum_steps = 8
for i, batch in enumerate(loader):
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
out = model(**batch)
loss = out.loss / accum_steps
loss.backward()
if (i + 1) % accum_steps == 0:
optimizer.step()
optimizer.zero_grad(set_to_none=True)
```
注意几点:
optimizer.zero_grad(set_to_none=True)比默认写法更省,它把梯度设成None而不是填零张量;- 梯度累积只省显存,不改变有效 batch 的数学结果,但 batch norm 的统计量会受影响,视觉任务要注意;
- 累加 loss 打日志时一定写
loss.item(),写total_loss += loss会把整个计算图钉在显存里,这是最常见的隐性泄漏。
5. 混合精度
```python
scaler = torch.cuda.amp.GradScaler()
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
out = model(**batch)
loss = out.loss
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
```
判断方法:把精度从 FP32 换成 BF16 后,显存和吞吐一般都会明显改善。BF16 数值范围大,通常不需要 GradScaler;FP16 需要,否则容易梯度下溢。
注意:有些算子(比如部分归一化、loss 计算)对精度敏感,需要留在 FP32,autocast 会自动处理大部分情况,遇到 NaN 再单独排查。
6. 梯度检查点(用时间换显存)
```python
Hugging Face 模型
model.gradient_checkpointing_enable()
或者手动包一层
from torch.utils.checkpoint import checkpoint
out = checkpoint(block, x)
```
判断方法:开启后激活值显存通常能降一半以上,代价是反向要重算前向,训练变慢。如果显存降了但速度降得离谱,说明检查点颗粒度太细,可以改到整层级别。
注意:AI 词典:梯度检查点">梯度检查点对 batch size 帮助明显,但如果单层参数本身就放不下(比如超大 embedding),它救不了。
7. 优化器状态与量化
全量微调时,混合精度加 Adam,经验上每参数要按十几字节估算(参数、梯度、优化器状态加起来),7B 量级模型的优化器状态就非常可观,具体数字以实测为准。
先换 8-bit 优化器:
```python
import bitsandbytes as bnb
optimizer = bnb.optim.Adam8bit(model.parameters(), lr=1e-4)
```
再考虑量化加载:
```python
from transformers import AutoModelForCausalLM, BitsAndBytesConfig
bnb_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.bfloat16,
)
model = AutoModelForCausalLM.from_pretrained(
model_name,
quantization_config=bnb_config,
device_map="auto",
)
```
判断方法:如果模型权重本身就占了大半张卡,量化是收益最直接的一步。4-bit 量化大约把权重显存降到原来的四分之一量级。
如果只是想让模型适配某个下游任务,优先考虑 LoRA 之类的参数高效微调,只训适配器,优化器状态随之大幅缩小。量化加 LoRA(常说的 QLoRA)是显存紧张时的常见组合。
注意:量化会带来精度损失,推理任务影响通常可接受,训练任务建议只在显存实在不够时用。
8. 多卡相关
先明确一件事:DDP 不会降低单卡显存。每张卡都有完整的模型、梯度和优化器状态,卡多了只提升吞吐。
```bash
torchrun --nproc_per_node=4 train.py
```
如果每张卡都 OOM,说明单卡就放不下,要先用前面几条优化。如果只有某张卡 OOM,检查:
- 数据是否按 rank 正确切分,别让 0 号卡拿到超长样本;
device_map="auto"的切分是否均衡,可以用hf_device_map打印出来看;- 是否有的卡同时还在跑评估或日志。
想真正降低单卡占用,要用分片类方案:
- ZeRO-2:分片优化器状态和梯度;
- ZeRO-3:连参数一起分片,单卡占用显著下降,通信开销上升;
- 张量并行 / 流水线并行:适合单层都放不下的超大模型。
DeepSpeed、FSDP、Accelerate 都能配这些,配置文件按官方文档当前版本写。
都不管用时的兜底方案
1. 把模型换小。同系列里选更小的尺寸,先跑通再谈效果。很多"必须用大模型"的需求,其实小模型加更好的数据就能满足。
2. 减少输入规模。序列长度从 4096 截到 1024,图片分辨率从 1024 降到 512,视频抽帧减半 —— 这些往往比调模型更立竿见影。
3. CPU / 磁盘 offload。Accelerate 的 device_map="auto" 配合 offload 目录,或者 DeepSpeed 的 CPU offload,把暂时不用的参数放到内存。代价是速度掉得厉害,适合推理和低频任务。
4. 把 batch 降到 1。配梯度累积把有效 batch 补回来。如果 batch=1 都 OOM,就只剩换卡或者换模型两条路。
5. 拆分推理。长文本切成多段分别处理再合并;大图切块推理再拼接。
6. 换更大的卡或临时扩容。云上按需开机,训练完就关,适合一次性任务。
如何预防再次发生
1. 上线前做 batch size 探测。用二分或者倍增的方式,从 1 开始逐步加大,记录每一步的 max_memory_allocated,找到能稳定跑通的最大值,然后留出 15%~20% 余量。
2. 固定输入尺寸。动态 shape 会让 cuDNN 缓存多套 kernel,也会让显存分配器留下更多碎片。训练前把序列长度、图片尺寸统一 padding 到固定值。
3. 把显存打点做进训练日志。每个 logging 步记录 torch.cuda.max_memory_allocated(),并且每个 epoch 重置一次。显存曲线缓慢上升,多半就是泄漏。
4. 代码规范写进 review 清单:zero_grad(set_to_none=True)、loss 一定 .item()、验证和推理一定 torch.no_grad()、del 掉不用的大张量。
5. 保存 checkpoint 前把张量移到 CPU,别在 GPU 上直接 torch.save。
6. 把环境变量固化到启动脚本,比如 PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,避免每次手工加。
7. 给 OOM 加捕获与重试。小规模重试比整段训练崩掉划算:
```python
try:
out = model(**batch)
except torch.cuda.OutOfMemoryError:
torch.cuda.empty_cache()
降级:跳过这个 batch,或者切小后重试
```
8. 小规模冒烟测试进 CI。用真实数据的一个极小切片跑通完整的前向、反向、保存流程,能在几分钟内暴露绝大部分显存问题。
显存优化说到底是一道预算题:把模型权重、梯度、优化器状态、激活值四项分别算清楚,再对着卡容量做减法。绝大多数 OOM,都能在"降精度、降 batch、开检查点、换量化"这四招里找到出路。
