🏠 总目录📚 本教程 显存与带宽墙 ← →
📑 本页目录(点开跳转)

03 · 显存与带宽墙

⏱ 30 分钟 | ⭐ 理解这一章,后面十几章都会变简单


🎯 一句话

你的 GPU 不是在算,大部分时间是在等数据。 这一章讲清楚数据存在哪、搬一次要多久, 之后 FlashAttention、量化、KV Cache、算子融合 —— 全都会变成"当然应该这么做"。

SRAM / 片上~19 TB/s~20 MBHBM 显存~3.3 TB/s80 GB主机内存~64 GB/s1 TB+网络 / 跨机~25 GB/s—容量越快的存储越小,越大的存储越慢 —— 差了近 1000 倍所以瓶颈常常不是「算得慢」,是「数据搬不过来」FlashAttention 之类的优化,本质都是在减少这堵墙前面的排队
越快的存储越小,越大的存储越慢,中间差了近 1000 倍。⭐ 所以真正的瓶颈常常不是「算得慢」,而是「数据搬不过来」 —— 后面十几章的优化,本质上都是在减少这堵墙前面的排队。

🏔️ 一、内存层次:一张必须记住的表

层级 容量 / 到哪 速度 带宽
寄存器 ~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 字节):

离 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{激活}$$

⚠️ 用 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 字节的来源

✅ 检查点

  1. 内存层次里,哪三个比值最该记住?
  2. 算术强度的定义是什么?A100 的平衡点是多少?
  3. 逐元素加法的算术强度是多少?离平衡点差多远?这推出了什么技术?
  4. batch=1 的矩阵乘算术强度是多少?这对推理意味着什么?
  5. 训练显存的四大块是什么?7B 模型各占多少?
  6. 优化器状态为什么是 12 字节/参数?训练显存的经验公式是什么?
  7. 四大块里哪一块和 batch 有关?这修正了什么错误直觉?
  8. 什么是"带宽墙"?它随硬件迭代是变好还是变坏?推论是什么?
👀 答案
  1. Shared Memory 比 HBM 快约 5 倍(FlashAttention 的依据)、HBM 比 PCIe 快约 40 倍(别随便 .cpu())、NVLink 比网络快约 20 倍(并行策略的依据)。
  2. 计算量(FLOP) ÷ 数据搬运量(字节),即"每搬一个字节能顺带做多少次运算"。A100 平衡点 312 TFLOPS ÷ 2.0 TB/s ≈ 156 FLOP/字节。
  3. 约 0.17(1 FLOP ÷ 6 字节:读 a、读 b、写 c)。离平衡点差 900 倍,算力 99.9% 闲置。推出了算子融合——把 N 个逐元素操作合成一个 kernel,数据只搬一次。
  4. 约 2,严重带宽瓶颈。意味着自回归解码(每次生成 1 个 token)天生是带宽瓶颈,和训练完全相反。
  5. 参数(14GB)、梯度(14GB)、优化器状态(84GB,最大)、激活(几 GB~几十 GB)。
  6. FP32 主权重副本 4 + Adam 的 m 4 + v 4 = 12。加上 BF16 参数 2 和梯度 2 共 16。经验公式:训练显存(GB) ≈ 16 × 参数量(B) + 激活。
  7. 只有激活和 batch 成正比,其他三块完全无关。修正了"OOM 就调小 batch"——如果 OOM 来自优化器状态,batch 调到 1 也没用,那时需要 ZeRO 或量化。
  8. 算力/带宽的比值不断上升(2012 年 20 → 2020 年 156 → 2024 年 295),越来越多操作变成带宽瓶颈。在变坏。推论:"少搬数据"的价值只会越来越高,今天勉强 compute-bound 的操作下一代卡上就变 memory-bound。

🛑 可以停在这里

⚡ 走神救援

先记住这几件事

下一节 👉 04-Roofline与MFU.md ⭐

打卡记录保存在你的浏览器里,首页能看到总进度