📑 本页目录(点开跳转)
03 · 显存与带宽墙
⏱ 30 分钟 | ⭐⭐ 理解这一章,后面十几章都会变简单
🎯 一句话
你的 GPU 不是在算,大部分时间是在等数据。 这一章讲清楚数据存在哪、搬一次要多久, 之后 FlashAttention、量化、KV Cache、算子融合 —— 全都会变成"当然应该这么做"。
🏔️ 一、内存层次:一张必须记住的表
| 层级 | 容量 / 到哪 | 速度 | 带宽 |
|---|---|---|---|
| 寄存器 | ~256 KB/SM | 极快 | (几乎免费) |
| Shared Memory / L1 | ~228 KB/SM | 很快 | ~10 TB/s |
| L2 Cache | ~50 MB | 快 | ~5 TB/s |
| HBM 显存 | 80 GB | ⚠️ 慢 | ~2–3 TB/s |
| —— 以下出卡了 —— | |||
| NVLink | 到另一张卡 | 很慢 | ~900 GB/s |
| PCIe | 到 CPU 内存 | 非常慢 | ~64 GB/s |
| 网络 | 到另一台机器 | 极慢 | ~50 GB/s |
⬇︎ 越往下:容量越大,速度越慢(差好几个数量级)。
🔑 记住这三个比值就够了: - Shared Memory 比 HBM 快约 5 倍(FlashAttention 的全部依据) - HBM 比 PCIe 快约 40 倍(不要随便
.cpu()) - NVLink 比网络快约 20 倍(并行策略选择的依据)
⚖️ 二、算术强度:判断瓶颈的唯一标准
$$\text{算术强度} = \frac{\text{计算量(FLOP)}}{\text{数据搬运量(字节)}}$$
💡 人话:每从显存搬一个字节,能顺带做多少次运算。
A100 的"平衡点":
312 TFLOPS ÷ 2.0 TB/s ≈ 156 FLOP/字节
⭐ 算术强度 > 156 → 【算力瓶颈】compute-bound(好,说明在真干活)
⭐ 算术强度 < 156 → 【带宽瓶颈】memory-bound(GPU 在等数据)
📊 常见操作的算术强度(照着这张表判断)
| 操作 | 算术强度 | 瓶颈 |
|---|---|---|
| 大矩阵乘(4096³) | ~1365 | ✅ 算力瓶颈 |
| 卷积(大通道数) | 几百 | ✅ 算力瓶颈 |
| 矩阵乘(batch=1,即向量×矩阵) | ~2 | 💀 严重带宽瓶颈 ⭐ |
| 逐元素操作(ReLU、加法) | ~0.25 | 💀 极度带宽瓶颈 |
| LayerNorm / Softmax | ~1 | 💀 带宽瓶颈 |
| Attention(长序列,朴素实现) | 低 | 💀 带宽瓶颈 |
📐 算一遍逐元素加法为什么这么惨(想看再点)
c = a + b,每个元素 FP16(2 字节):
- 计算:1 次加法 = 1 FLOP
- 搬运:读 a(2 字节)+ 读 b(2 字节)+ 写 c(2 字节)= 6 字节
- 算术强度 = 1 / 6 ≈ 0.17 FLOP/字节
离 156 的平衡点差了 900 倍 —— 也就是说,做逐元素加法时,A100 的算力有 99.9% 在闲置。 💀
⭐ 这就是算子融合存在的全部理由: 把 10 个逐元素操作合成 1 个 kernel,数据只搬一次而不是 10 次。
🔑 注意表里第三行:batch=1 的矩阵乘算术强度只有 ~2。 这就是第 15 章的核心 —— 自回归解码每次只生成 1 个 token,就是 batch=1 的矩阵乘, 推理天生是带宽瓶颈,和训练完全相反。
💾 三、训练显存到底花在哪
四大块(以 7B 模型、BF16 混合精度、Adam 为例):
| 项目 | 计算方式 | 7B 模型 |
|---|---|---|
| 模型参数 | P × 2 字节(BF16) | 14 GB |
| 梯度 | P × 2 字节 | 14 GB |
| 优化器状态 ⭐ | P × 12 字节 | 84 GB |
| 激活值 | 取决于 batch/序列 | 几 GB ~ 几十 GB |
📐 优化器状态为什么是 12 字节/参数(想看再点)
标准的 BF16 混合精度训练(第 6 章)需要维护:
| 内容 | 精度 | 字节/参数 |
|---|---|---|
| FP32 主权重副本 | FP32 | 4 |
| Adam 一阶动量 m | FP32 | 4 |
| Adam 二阶动量 v | FP32 | 4 |
| 合计 | 12 ⭐ |
加上 BF16 的参数(2)和梯度(2),总共 16 字节/参数。
💡 一个好记的经验公式: $$\text{训练显存(GB)} \approx 16 \times \text{参数量(B)} + \text{激活}$$
- 7B → 112 GB(一张 80G 卡装不下)
- 70B → 1120 GB(至少 14 张 80G 卡,还没算激活)
⚠️ 用 SGD 而不是 Adam 的话会小很多(没有 m、v),但大模型基本都用 Adam。
激活值:唯一和 batch 相关的一块
Transformer 每层每个 token 的激活 ≈ 十几到几十倍的 hidden_size
总激活 ≈ batch × 序列长度 × hidden × 层数 × 常数
⭐ 关键性质:它和 batch【成正比】
→ 这是唯一能靠"调小 batch"缓解的一块
→ 其他三块和 batch 完全无关 ⭐
🔑 这解释了第 1 章那个错误直觉: "OOM 就调小 batch"只对激活有用。 如果你的 OOM 是因为优化器状态装不下,调 batch 到 1 也没用 —— 那时候需要的是 ZeRO 或量化。
🧮 四、动手:估算你的模型要多少显存
def estimate_training_memory(params_B, batch, seq_len, hidden, layers,
optimizer="adam", grad_ckpt=False):
"""返回 GB。粗估,用于判断量级和选并行策略。"""
P = params_B * 1e9
weights = P * 2 / 1e9 # BF16 参数
grads = P * 2 / 1e9 # BF16 梯度
opt = P * (12 if optimizer == "adam" else 4) / 1e9
# 激活:不开检查点约 ~34 倍 hidden/层/token(Megatron 的经验系数)
per_layer = batch * seq_len * hidden * 34 * 2 / 1e9
acts = per_layer * (layers ** 0.5 if grad_ckpt else layers)
# ↑ 检查点后约降到 √层数 量级
total = weights + grads + opt + acts
print(f"参数 {weights:6.1f} GB | 梯度 {grads:6.1f} GB | "
f"优化器 {opt:6.1f} GB | 激活 {acts:6.1f} GB")
print(f"合计 ≈ {total:.1f} GB")
return total
# 7B,batch=4,seq=2048
estimate_training_memory(7, 4, 2048, 4096, 32)
# 参数 14.0 GB | 梯度 14.0 GB | 优化器 84.0 GB | 激活 35.0 GB
# 合计 ≈ 147.0 GB ⭐ 需要至少 2 张 80G 卡,且必须用 ZeRO
⚠️ 这是粗估,实际还要加:显存碎片(5–10%)、 通信缓冲区、CUDA context(几百 MB)、临时工作区。 留 15% 余量是常见做法。
🚧 五、"带宽墙"这个词的含义
历史趋势(大致量级):
算力 带宽 比值(算力/带宽)
2012 ~4 TFLOPS ~200 GB/s 20
2020 ~312 TFLOPS ~2000 GB/s 156 ⭐ 恶化了 8 倍
2024 ~989 TFLOPS ~3350 GB/s 295 ⭐ 继续恶化
→ 平衡点越来越高
→ 越来越多的操作变成【带宽瓶颈】
→ 这就是"带宽墙"(memory wall)
🔑 一个重要的推论: 随着硬件迭代,"少搬数据"的价值只会越来越高。 今天勉强算 compute-bound 的操作,下一代卡上就变成 memory-bound 了。
💡 这也解释了为什么 FP8、KV Cache 量化、算子融合这些技术 越来越受重视 —— 它们都是在直接对抗带宽墙。
🧰 六、由此推出的所有优化(后面章节的地图)
| 技术 | 它在对抗什么 | 章节 |
|---|---|---|
| 算子融合 | 逐元素操作反复往返 HBM | 07 |
| FlashAttention | attention 中间矩阵往返 HBM | 08 ⭐ |
| 混合精度 / FP8 | 数据量减半 → 带宽压力减半 | 06 |
| 量化 | 权重和 KV Cache 的搬运量 | 18 |
| 梯度检查点 | 用重算换激活显存 | 09 |
| ZeRO | 优化器状态那 84GB | 11 ⭐ |
| 连续批处理 | 提高 batch → 提高算术强度 | 17 |
⭐ 看这张表:七个技术,六个是在对抗带宽和显存,只有一个和算力有关。 这就是为什么说这一章是钥匙。
🔗 和站内其他章的关系
| 相关的地方 | 这里的位置 |
|---|---|
| 第 2 章 内存层次 | 本章给出了具体数字 |
| 第 1 章 少搬点 | 算术强度是它的量化形式 ⭐ |
| ML 基础第 8 章 要缓存激活 | 激活显存的来源 |
| ML 基础第 9 章 Adam 的 m、v | 优化器状态 12 字节的来源 |
✅ 检查点
- 内存层次里,哪三个比值最该记住?
- 算术强度的定义是什么?A100 的平衡点是多少?
- 逐元素加法的算术强度是多少?离平衡点差多远?这推出了什么技术?
- batch=1 的矩阵乘算术强度是多少?这对推理意味着什么?
- 训练显存的四大块是什么?7B 模型各占多少?
- 优化器状态为什么是 12 字节/参数?训练显存的经验公式是什么?
- 四大块里哪一块和 batch 有关?这修正了什么错误直觉?
- 什么是"带宽墙"?它随硬件迭代是变好还是变坏?推论是什么?
👀 答案
- Shared Memory 比 HBM 快约 5 倍(FlashAttention 的依据)、HBM 比 PCIe 快约 40 倍(别随便
.cpu())、NVLink 比网络快约 20 倍(并行策略的依据)。 - 计算量(FLOP) ÷ 数据搬运量(字节),即"每搬一个字节能顺带做多少次运算"。A100 平衡点 312 TFLOPS ÷ 2.0 TB/s ≈ 156 FLOP/字节。
- 约 0.17(1 FLOP ÷ 6 字节:读 a、读 b、写 c)。离平衡点差 900 倍,算力 99.9% 闲置。推出了算子融合——把 N 个逐元素操作合成一个 kernel,数据只搬一次。
- 约 2,严重带宽瓶颈。意味着自回归解码(每次生成 1 个 token)天生是带宽瓶颈,和训练完全相反。
- 参数(14GB)、梯度(14GB)、优化器状态(84GB,最大)、激活(几 GB~几十 GB)。
- FP32 主权重副本 4 + Adam 的 m 4 + v 4 = 12。加上 BF16 参数 2 和梯度 2 共 16。经验公式:训练显存(GB) ≈ 16 × 参数量(B) + 激活。
- 只有激活和 batch 成正比,其他三块完全无关。修正了"OOM 就调小 batch"——如果 OOM 来自优化器状态,batch 调到 1 也没用,那时需要 ZeRO 或量化。
- 算力/带宽的比值不断上升(2012 年 20 → 2020 年 156 → 2024 年 295),越来越多操作变成带宽瓶颈。在变坏。推论:"少搬数据"的价值只会越来越高,今天勉强 compute-bound 的操作下一代卡上就变 memory-bound。
🛑 可以停在这里
⚡ 走神救援
⭐⭐这章是钥匙。内存层次:寄存器 → Shared/L1(228KB,~10TB/s)→ L2(50MB)→ HBM(80GB,~2-3TB/s)→ NVLink(900GB/s)→ PCIe(64GB/s)→ 网络。⭐记三个比值:Shared 比 HBM 快 5 倍(FlashAttention 的依据)、HBM 比 PCIe 快 40 倍(别随便 .cpu())、NVLink 比网络快 20 倍(并行策略的依据)。⭐⭐算术强度 = FLOP ÷ 搬运字节数,A100 平衡点 156:大于它是算力瓶颈,小于是带宽瓶颈。关键数字:大矩阵乘 ~1365 ✅;逐元素加法 ~0.17(离平衡点差 900 倍,算力 99.9% 闲置 → 这就是算子融合的全部理由);⭐batch=1 的矩阵乘只有 ~2 → 自回归解码天生是带宽瓶颈,和训练完全相反(第15章核心)。训练显存四大块(7B):参数 14GB + 梯度 14GB + ⭐优化器状态 84GB(最大) + 激活;优化器 12 字节/参数 = FP32主权重4 + Adam的m 4 + v 4;⭐经验公式:训练显存 ≈ 16 × 参数量(B) + 激活(7B→112GB,一张80G装不下;70B→1120GB)。⭐只有激活和 batch 成正比 → 修正"OOM就调小batch":如果 OOM 来自优化器状态,batch 调到 1 也没用,要用 ZeRO 或量化。带宽墙:算力/带宽比值 2012 年 20 → 2020 年 156 → 2024 年 295,在持续恶化 → ⭐"少搬数据"的价值只会越来越高。后面七个技术里六个都在对抗带宽和显存。
下一节 👉 04-Roofline与MFU.md ⭐⭐