目录
第四课中,DataLoader 把形状为 $[B,T]$ 的 Token IDs 送进 GPU。接下来常见的问题是:模型文件明明只有十几 GB,为什么训练时几十 GB 显存仍然不够?
因为 GPU 里放的不只是模型参数。训练还要保存梯度、优化器状态和反向传播需要的激活值。
模型大小只是显存账单的一部分;训练显存更像“四笔账”的总和。
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 参数外,往往还有:
- 梯度:通常每个参数约 2 Bytes。
- FP32 主权重:用于更稳定地更新参数,约 4 Bytes。
- 一阶动量 $m$:约 4 Bytes。
- 二阶动量 $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 | 可训练参数及相关状态 | 适合微调,不等于完整预训练 |
这些方法没有“免费午餐”:省显存通常会换来更多计算、通信、工程复杂度,或限制训练方式。
容易混淆的三件事
- Gradient Accumulation 不会让一次前向能容纳更长的序列。 它降低的是 Micro Batch,而不是单条样本的长度。
- 量化模型能省推理显存,不代表可以直接用同样方式完整训练。 训练还涉及梯度和优化器状态。
- 显存占用不只由参数量决定。 长上下文和大 Batch 会让激活值成为重要部分。
估算显存时看什么
| 类别 | 主要受什么影响 |
|---|---|
| 参数 | 参数量、参数精度 |
| 梯度 | 可训练参数量、梯度精度 |
| 优化器状态 | 优化器类型、状态精度 |
| 激活值 | $B$、$T$、层数、Hidden Size |
| 临时开销 | Kernel、通信、框架、显存碎片 |
第五课复习总图
下一课,我们可以继续看:一张 GPU 放不下模型时,数据并行、张量并行和流水线并行分别在拆什么?