结论先行:CUDA OOM是训练最常见报错,按降低batch size、开混合精度、开梯度检查点、分片优化器四步走,80%的情况都能解决。
一、OOM的四个来源
模型权重、梯度、优化器状态、激活值。训练时这四项都要占显存,7B模型FP16训练仅前三项就约需56GB。激活值随batch size和序列长度线性增长,是最容易爆的部分。
二、四步排查法
第一步:降低batch size
从32降到8再降到1,看是否还OOM。如果降到1还OOM,说明模型本身就装不下,需要并行或量化。
第二步:开混合精度
torch.cuda.amp.GradScaler开启FP16,显存立减约一半。这是最简单有效的优化,几乎所有训练都应该开。
第三步:开梯度检查点
model.gradient_checkpointing_enable(),用重算换显存,速度慢约20%但显存省很多。长序列训练必开。
第四步:分片优化器
用FSDP或DeepSpeed ZeRO-3把优化器状态和梯度切分到多卡,单卡显存压力大幅降低。
三、推理OOM怎么办
推理OOM通常是KV Cache太大。降低max_new_tokens、用vLLM的PagedAttention、或开INT4量化即可。70B模型INT4推理约需40GB显存。
五、OOM后定位方法
OOM报错会打印已分配的显存和尝试分配的显存,根据这个数字判断差多少。用torch.cuda.memory_summary()打印详细显存分配表,看哪个tensor占用最大。常见元凶是Dataset没开num_workers导致预加载太多batch、或者optimizer没offload。定位后针对性优化,不要盲目加卡。
除了四步法,还有两个技巧:用checkpoint功能自动保存中间状态,OOM后从最近检查点恢复而非从头跑;用torch.cuda.empty_cache()清理碎片显存。
不确定显存怎么分配,可在GPU算力平台先用小卡跑通再逐步升配。




