🏠 总目录📚 本教程 算子融合与编译 ← →
📑 本页目录(点开跳转)

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:先试这个

# 🧩 骨架:`model` 来自你自己的代码,这一段只看写法
import torch
model = torch.compile(model)          # ⭐ 就这一行

它做的事:

操作步骤

  1. TorchDynamo:把 Python 字节码抓成计算图
  2. AOTAutograd:连反向图一起抓
  3. 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 参数的优化器
# 🧩 骨架:`model` 来自你自己的代码,这一段只看写法
# ⭐ 融合优化器:把 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%

⚠️ 限制:形状必须固定,不能有动态控制流

# 🧩 骨架:`model` 来自你自己的代码,这一段只看写法
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,在这里会直接把图切断 ⚠️

✅ 检查点

  1. 为什么 kernel 之间必须过 HBM?这造成什么浪费?
  2. 除了带宽,融合还省了什么开销?什么时候这一项更致命?
  3. torch.compile 的三个阶段是什么?典型收益多少?
  4. reduce-overhead 模式做了什么?适合什么场景?
  5. "编译了但没变快"最常见的原因是什么?怎么查?
  6. Triton 相比 CUDA,你不用管什么?性能能到多少?
  7. 什么时候才值得手写 kernel?(三个条件)
  8. 怎么从 profiler 判断"CPU 喂不饱 GPU"?
👀 答案
  1. 因为每个 kernel 的输出必须写回 HBM 供下一个 kernel 读取。造成中间结果反复往返显存——一串 4 个 kernel 要 7 次 HBM 读写,而每个 kernel 只做一点点计算,纯带宽浪费。
  2. kernel 启动开销(每次 3-10 微秒)。小模型上这一项比带宽更致命——一步跑几千个 kernel,光启动就几十毫秒。
  3. ①TorchDynamo 抓计算图 ②AOTAutograd 连反向一起抓 ③Inductor 生成融合的 Triton kernel。收益:训练 10-30%,小 batch 推理可达 2 倍。
  4. 用 CUDA Graph 把一串 kernel 调用录制下来一次性提交,消除启动开销。适合小模型/小 batch 推理。
  5. graph break——一个 break 把图切成两半,跨边界的融合机会全丢。用 TORCH_LOGS="graph_breaks,recompiles" 查。常见诱因:.item()、print、数据依赖的控制流。
  6. 不用管内存合并、shared memory、warp 调度(只管 block 级别)。性能能到手写 CUDA 的 80-95%。
  7. ①torch.compile 试过了,profiler 显示某 kernel 仍是瓶颈 ②这个模式框架里没有 ③你能算出理论上限且当前实现离它很远。三条不全满足就别写。
  8. GPU 行有大量空隙,CPU 行满满当当。这时算力优化毫无意义,要解决的是启动开销(CUDA Graph)。

🛑 可以停在这里

⚡ 走神救援

⭐ 十个逐元素操作要往返显存十趟,融合成一个 kernel 只往返一趟——这是「逐元素运算算术强度极低」那条的解药。一串小操作在 eager 下每一步都把中间结果写出去又读回来,融合之后往返次数直接降到一次;⚠️ 层数一多,省下的就是几毫秒到几十毫秒的量级。

💡 还有第二笔账:kernel 启动开销。单次只有几微秒,⭐ 但一步几千个 kernel 就是几十毫秒——小模型上这一项比带宽更致命。

⭐ 先试一行 torch.compile(抓图、抓反向、生成 kernel),训练能拿到一两成、小 batch 推理能到成倍;专门的模式还会用 CUDA Graph 把启动开销也消掉。

⚠️⭐ 「编译了没变快」最常见的原因是 graph break:一个 break 就切断跨越它的融合机会,而一个训练循环有几十个 break 很常见。⭐ 别猜,开日志看。常见诱因:把张量取成 Python 标量、打印、数据依赖的控制流、⚠️ 动态形状反复触发重编译(分桶补齐)。

⭐ 什么时候才轮到自己写 kernel:三条必须同时成立——torch.compile 试过了、profiler 显示它仍是瓶颈、而且你能算出理论上限并确认离得很远。⭐ 优先级永远是:现成库 > 编译器 > Triton > 手写 CUDA。

⭐ 最后一条判据最省事:profiler 时间轴上 GPU 那行有空隙、CPU 那行是满的,就说明 CPU 喂不饱 GPU——这时候做任何算力优化都毫无意义。

下一节 👉 08-FlashAttention.md ⭐

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