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

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在这里会直接把图切断 ⚠️

✅ 检查点

  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 只往返一趟——这是第 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()print、数据依赖控制流、动态形状反复重编译(分桶补齐)。Triton:Python 语法、只管 block 级别(内存合并/shared memory/warp 调度自动处理)、性能到手写 CUDA 的 80-95%、几小时上手。⭐什么时候才手写:compile 试过了+profiler 显示仍是瓶颈+能算出理论上限且离得很远——三条不全就别写。⭐优先级:现成库 > torch.compile > Triton > CUDA;现成的有 FlashAttention / F.scaled_dot_product_attention(自动选 FA 后端)/ Apex 融合 LayerNorm / AdamW(fused=True)(5-10%)。⭐判断 CPU 喂不饱 GPU:profiler 时间轴上 GPU 行有空隙、CPU 行满——这时算力优化毫无意义。

下一节 👉 08-FlashAttention.md ⭐⭐

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