📑 本页目录(点开跳转)
07 · 算子融合与编译
⏱ 24 分钟 | 🎁 一行代码常常能拿到 20-40%
🎯 一句话
十个逐元素操作,数据往返显存十趟;融合成一个 kernel,只往返一趟。 第 3 章算过逐元素操作的算术强度只有 0.17 —— 融合就是把这个数字乘以 N 的手段。
🔥 一、问题:kernel 之间必须过 HBM
y = gelu(x @ W + b) * scale
PyTorch eager 模式下会跑【四个】kernel:
① tmp1 = x @ W → 写 HBM
② tmp2 = tmp1 + b → 读 HBM,写 HBM
③ tmp3 = gelu(tmp2) → 读 HBM,写 HBM
④ y = tmp3 * scale → 读 HBM,写 HBM
⭐ 中间结果被写出去又读回来【三次】
而每次 kernel 只做一点点计算 → 纯带宽浪费
融合后:
① 一个 kernel:读 x 和 W → 在寄存器/Shared Memory 里
完成 matmul、加 bias、gelu、乘 scale → 写 y
⭐ HBM 往返从 4 次降到 1 次
收益估算:
一个 [8192, 8192] 的 BF16 张量 = 134 MB
融合前(4 个 kernel):约 7 次 HBM 读写 ≈ 940 MB
融合后(1 个 kernel):约 2 次 ≈ 270 MB
→ A100 2TB/s 下,节省约 0.34 毫秒
→ 一个 32 层的模型每步就是 ~10 毫秒 ⭐
💡 还有一个常被忽略的收益:kernel 启动开销。 每次 kernel 启动约 3–10 微秒。一个模型一步跑几千个 kernel, 光启动开销就能占几十毫秒 —— 小模型上这一项甚至比带宽更致命。
🪄 二、torch.compile:先试这个
import torch
model = torch.compile(model) # ⭐ 就这一行
它做的事:
① TorchDynamo:把 Python 字节码抓成计算图
② AOTAutograd:连反向图一起抓
③ Inductor:生成融合后的 Triton kernel
| 模式 | 说明 |
|---|---|
default |
平衡,编译快 |
reduce-overhead |
用 CUDA Graph 消除启动开销,小模型/推理效果显著 ⭐ |
max-autotune |
自动搜索最优配置,编译很慢但运行最快 |
典型收益:训练 10–30%,推理(小 batch)可达 2 倍。
⚠️ 五个会让它失效的坑
| 坑 | 症状 | 解法 |
|---|---|---|
| graph break ⭐ | 编译了但几乎没提速 | TORCH_LOGS="graph_breaks" 查在哪断的 |
| 动态形状反复重编译 | 前几十步极慢 | dynamic=True,或把序列长度分桶补齐 ⭐ |
代码里有 .item() / print |
强制同步 → 必然 graph break | 训练循环里删掉 |
| 数据依赖的控制流 | if x.sum() > 0: 无法编译 |
改写成张量运算或接受断点 |
| 首次编译很慢 | 以为卡死了 | 正常,几十秒到几分钟;别在 benchmark 里算进去 |
# ⭐ 排查 graph break 的标准做法
TORCH_LOGS="graph_breaks,recompiles" python train.py
🔑
graph break是最常见的"我编译了但没变快"的原因: 一个 break 会把图切成两半,跨边界的融合机会全部丢失。 一个训练循环里有几十个 break 是很常见的。
✍️ 三、Triton:需要手写时的选择
import triton, triton.language as tl
@triton.jit
def fused_add_relu(x_ptr, y_ptr, out_ptr, n, BLOCK: tl.constexpr):
pid = tl.program_id(0)
offs = pid * BLOCK + tl.arange(0, BLOCK)
mask = offs < n
x = tl.load(x_ptr + offs, mask=mask) # ⭐ 一次读
y = tl.load(y_ptr + offs, mask=mask)
out = tl.maximum(x + y, 0.0) # 加法和 relu 在寄存器里完成
tl.store(out_ptr + offs, out, mask=mask) # ⭐ 一次写
Triton vs CUDA:
| Triton | CUDA | |
|---|---|---|
| 语言 | Python | C++ |
| 你要管的 | block 级别 | thread 级别 |
| 自动处理 | 内存合并、shared memory、warp 调度 ⭐ | 全部手动 |
| 性能 | 通常能到手写 CUDA 的 80–95% | 100% |
| 上手 | 几小时 | 几周 |
⭐ 什么时候值得手写:
① torch.compile 试过了,profiler 显示某个 kernel 仍是瓶颈 ② 这个模式框架里没有(比如自定义的稀疏注意力) ③ 你能算出理论上限,且当前实现离它很远(第 4 章) ❌ 不满足这三条 → 别写,成本不划算
📦 四、现成的融合实现(优先用这些)
⭐ 顺序:现成库 > torch.compile > 手写 Triton > 手写 CUDA
| 来源 | 提供什么 |
|---|---|
| FlashAttention | 融合的注意力(第 8 章)⭐ |
F.scaled_dot_product_attention |
PyTorch 内置,自动选 FlashAttention 后端 ⭐ |
| Apex / TransformerEngine | 融合 LayerNorm、融合 Adam、FP8 支持 |
| xFormers | 各种注意力变体 |
torch.nn.functional 的融合版 |
如 fused 参数的优化器 |
# ⭐ 融合优化器:把 Adam 的几十个小 kernel 合成一个
import torch
opt = torch.optim.AdamW(model.parameters(), lr=3e-4, fused=True)
# ↑ 常有 5-10% 提升
🧊 五、CUDA Graph:消灭启动开销
问题:每个 kernel 启动都有 CPU 开销(3-10 微秒)
小模型 / 推理时,CPU 可能【喂不饱 GPU】
CUDA Graph:把一整串 kernel 调用【录制】下来,
之后一次性提交给 GPU
⭐ 收益:小 batch 推理可达 20-50%
⚠️ 限制:形状必须固定,不能有动态控制流
import torch
model = torch.compile(model, mode="reduce-overhead") # ⭐ 自动用 CUDA Graph
💡 怎么判断你是不是"CPU 喂不饱 GPU": profiler 的时间轴上,GPU 行有大量空隙,CPU 行满满当当 —— 就是它。 这种情况下算力优化毫无意义,要解决的是启动开销。
🔗 和站内其他章的关系
| 相关的地方 | 这里的位置 |
|---|---|
| 第 3 章 逐元素算术强度 0.17 | 融合要解决的就是它 ⭐ |
| 第 2 章 Shared Memory | 融合后中间结果留在这里 |
| 第 4 章 一堆小 kernel | profiler 里融合机会的信号 |
| 第 8 章 FlashAttention | 融合思想的最成功案例 ⭐ |
| 《机器学习与深度学习基础》15 PyTorch 实战手册 提速三件套 | 那里教你加上 torch.compile,本章讲它为什么有效、以及五个让它白加的坑 ⭐ |
| 《机器学习与深度学习基础》11 训练调试手册 训练循环里打印中间值 | 调试时随手加的 .item() / print,在这里会直接把图切断 ⚠️ |
✅ 检查点
- 为什么 kernel 之间必须过 HBM?这造成什么浪费?
- 除了带宽,融合还省了什么开销?什么时候这一项更致命?
torch.compile的三个阶段是什么?典型收益多少?reduce-overhead模式做了什么?适合什么场景?- "编译了但没变快"最常见的原因是什么?怎么查?
- Triton 相比 CUDA,你不用管什么?性能能到多少?
- 什么时候才值得手写 kernel?(三个条件)
- 怎么从 profiler 判断"CPU 喂不饱 GPU"?
👀 答案
- 因为每个 kernel 的输出必须写回 HBM 供下一个 kernel 读取。造成中间结果反复往返显存——一串 4 个 kernel 要 7 次 HBM 读写,而每个 kernel 只做一点点计算,纯带宽浪费。
- kernel 启动开销(每次 3-10 微秒)。小模型上这一项比带宽更致命——一步跑几千个 kernel,光启动就几十毫秒。
- ①TorchDynamo 抓计算图 ②AOTAutograd 连反向一起抓 ③Inductor 生成融合的 Triton kernel。收益:训练 10-30%,小 batch 推理可达 2 倍。
- 用 CUDA Graph 把一串 kernel 调用录制下来一次性提交,消除启动开销。适合小模型/小 batch 推理。
- graph break——一个 break 把图切成两半,跨边界的融合机会全丢。用
TORCH_LOGS="graph_breaks,recompiles"查。常见诱因:.item()、print、数据依赖的控制流。 - 不用管内存合并、shared memory、warp 调度(只管 block 级别)。性能能到手写 CUDA 的 80-95%。
- ①torch.compile 试过了,profiler 显示某 kernel 仍是瓶颈 ②这个模式框架里没有 ③你能算出理论上限且当前实现离它很远。三条不全满足就别写。
- GPU 行有大量空隙,CPU 行满满当当。这时算力优化毫无意义,要解决的是启动开销(CUDA Graph)。
🛑 可以停在这里
⚡ 走神救援
⭐十个逐元素操作往返显存十趟,融合成一个 kernel 只往返一趟——这是第 3 章"逐元素算术强度只有 0.17"的解药。
y=gelu(x@W+b)*scale在 eager 下跑 4 个 kernel、中间结果写出去又读回来 3 次;融合后 HBM 往返从 4 次降到 1 次(8192² 张量上一层省 0.34ms,32 层就是 ~10ms)。💡还有 kernel 启动开销(每次 3-10μs,一步几千个 kernel = 几十毫秒),小模型上这项比带宽更致命。⭐先试torch.compile(model)一行(Dynamo 抓图 → AOTAutograd 抓反向 → Inductor 生成 Triton kernel),训练 10-30%、小 batch 推理可达 2 倍;mode="reduce-overhead"用 CUDA Graph 消除启动开销。⚠️⭐"编译了没变快"最常见原因是 graph break(一个 break 切断跨边界的融合机会,一个训练循环几十个 break 很常见)→TORCH_LOGS="graph_breaks,recompiles";诱因:.item()、F.scaled_dot_product_attention(自动选 FA 后端)/ Apex 融合 LayerNorm /AdamW(fused=True)(5-10%)。⭐判断 CPU 喂不饱 GPU:profiler 时间轴上 GPU 行有空隙、CPU 行满——这时算力优化毫无意义。
下一节 👉 08-FlashAttention.md ⭐⭐