☰
大模型训练显存估算与混合精度实战指南
2026/9/29 16:21:15 网站建设 项目流程

跑大模型训练的人,谁没被 “CUDA out of memory” 那行红字折磨过?大模型训练做到中后期,显存就不只是“够不够用”的问题,而是必须带着计算器去预判的问题。这篇是这个系列的第三篇,我把大模型训练的显存估计和混合精度训练放在一起聊,因为它们在实战中几乎总是绑在一起出现:显存不够用的时候,第一反应不是换卡,而是先看看有没有把精度浪费在不该浪费的地方,再考虑怎么用混合精度把空间省出来。

这篇文章先给出一套能直接套用的显存估算方法,把 7B、13B 这种常见规模算给你看;再讲清 FP32、FP16、BF16 在训练中的真实差异,以及大模型训练混合精度为什么既能省显存又能提速;最后是一份带代码的实操清单和踩坑记录。不管你正在准备训自己的模型,还是单纯想搞清楚为什么隔壁团队能在一张 A100 上跑起来而你不能,这篇都值得存一下。

1. 训练显存不是只有模型权重

1.1 静态开销:权重、梯度、优化器状态三件套

绝大多数人对显存的直观理解就是“模型多大,显存就多大”。这个直觉在推理场景下基本成立,但在训练时误差很大。训练时显存实际上由四个部分构成:模型权重、梯度、优化器状态、激活值。前三个属于静态开销,基本不随训练轮次变化;最后一个是动态大头,会在前向过程中持续波动。

先说静态三件套。模型权重必须常驻显存,这个不用解释。梯度呢?反向传播逐层算出梯度后并不是马上丢掉的,优化器要做整体更新,分布式训练里要跨卡聚合,梯度在权重更新之前必须完整保存在显存里。梯度的数据量和模型权重完全一样,相当于把模型本体再存了一份。优化器状态则是很多人容易漏掉的一项,以最常用的 AdamW 为例,它会为每个参数维护两个状态:一阶动量(momentum)和二阶动量(variance),而且这两个状态通常都保留为 FP32 精度。也就是说,一个 7B 模型用 AdamW 训练,光是优化器状态就相当于把 7B 个 FP32 数存了两次,算下来比模型本身还大。

给你一个厨房备菜的类比:权重是你灶台上正在用的食材,梯度是切好还没下锅的菜码,优化器状态是贴在墙上的配方笔记,三者缺一不可,但只有真正做完菜你才会发现,最占台面的其实是临时堆在那里等着下锅的备菜——那就是激活值。

1.2 激活值:被忽略的显存大头

激活值(activation)指的是前向传播过程中每一层算出来的中间张量。Transformer 的每一层要做自注意力、LayerNorm、MLP 等一连串计算,反向传播时需要这些中间结果才能回传梯度,所以不能算完就扔。层数越深、序列越长、batch 越大,激活值的体积就滚得越大。

很多人的 OOM 不是发生在加载模型的时候,而是训练跑到第几百步才突然炸掉,原因就在这里:加载阶段只有静态三件套,训练开始后激活值叠加上来,峰值一冒,显存就崩了。明白这个结构之后,你拿到一个新模型,就不该再笼统地问“这个模型要多少显存”,而应该拆成两个问题:静态三件套占多少,激活值峰值占多少。这两个问题的估算方法,就是下一章要说的重点。

2. 显存估算:先算静态,再算动态

2.1 一张表看懂精度与字节数

在估算之前,先记住不同数据类型对应的字节数。大模型训练里最常见的精度就是 FP32、FP16、BF16,偶尔还会见到 INT8 精度的优化器。

数据类型字节数指数位尾数位训练中的典型用途
FP324823优化器状态、主权重
FP162510旧卡上的混合精度前向计算
BF16287新卡上的混合精度前向计算
INT81无无量化优化器状态

一个参数占用几个字节,拿参数量 P 一乘就出来了。假设参数量是 700 亿,也就是 P=70×10^9,那么 FP32 下光权重就是 70×4=280GB,看到这个数字你马上就能理解,为什么 70B 模型训练几乎不可能用纯 FP32 单卡完成。

2.2 常见训练组合的系数表

大模型训练里最常见的优化器组合就那么几种,我直接把“每参数字节数”的系数列出来。这里的 P 表示参数量,单位是字节。

训练配置权重梯度优化器状态每参数总字节7B 模型静态显存
FP32 + SGD4P4P08P约 56GB
FP32 + AdamW4P4P8P16P约 112GB
混合精度 + AdamW2P2P8P12P约 84GB

