有个挺常见的误区:模型权重只有十几 GB,显卡有 24GB 显存,那训练应该没问题。
真跑起来却经常第一步就报:
torch.OutOfMemoryError: CUDA out of memory
这不一定是环境坏了,也不一定要先换更大的卡。模型权重只是显存账单里的一项,训练时还要放梯度、优化器状态、激活值和各种临时缓冲区。推理能跑,不代表训练也能跑;空模型能加载,不代表第一个 batch 能完成反向传播。
与其看到 OOM 就反复清缓存,不如先把显存到底花在哪里算清楚。
一、训练显存主要有 5 笔账
可以先用一个不严谨但很好用的式子理解:
训练峰值显存 ≈ 参数 + 梯度 + 优化器状态 + 激活值 + 临时工作区
1. 模型参数
参数量乘以每个参数占用的字节数,就是最基础的一笔。
| 参数精度 | 每个参数的理论存储量 |
|---|---|
| FP32 | 4 Byte |
| FP16 / BF16 | 2 Byte |
| INT8 | 1 Byte |
| INT4 | 0.5 Byte |
例如 7B 参数只按 BF16 权重粗算,大约是:
7 × 10^9 × 2 Byte ≈ 14 GB
但这 14GB 只能说明权重本身大概能不能放进去,不能说明能不能训练。
量化模型也不能简单按表格数字下结论。实际还可能存在量化元数据、分组缩放参数、未量化层、计算时反量化缓冲区等额外占用。
2. 梯度
需要训练的参数通常还要保存对应梯度。如果全部参数都参与训练,梯度可能又接近一份参数规模。
LoRA 之所以能明显降低训练门槛,关键并不只是“用了低秩矩阵”,而是冻结了大部分原始参数,只给少量可训练参数保存梯度和优化器状态。
所以判断一个微调任务时,不能只问模型是多少 B,还要问:
- 全参数训练还是参数高效微调;
- 哪些模块参与训练;
- 可训练参数到底有多少。
可以直接查看:
trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
total = sum(p.numel() for p in model.parameters())
print(f"trainable: {trainable:,}")
print(f"total: {total:,}")
print(f"ratio: {trainable / total:.4%}")
3. 优化器状态
Adam、AdamW 这类优化器会为可训练参数维护一阶矩和二阶矩。只看权重精度,很容易漏掉这里。
以最简单的 FP32 参数训练粗算:
参数 4 Byte
+ 梯度 4 Byte
+ Adam 一阶矩 4 Byte
+ Adam 二阶矩 4 Byte
= 约 16 Byte / 可训练参数
混合精度训练的情况会更复杂。有些方案还会保留 FP32 主权重,具体占用取决于框架、优化器和训练策略。因此“开启 FP16 后显存一定减半”这个说法并不准确。
4. 激活值
激活值通常是最难提前估算、也最容易突然把显存顶满的一项。
它与下面这些变量有关:
- batch size;
- 序列长度;
- 图像分辨率;
- 网络层数和隐藏维度;
- 是否保存中间结果用于反向传播;
- 算子的具体实现。
这也是为什么同一个模型,batch size 从 1 调到 2,显存可能不只是多一点;把上下文长度或图片边长翻倍,显存也可能明显上升。
推理阶段不需要保存完整反向传播所需的激活,但大模型生成还会有 KV Cache。它又与 batch、序列长度、层数和注意力结构有关。训练和推理不能共用一张简单的显存表。
5. 临时工作区与显存分配器
矩阵乘法、卷积、注意力和自定义 CUDA 算子运行时,可能临时申请额外工作区。PyTorch 的缓存分配器还会保留一部分已经申请过的显存,以减少频繁分配带来的开销。
所以你会看到几组不同的数字:
memory_allocated:张量当前实际占用;memory_reserved:PyTorch 已向 CUDA 申请并保留的显存;nvidia-smi:进程在设备侧呈现的总体占用。
它们不相等是正常的。
二、先测峰值,不要只盯着 nvidia-smi
可以在最小训练步骤前后加一组记录:
import torch
def gib(value):
return value / 1024**3
torch.cuda.reset_peak_memory_stats()
# 放入一次最小训练步骤
# loss = train_one_step(...)
torch.cuda.synchronize()
print(f"allocated: {gib(torch.cuda.memory_allocated()):.2f} GiB")
print(f"reserved: {gib(torch.cuda.memory_reserved()):.2f} GiB")
print(f"peak allocated: {gib(torch.cuda.max_memory_allocated()):.2f} GiB")
print(f"peak reserved: {gib(torch.cuda.max_memory_reserved()):.2f} GiB")
想看更细的分配情况,可以打印:
print(torch.cuda.memory_summary())
记录时至少把这些变量一起写下来:
GPU 型号与显存:
PyTorch 版本:
模型与参数精度:
训练方式:全参数 / LoRA / 其他
batch size:
序列长度或图像分辨率:
梯度累积步数:
是否开启 AMP:
是否开启 gradient checkpointing:
峰值 allocated:
峰值 reserved:
没有这些上下文,单独一句“这个模型需要多少显存”通常很难得到可靠答案。
三、OOM 以后,按这个顺序处理
第一步:先把最小 batch 跑通
先把单卡 batch size 调到 1,缩短序列或降低图片分辨率,确认最小训练闭环能够完成:
读取数据 → 前向传播 → 计算 loss → 反向传播 → optimizer.step → 清零梯度
如果 batch size 已经是 1 仍然 OOM,继续看模型本身、序列长度、激活值和优化器,而不是一直减 batch。
第二步:用梯度累积补回有效 batch
显存不够时,可以用多个小 batch 累积梯度:
有效 batch size
= 单卡 batch size × 梯度累积步数 × GPU 数量
例如单卡 batch 为 1,梯度累积 8 步,两张 GPU,在常见的数据并行语义下有效 batch 可以按 16 理解。
不过梯度累积只是降低单步激活占用,不会让训练完全等价于任意的大 batch。BatchNorm、学习率调度、梯度裁剪和日志步数都要重新确认。
第三步:开启混合精度
支持 BF16 的 GPU 上,可以优先评估 BF16;其他场景再考虑 FP16。一个简化写法是:
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
outputs = model(**batch)
loss = outputs.loss
混合精度通常能减少部分参数、激活和计算开销,但实际节省多少取决于训练框架。改完以后要重新测峰值,也要看 loss 是否正常,有没有 NaN 或溢出。
第四步:用 gradient checkpointing 换显存
Gradient checkpointing 不保存所有中间激活,而是在反向传播时重新计算一部分结果。
它解决的是激活值占用,代价是增加计算时间。模型越深、序列越长,收益通常越明显,但不能把它当成没有成本的开关。
第五步:减少真正参与训练的参数
如果目标不是从头训练整个模型,可以考虑冻结主干、LoRA、QLoRA 或其他参数高效微调方案。
这一层解决的是梯度和优化器状态,而不只是激活值。选择前要先明确任务是否允许参数高效微调,不能为了省显存把训练目标也改了。
第六步:最后再处理碎片和残留引用
torch.cuda.empty_cache() 只会释放缓存分配器中当前未被张量使用的缓存,不会释放仍被 Python 对象引用的张量,也不会降低模型本身需要的峰值显存。
如果循环中把带计算图的 loss 或输出一直塞进列表,显存也可能逐步增长:
# 容易保留计算图
history.append(loss)
# 只保存数值
history.append(loss.item())
如果 OOM 信息显示 reserved 明显大于 allocated,可以进一步检查分配模式和碎片问题。但这应放在 batch、长度、精度和训练参数之后,不要一开始就靠环境变量碰运气。
四、几个经常无效的操作
反复执行 empty_cache()
它可能让显存显示暂时下降,但不会改变下一次前向和反向传播的真实峰值需求。
只降低 DataLoader 的 num_workers
num_workers 主要影响 CPU 侧的数据加载进程、内存和吞吐,通常不是 CUDA OOM 的直接原因。它可能改善主机内存压力,却不等于减少模型训练显存。
看见 OOM 就换 CUDA
版本不兼容和显存不足是两类问题。错误已经明确写着申请多少、现有多少时,先做显存账本;如果是算子不存在、二进制不兼容或扩展编译失败,再查 CUDA 和 PyTorch 版本。
只记录“24GB 能跑”
这句话几乎没法复现。至少还要补上模型版本、精度、训练方式、batch、长度或分辨率、是否使用 checkpointing,以及峰值显存。
五、最终要留下的不是一句经验,而是一张账单
一次比较完整的训练显存记录可以长这样:
模型:
可训练参数量:
权重精度:
优化器:
单卡 batch:
梯度累积:
序列长度 / 图像分辨率:
AMP:
gradient checkpointing:
峰值 allocated:
峰值 reserved:
是否完成保存和重新加载 checkpoint:
判断显存够不够,最好不要从“模型文件多大”开始,也不要只问别人某张卡能不能跑。
先把参数、梯度、优化器状态、激活值和临时工作区拆开,再用最小训练步骤测峰值。这样即使最后确实需要更大的显存,也能知道钱花在了哪一笔,而不是换完卡继续 OOM。