模型能装进显存,训练为什么还是 OOM?把 PyTorch 显存拆成 5 笔账

简介: 模型权重能够装进显存,不代表训练就能正常运行。本文拆解 PyTorch 训练中的参数、梯度、优化器状态、激活值和临时工作区,并给出显存监测代码及 OOM 排查顺序,覆盖梯度累积、混合精度、Gradient Checkpointing 和 LoRA 等常用方案。

有个挺常见的误区:模型权重只有十几 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。

相关文章
人工智能 JavaScript 开发工具
2227 1
Ubuntu 测试技术 网络安全
96 2
人工智能 Kubernetes Cloud Native
46 3
人工智能 自然语言处理 安全
46 0
弹性计算 小程序 C++
37 1
人工智能 弹性计算 Cloud Native
26 0
机器学习/深度学习 机器人 定位技术
51 0
存储 人工智能 安全
29 0
IDE Java Linux
40 0
存储 弹性计算 人工智能
39 0