离线模型推理时遇到 CUDA error: out of memory,通常不是显存真的被模型占满,而是推理脚本里保留了不该保留的中间状态,或者批处理大小设置超过了当前显卡容量。排查时先确认两个方向:推理路径是否关闭了梯度记录,以及批处理张量是否被一次性放大。
离线推理遇到显存不足,先检查是否对模型执行了 requires_grad_ 或在非训练模式下未使用 torch.no_grad()。若已关闭梯度,则逐步下调 batch_size,并记录每个批次显存峰值。直接加大显卡或升级驱动不是首选,先验证脚本中是否存在显存泄漏和缓存堆积。
先分清推理场景还是训练场景
标题同时提到“梯度释放”和“批处理大小调整”,说明这条报错可能出现在两种场景:一种是把训练好的模型拿去做批量推理,另一种是微调或继续训练时显存不足。两者处理思路不同。
纯推理场景下,模型参数梯度不应存在。使用 PyTorch 时,需要在推理前加上 model.eval(),并在推断的代码块外层包裹 torch.no_grad()。只调用 model.eval() 不会关闭梯度计算,它只改变 dropout 和 batch norm 的行为。
model.eval()
with torch.no_grad():
output = model(input_batch)训练场景下,模型需要保留梯度和中间激活值,显存占用通常是推理模式的数倍。如果训练时出现 OOM,缩小 batch_size 代价最小,但还需要检查优化器状态、梯度累积变量等是否存在额外占用。
释放梯度和中间张量的具体做法
如果代码在循环中每批次都累积张量,比如把每个 output 都保存到 list 里,list 会持续占用显存直到进程结束。先检查是否有这类累积逻辑。
- 确认输入张量是否需要梯度,若不需要,将其从计算图分离:
input_batch = input_batch.detach()。 - 每个批次结束后,对可能残留的计算图调用
del variable并torch.cuda.empty_cache()。但注意 empty_cache 只释放未使用的缓存块,并不是把显存归还给系统,频繁调用反而降低速度。 - 如果使用自定义循环做梯度累积,确认在反向传播后调用了
optimizer.zero_grad(),否则梯度会叠加在旧变量上,显存同步增长。 - 对于推理,尽可能把数据分片,用生成器或队列逐批读取,避免把整个数据集的张量一次性放到 GPU。
批处理大小调整的显存估算
显存占用主要由三部分组成:模型参数、优化器状态(训练时)、前向过程中的激活值。激活值与 batch_size 近似线性相关,所以把 batch_size 从 32 降到 16,激活显存可以减半,但参数和优化器状态不变。
做一个快速估算:先设置 batch_size=1,记录当前的显存占用 base_mem,再用 batch_size=8 运行并记录峰值 peak_mem,估算单个样本的增量 delta = (peak_mem - base_mem) / 7。这样可以反推当前显卡可接受的最大 batch_size,而不是盲目从 32 开始逐次减半。
验证顺序建议
按照先排除脚本问题、再调整 batch_size、最后考虑模型改动的顺序排查。
- 先打开 nvidia-smi 观察 GPU 显存占用和进程,确认没有旧进程占住显存。
nvidia-smi - 在代码中给每个关键块加显存峰值打印,确认 OOM 发生在哪一步。
- 若推理时仍 OOM,检查是否在模型 forward 里意外调用了
backward()或保留了 loss。 - 若训练时 OOM,先试 batch_size 减半,观察 loss 曲线是否仍能正常下降,再决定是否需要调整学习率。
- 如果 batch_size 已经很小,仍然 OOM,考虑将模型切到半精度推理,或用梯度检查点来降低激活值缓存。
还需要注意,不需要把所有显存占用都优化到极低。PyTorch 的 CUDA 缓存机制会预留一部分显存,通常表现为减少 batch_size 后显存占用没有立刻下降,这是正常现象。真正需要关注的是 OOM 是否消失,以及任务能否稳定跑完。
常见的两个追问
为什么显存足够,还是报 out of memory?
PyTorch 的显存分配器会预分配缓存块,这个缓存块不会自动归还给 CUDA。如果之前跑过一个占用较高的任务,新任务申请显存时,CUDA 可能报告不足,尽管 nvidia-smi 显示空闲。可用 torch.cuda.empty_cache() 释放缓存,或更换测试任务前先重启 Python 进程。
batch_size 调小后,模型输出变差怎么办?
先确认这属于训练过程,推理时 batch_size 不影响单样本结果。训练时减小 batch_size 会改变 batch norm 统计和梯度噪声,可尝试同时降低学习率,或使用梯度累积来近似原来的 batch_size。但梯度累积不会减少显存占用,它只用来模拟更大的 batch_size。