有点反直觉的是,混合精度训练虽然把权重和梯度降到了 FP16/BF16,但优化器状态依然是 FP32 的两份,所以总字节系数是 2+2+8=12P,而不是单纯的一半。这也是为什么大家普普通通说“7B 混合精度训练需要 80 多 GB 显存”的来源。

顺便说一句,纯 FP32 的 AdamW 需要 16P,一个 7B 模型就吃掉 112GB,单卡基本没戏;而混合精度把权重和梯度各减半,降到 84GB,配合激活值控制和梯度检查点,就有机会塞进单张 80GB 的卡里。省下来的这 28GB,就是混合精度的最大意义之一。

2.3 手把手算一次 7B、13B、70B

套用 12P 这个系数,你拿计算器直接乘就行:

  • 7B 模型:12 × 7 = 84GB,不含激活值;
  • 13B 模型:12 × 13 = 156GB,不含激活值;
  • 70B 模型:12 × 70 = 840GB,不含激活值。

看到 840GB 不要慌,这是单卡全量放置的数字;实际训练 70B 都会走多卡并行加 ZeRO 分片,把权重、梯度、优化器状态平均拆到每张卡上。后面第 5 章会讲怎么拆。

还要特别注意,这里的数字只是静态开销。一个 7B 模型的 84GB 算出来,如果手里只有一张 80GB 的 A100,听起来勉强够,但一旦前向传播的激活值冲上来,100GB 照样瞬间爆掉。所以估算永远是“先算静态,再给动态留余量”,我自己的习惯是按照公式结果的 1.2 到 1.3 倍来选卡,宁可富余别赌运气。

2.4 激活值怎么估:经验公式加实测

激活值没有静态三件套那么规整,但可以按 Transformer 的常见实现估一个量级。经验上,每个 Transformer 层要保存的中间激活元素数大致是:

$$batch \times seq_len \times hidden_size \times k$$

其中 k 是一个和实现相关的经验系数,常见实现不开梯度检查点时大约在 30 到 40 之间,按 2 字节存储。以 Llama-7B 级别的配置为例,hidden_size 取 4096,层数取 32,序列长度 2048,batch_size 为 1,带入后:

$$1 \times 2048 \times 4096 \times 32 \times 34 \times 2 \approx 18GB$$

也就是说,单序列长度 2048 的 7B 模型,光激活值峰值就逼近 18GB。一旦 batch 开到 4,这一项就变成 72GB,整张 80GB 的卡直接就没有余量了。这还只是保守估计,实际因为 dropout mask、临时变量、注意力矩阵的实现差异,峰值可能更高。

这个数字也解释了一个现象:为什么大模型训练里大家那么怕“长序列”,sequence length 一翻倍,激活值和序列长度的平方项一起涨,对显存的杀伤力远比增加 hidden size 要大。所以做显存预算的时候,激活值必须当成头号变量来对待。最可靠的方法是本地跑一个小配置,用 PyTorch 的 torch.cuda.max_memory_allocated() 直接量出真实峰值,经验公式只能用来做出发前的预判。

3. 混合精度训练的原理与选型

3.1 FP16 和 BF16 到底差在哪

混合精度训练的核心思路,是把前向计算和梯度计算中那些对精度不太敏感的算子从 FP32 降到半精度,从而省显存、降带宽压力、利用 Tensor Core 加速。但半精度有两个兄弟,FP16 和 BF16,长相相近,脾气完全不同。

FP16 是 1 位符号、5 位指数、10 位尾数,最大能表示到 65504,最小正规数大约是 6.1e-5。问题就出在这个范围上:训练时反向传播的梯度经过层层乘法链式法则,很多数值会掉到 1e-5 以下,一旦小于最小正规数,FP16 里存的就是 0。梯度变成 0,意味着权重根本不会更新,模型原地踏步。反过来,碰到中间结果稍大一点,超过 65504 就变成无穷大,直接炸掉整个训练。上下两个方向都容易出问题,这是 FP16 最麻烦的地方。

BF16 是 1 位符号、8 位指数、7 位尾数。指数位和 FP32 一样多,所以它的表示范围和 FP32 几乎一致,极小值到极大值都能覆盖,天然不会出现 FP16 那种动不动溢出、下溢的情况。代价是尾数只有 7 位,十进制有效数字大约只有 2 到 3 位,精度比 FP16 还低。这就带来一个很有意思的局面:BF16 范

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询