模型能装进显存,训练为什么还是 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。

相关文章
人工智能 缓存 前端开发
4328 2
|
10天前
|
存储 弹性计算 缓存
阿里云服务器租赁费用:新版租赁收费标准及活动报价参考
本文更新了2026年阿里云全系列云服务器租赁活动报价,所有特惠资源均可前往阿里云活动中心选购,整体覆盖从个人入门到企业级高性能场景的全梯度需求。其中轻量应用服务器主打极致性价比,2核2G峰值200M带宽配置每日10点、15点限时抢购价仅38元/年,2核4G配置379元/年起;高性价比的经济型e实例、通用算力型u2i实例覆盖2核4G至4核32G全档位,适配开发测试与中小型企业业务;搭载英特尔至强6处理器的第九代c9i企业级实例算力较上代提升20%,支撑高并发生产环境,不同实例规格价差清晰,用户可根据自身业务负载与预算灵活选型。
1959 119
阿里云服务器租赁费用:新版租赁收费标准及活动报价参考
人工智能 JavaScript 开发工具
1682 1
|
11天前
|
人工智能 程序员 API
Codex 接入 DeepSeek-V4-Flash:还能补上识图,提供两套方案
Codex 接入 DeepSeek-V4-Flash 怎么配?本文覆盖 CLI 与桌面端,再用 qwen3-vl-flash 补识图,两套方案可直接照做
1516 13
缓存 人工智能 算法
444 0
|
8天前
|
编解码 弹性计算 云计算
MiniMax-H3 视频生成模型 — 一键部署与使用指南
MiniMax-H3是MiniMax开源的33B全模态视频生成模型,支持文生视频、图生视频、参考生视频三种模式,原生输出2K/15秒带立体声音频视频,已原生适配ComfyUI,并可通过阿里云计算巢一键部署。(239字)
|
17天前
|
云安全 人工智能 运维
阿里云联动百位企业安全专家,共识Agent防御最佳实践
当Agent成为新员工,你的安全边界在哪里?
1974 10
阿里云联动百位企业安全专家,共识Agent防御最佳实践
|
9天前
|
人工智能 API 开发工具
2026 零基础本地 AI 漫剧完整实操教程(8G 笔记本显卡可用|附可直接复制命令与代码)
本方案提供完全离线、本地运行的漫剧全自动制作流程:RTX3060/4050 8G显卡即可驱动,涵盖Qwen写分镜→ComfyUI统一角色绘图→LTX2.3图生微动画→Qwen3-TTS本地配音→FFmpeg自动合成,全程无水印、免API、不限次。专为低显存优化,解决变脸、闪烁、爆内存三大痛点。(239字)