遇到显存不足时,直接减小批次通常能让程序继续运行,却没有回答显存究竟花在何处。若主要成本来自参数和优化器,缩小输入只能有限缓解;若峰值由高分辨率激活造成,删除少量权重也未必有效。可靠的规划方法是先列账目,再观察这些账目在训练时间线上的重叠,最后用框架统计修正理论值。
| 显存项目 | 主要决定因素 | 典型存在阶段 | 常见误判 |
|---|---|---|---|
| 模型参数 | 参数个数与数据类型 | 模型驻留期间 | 把磁盘文件大小当训练总量 |
| 梯度 | 需要更新的参数 | 反向计算至更新 | 认为推理占用可代表训练 |
| 优化器状态 | 算法保存的状态份数 | 多次更新之间 | 忽略动量与统计量 |
| 中间激活 | 批次、序列或分辨率、层宽 | 前向保存至反向消费 | 只统计有权重的层 |
| 工作区与缓存 | 算子实现、通信和分配器 | 随执行过程变化 | 把保留块都认作泄漏 |
第一笔账从张量形状与字节数开始
张量的理论容量等于各维度长度的乘积,再乘每个元素占用的字节。例如一个形状为 N×C×H×W 的浮点张量,先算元素总数,再根据数据类型换算。float32 通常每个元素四字节,float16 与 bfloat16 通常两字节,int64 则为八字节。这里的通常很重要,因为实际表示与框架对象还可能包含元数据或对齐开销。
单位也要统一。十进制 GB 以十的幂换算,GiB 以二的幂换算,不同工具显示方式可能不同。硬件标称、操作系统界面和训练框架若采用不同单位,数字会有可见差异。先换成字节比较,才能避免把单位差当成额外占用。理论预算还应留出余量,不能把标称容量每个字节都分配给张量。
- 卷积权重的主体数量可由输入通道、输出通道与卷积核尺寸相乘得到,偏置另计。
- 全连接层的主体参数由输入宽度与输出宽度相乘,是否有偏置取决于定义。
- 嵌入表通常与词表规模和嵌入宽度成正比,常成为大词表模型的重要常驻项。
- 归一化层可能有少量可训练参数,也可能维护不参与梯度的统计量。
模型文件为什么不能直接当作驻留参数
保存文件可能经过压缩,只包含部分权重,或以不同精度写出;加载后还可能发生类型转换、设备复制和权重共享。反过来,文件里也可能包含训练状态,而推理只加载其中一部分。最稳妥的做法是在模型构建后按参数对象的实际形状和类型统计,而不是从下载包大小推断。

训练常驻项会沿参数数量继续放大
设可训练参数共有 P 个,单份权重的理论容量容易计算,但训练还要为这些参数保存梯度。带动量的方法通常需要额外状态,保存一阶、二阶统计量的优化器会再增加多份与参数形状相近的张量。混合精度能够降低部分计算与激活成本,却可能保留高精度主权重或状态,因此不能简单把全部显存除以二。
估算优化器成本时,应查看它为每个可训练参数实际维护多少状态、这些状态使用什么类型,而不是背诵一个对所有配置都成立的倍数。
冻结一部分参数会减少相应梯度和优化器状态,但不一定消除前向计算产生的激活。参数共享也需要分清逻辑引用与真实存储:同一权重被多处调用,未必在显存中复制多份;相反,指数滑动平均、教师模型或多个检查点驻留,可能带来容易遗漏的完整副本。
激活预算跟着批次与空间尺寸走
自动求导为了计算梯度,需要保留前向阶段的某些输入、输出、掩码或索引。早期卷积层参数可能很少,但输出拥有较大的高和宽;序列模型的激活则会随批次、序列长度、隐藏宽度和层数增长。没有可训练权重的激活、池化或随机失活操作,也可能为了反向保存必要信息,所以无参数不等于零显存。
- 用代表性输入记录每个关键张量的形状与数据类型。
- 标记哪些张量必须跨越到反向阶段,哪些可及时释放。
- 考虑分支、残差和多输出使多个激活同时存活的情况。
- 分别测量只有前向、完成反向、执行更新三个阶段的峰值。
把反向显存粗略说成前向的某个固定倍数,只能用于非常早期的保守预估。实际峰值取决于计算图、是否原地操作、反向公式保存什么、梯度何时清除以及是否进行重计算。梯度检查点正是利用这个关系:不长期保留部分激活,反向时重新执行相应前向,以更多计算换取较低存储。
峰值来自同时存活,而不是历史累计
显存总量不应把运行期间出现过的所有张量机械相加。真正触发不足的是某一时刻仍存活的张量、通信缓冲和算子工作区之和。优化器更新阶段可能在梯度尚未释放时创建临时结果,分布式训练也可能出现梯度桶与计算重叠。沿时间线找峰值,比只在迭代结束时读取当前占用更有解释力。
| 看到的现象 | 优先检查 | 有针对性的方向 |
|---|---|---|
| 批次减半后明显下降 | 激活与输入 | 微批、累积、降分辨率 |
| 换优化器后大幅上升 | 参数状态 | 状态分片、卸载或算法选择 |
| 首次执行峰值异常高 | 上下文与算子工作区 | 预热后重复测量 |
| 循环中持续增长 | 计算图或张量被意外保留 | 检查容器、日志与梯度引用 |
理论值与监控值之间还有实现开销
设备上下文、已加载内核、算子库工作区、通信缓冲和内存分配器都会占据显存。框架为了减少反复申请,常把不再由活跃张量使用的块保留以便复用;因此已分配给张量的容量、分配器保留容量和设备进程总占用是三个不同概念。看到保留值较高,并不能直接认定存在泄漏。
测量时应先固定输入,完成必要预热,清零峰值统计,再分别运行前向、反向和更新。异步执行环境中,读取计数前还要确保目标阶段真正完成。若理论值与峰值仍差很多,可用内存快照或分析器查看具体算子和对象生命周期,而不是在训练结束后凭一个数字猜原因。
把优化动作对准对应账目
减小微批或输入尺寸主要压缩激活;梯度累积在保持较大有效批次的同时降低单次激活峰值,但会增加执行次数。较低精度可能减少部分张量容量,却需要核对主权重和状态类型。重计算用时间换激活空间,冻结参数减少梯度与优化器状态,分片或卸载则把参数相关账目分散到其他设备或主存。每项措施都有计算、通信或精度代价。
一份可用的显存预算最终应包含公式、输入假设、实测峰值和安全余量。模型结构变化后重新统计,部署推理时另建一份账,不沿用训练数字。能够说明每一类占用随哪个维度增长,显存不足就从突发错误变成可预见的容量问题;优化也不再是盲目缩小一切,而是有证据地削减真正形成峰值的那一项。
本文《训练前先算清显存账:参数、激活、梯度和优化器状态逐项估算》由 xkmchenmu 发布于 xkmchenmu Blog。 转载请保留原文链接并注明出处。
支付宝扫一扫