跳到主内容
快讯直播
AI智模界
教程

CUDA out of memory 显存不足的排查与优化

报错现象

报错原文(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=Truenum_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、开检查点、换量化"这四招里找到出路。

AI 生成本文由 AI 基于公开信息自动生成,仅供参考。