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

相关文章
|
18天前
|
人工智能 安全 测试技术
【AI时代软件项目管理系列】4. AI 时代的软件项目启动:从立项评估到 AI 可行性分析
AI 并没有让瀑布型项目失效,尤其在政企、定制化和交付型项目中,合同范围、预算、里程碑、阶段评审和最终验收仍然需要明确治理。真正需要改变的是传统严格串行的管理方式。AI 时代的瀑布项目应保留阶段、基线和责任机制,同时引入快速反馈、持续验证和 AI/Agent 赋能,让需求、设计、开发、测试和交付形成更高效的闭环,实现“治理不失控、执行更敏捷”。
117 1
|
18天前
|
安全 前端开发 JavaScript
基于 BitB 技术的 RecruitTrap 招聘主题规模化钓鱼攻击机理与全域防御研究
本文剖析2026年RecruitTrap招聘钓鱼攻击:利用BitB(浏览器内嵌)伪造登录弹窗、实时中继MFA验证码,两个月仿冒50+企业、投放3000+钓鱼链接,专攻营销岗谷歌/Facebook账号。揭示传统URL校验与软MFA失效本质,提出以FIDO2通行密钥为核心的分层防御体系。(239字)
65 1
|
20天前
|
人工智能 算法 搜索推荐
3款AI引擎引用监测实录:固定词库追踪竞品内容被引变化
本文提供一套零成本AI搜索竞品监测方案:基于10-15个固定提问词,每周30分钟手动测试豆包、DeepSeek、秘塔三款引擎,记录竞品出现频次、引用平台、内容类型及语境。通过交叉印证分析,识别其内容策略与平台布局,指导自身内容优化方向。
91 1
|
监控 Python
手把手教你用 Python 制作一场炫酷烟花秀
本篇文章,带大家用 Python 制作一个炫酷烟花秀,来迎接即将到来的元旦佳节。开始之前先看一下最终效果
手把手教你用 Python 制作一场炫酷烟花秀
|
20天前
|
Shell API 调度
DeepSeek Harness 一切皆插件:开源Agent框架强在哪,怎么装
DeepSeek Harness 开发者预览版 2026 年 8 月开源,口号"一切皆插件"。本文拆解插件化架构的 4 个好处,并带你跑通安装命令。
1249 3
DeepSeek Harness 一切皆插件:开源Agent框架强在哪,怎么装
|
20天前
|
人工智能 JavaScript API
#Codex接入DeepSeek-V4-Flash完整实操指南:搭配qwen3-vl-flash补齐图像识别两套落地方案
在AI编程工具快速普及的当下,Codex作为终端与桌面端一体化代码智能体,凭借读写本地文件、执行终端命令、多步骤代码重构、工具调用等能力,成为大量开发者日常开发的核心辅助工具。但原生Codex依赖官方模型订阅,长期使用成本较高,不少开发者开始寻找性价比更高的第三方推理基座,DeepSeek-V4-Flash凭借原生适配Codex所需的Responses API、百万级上下文窗口、低廉的Token计费标准、完善的Agent工具调用能力,成为替换原生模型的最优选择之一。
228 2
|
20天前
|
存储 监控 API
基于 RAG + LangChain 搭建企业级私有知识库问答系统(2026 实战版)
本文是作者基于多个企业RAG知识库落地经验的实战总结,提供完整可运行代码与十年避坑指南。涵盖文档解析、混合检索、向量存储、DeepSeek接入、结果重排、拒答机制及效果评估,助你构建本地可运行、生产可扩展的企业级私有知识库系统。(239字)
304 1
|
20天前
|
数据采集 人工智能 算法
45条AI引用源实测:内容平台权重分布与信息块拆解
本文拆解豆包AI的45条引用源,揭示CSDN、头条、搜狐占国内引用近半;剖析被高频引用的CSDN文章结构参数(如数字密度、列表数、H2标题),提出“平台推荐→AI抓取→被引用”链路及可复现的监测方法。
152 1
|
20天前
|
缓存 安全 程序员
智谱GLM-5.3发布同基座纯靠后训练编程涨50还点亮网安技能树
智谱 8 月 14 日发布 GLM-5.3,与 5.2 同基座、纯后训练,编程内部基准提升 50%,CyberGym 拿下开源第一,两周后开源权重
智谱GLM-5.3发布同基座纯靠后训练编程涨50还点亮网安技能树
|
20天前
|
人工智能 缓存 算法
最新版通义千问(Qwen3.8‑Max‑Preview)功能介绍
随着AI智能体从简单问答走向长周期自主任务执行,市场对大模型的综合能力提出更高要求,不仅需要强悍的文本推理,还需要原生多模态理解、百万级超长上下文、稳定的链式工具调用、大型工程项目完整交付能力。Qwen3.8‑Max‑Preview作为通义千问系列新一代旗舰预览基座,总参数规模达到2.4万亿,采用MoE混合专家架构,定位为**代码工程+专业办公**双核旗舰模型,完成从纯文本向原生多模态的跨越,在长链路Agent自治、全栈软件开发、大批量复杂文档分析、多模态专业办公场景实现能力跨越式提升。该预览版本率先开放于百炼平台,支持Token Plan订阅模式调用,适配OpenClaw、Hermes Ag
675 2