训练大模型为什么这么吃显存?

参数、梯度、优化器状态与激活值,到底谁占得最多

目录

  第四课中,DataLoader 把形状为 $[B,T]$ 的 Token IDs 送进 GPU。接下来常见的问题是:模型文件明明只有十几 GB,为什么训练时几十 GB 显存仍然不够?

  因为 GPU 里放的不只是模型参数。训练还要保存梯度、优化器状态和反向传播需要的激活值。

模型大小只是显存账单的一部分;训练显存更像“四笔账”的总和。

大模型训练显存构成

训练显存的四笔账,点击查看原始 SVG

1、先分清训练和推理

  推理只需要使用已经训练好的参数完成前向计算;训练还要反向传播并修改参数,因此需要保存更多中间数据。

内容 推理 训练
模型参数 ✓ ✓
激活值 部分、短暂 ✓,供反向传播使用
梯度 — ✓
优化器状态 — ✓
KV Cache 生成时常用 通常不是主要项

  因此,同一个模型能在一张卡上推理,不代表也能在这张卡上完整训练。

2、参数本身占多少显存

  参数显存可以粗略估算为:

\[参数量 \times 每个参数的字节数\]
数据类型 每个参数大小 10 亿参数约占
FP32 4 Bytes 4 GB
FP16 / BF16 2 Bytes 2 GB
INT8 1 Byte 1 GB
INT4 0.5 Byte 0.5 GB

  一个 7B 模型若以 BF16 保存权重,参数本身大约需要 $7\times2=14$ GB。这里的 GB 是便于理解的近似值,实际还会受单位换算、量化分组和框架开销影响。

3、训练还要保存什么

  以常见的 AdamW 和混合精度训练为例,除 BF16 参数外,往往还有:

  1. 梯度:通常每个参数约 2 Bytes。
  2. FP32 主权重:用于更稳定地更新参数,约 4 Bytes。
  3. 一阶动量 $m$:约 4 Bytes。
  4. 二阶动量 $v$:约 4 Bytes。

  于是每个参数可能需要约 $2+2+4+4+4=16$ Bytes。7B 模型仅这些“按参数线性增长”的状态,粗略就达到 112 GB。

  不同框架和训练配置会不同,所以 16 Bytes 不是永恒公式,而是一种常用估算模型。

4、激活值为什么会变大

  前向计算中,每一层都会产生中间结果。反向传播需要它们计算梯度,因此训练时不能马上全部丢掉。

  激活值通常随这些因素增长:

  • Batch Size $B$ 越大,占用越多。
  • Sequence Length $T$ 越长,占用越多。
  • Hidden Size 越大,占用越多。
  • Transformer 层数越多,需要保留的中间结果越多。

  可以先记住一个方向性的关系:

\[Activation\ Memory \propto B\times T\times d_{model}\times Layers\]

  某些 Attention 中间矩阵还可能与 $T^2$ 相关。因此上下文长度翻倍时,显存和计算量不一定只翻倍。

5、一个简化算例

  假设训练一个 7B 模型,使用 BF16 参数、BF16 梯度和带 FP32 状态的 AdamW:

项目 每参数字节 7B 粗略占用
BF16 参数 2 B 14 GB
BF16 梯度 2 B 14 GB
FP32 主权重 4 B 28 GB
Adam 一阶动量 4 B 28 GB
Adam 二阶动量 4 B 28 GB
合计(未含激活等) 16 B 112 GB

  还没算激活值、临时 Buffer、通信缓存和框架本身的开销,单卡 80 GB 已经放不下。这就是大模型训练需要多卡并行和显存优化的直接原因。

6、常见省显存方法

方法 主要省掉什么 代价或特点
混合精度 参数、梯度、激活 需要处理数值稳定性
Gradient Checkpointing 激活值 反向时重算,速度变慢
Gradient Accumulation 单步激活值 多次小 Batch 后再更新
ZeRO / FSDP 参数、梯度、优化器状态 分片到多张 GPU,增加通信
8-bit Optimizer 优化器状态 依赖实现与任务适配
LoRA / QLoRA 可训练参数及相关状态 适合微调,不等于完整预训练

  这些方法没有“免费午餐”:省显存通常会换来更多计算、通信、工程复杂度,或限制训练方式。

容易混淆的三件事

  1. Gradient Accumulation 不会让一次前向能容纳更长的序列。 它降低的是 Micro Batch,而不是单条样本的长度。
  2. 量化模型能省推理显存,不代表可以直接用同样方式完整训练。 训练还涉及梯度和优化器状态。
  3. 显存占用不只由参数量决定。 长上下文和大 Batch 会让激活值成为重要部分。

估算显存时看什么

类别 主要受什么影响
参数 参数量、参数精度
梯度 可训练参数量、梯度精度
优化器状态 优化器类型、状态精度
激活值 $B$、$T$、层数、Hidden Size
临时开销 Kernel、通信、框架、显存碎片

第五课复习总图

大模型第五课复习总图

第五课复习总图,点击查看原始 SVG

  下一课,我们可以继续看:一张 GPU 放不下模型时,数据并行、张量并行和流水线并行分别在拆什么?