Skip to content

算子融合与自定义核

本页速览 算子融合把多个 elementwise/reduction 算子合成一个 kernel,避免中间张量在 HBM 往返——LLM 推理优化的核心手段。本文讲透经典融合、FlashAttention v1/v2/v3 的 SRAM tiling、FlashInfer 块稀疏、Triton/CUDA 编程模型、融合边界判定。

算子融合与自定义核

概念定义:少一次 HBM 读写就快一倍

GPU 性能优化的核心战场是减少 HBM 访问——算力(989 TFLOPS)远超带宽(3.35 TB/s),导致 memory-bound 算子跑不满算力。**算子融合(Operator Fusion)**是最直接的减 HBM 读写手段:把多个本来独立的 kernel 合成一个,让中间张量留在 SRAM/L2,不出 HBM。

理解算子融合的两个关键认知:

  1. 融合的收益来自少几遍 HBM 读写——3 个 elementwise 算子朴素实现要读写 3 遍中间张量,融合后只读写 1 遍输入输出,HBM 带宽省 3 倍
  2. 不是所有算子都能融合——融合受限于算子依赖(下游要上游全算完才能开始)和中间张量大小(中间结果太大塞不进 SRAM 就必须写回 HBM)。

FlashAttention 是融合思想的极致代表——把 attention 的 softmax 中间值留在 SRAM,不写回 HBM,把 attention 从 O(N²) HBM 读写降到 O(N)。

一、为什么融合能加速:算一笔账

考虑一个典型片段:Y = ReLU(W·X + b)。朴素实现:

text
1. GEMM:   C = W · X                → 读 X、W;写 C 到 HBM
2. Bias:   D = C + b                → 读 C、b;写 D 到 HBM
3. ReLU:   Y = max(0, D)            → 读 D;写 Y 到 HBM

每步要读一次 HBM、写一次 HBM。假设 X、C、D、Y 都是 N×M 大小(FP16),三次 kernel 共 6 次 HBM 读写(每次 2N·M bytes)。

融合后:在 GEMM kernel 末尾,C 一算出来立刻加 b、过 ReLU,输出 Y。HBM 读写只剩:读 X、W 一次,写 Y 一次。省了 4 次 N·M bytes 的 HBM 读写——3 倍带宽节省。

text
朴素:   X,W → HBM → C → HBM → D → HBM → Y    (6 次读写)
融合:   X,W → HBM → [GEMM+bias+ReLU 在 SRAM 内] → Y → HBM  (2 次读写)

融合 vs 量化:哪个收益更大

  • 量化:减字节(INT4 比 FP16 省 4 倍字节),收益直接是带宽 × 4;
  • 融合:减读写次数(3 算子融合省 3 遍),收益是 HBM 往返次数;
  • 两者叠加效果累乘:INT4 量化 + 算子融合 = 字节 ×4 省 + 读写 ×3 省 = 12 倍带宽压力降低。

二、经典融合模式

1. Elementwise + Elementwise(最简单)

text
ReLU(scale(x) + bias) → 一个 kernel

任意 elementwise 算子链都能融合(add、mul、scale、bias、ReLU、GELU、SiLU、Tanh 等)。

2. Reduction + Elementwise

text
LayerNorm = (x - mean) / sqrt(var + ε) · γ + β
         = 1. mean(x)                  # reduction
         = 2. var(x)                   # reduction
         = 3. x - mean                 # elementwise
         = 4. / sqrt(var + ε)          # elementwise
         = 5. · γ + β                  # elementwise

朴素实现 5 个 kernel,融合后 1 个 kernel 完成(前两步 reduction 留在 shared memory,后 3 步在 SRAM 内连续算)。

3. GEMM + Bias + Activation

text
Y = GELU(W · X + b)   # Transformer FFN 标配

cuBLASLt 的 epilogue fusion 直接支持:matmul 内部算完 → 在 epilogue 加 bias + GELU → 写回 HBM 一次。

4. LayerNorm + Linear

text
y = LayerNorm(x) → Linear(y)

把 LayerNorm 与随后的 Q/K/V projection 融合:LayerNorm 算完留在 SRAM,直接喂给 Linear——少了 LayerNorm 输出写回 HBM 这一步。

5. Attention 内部融合(FlashAttention 革命)

传统 attention:

text
Q·K → S → softmax(S) → P → P·V → O
              ↑              ↑
       中间 S、P 是 N×N 矩阵,要写回 HBM

S 是 N×N(序列长度的平方),N=2048 → 4M elements × 4 bytes = 16MB per attention head。每个 head 都要写回 HBM 再读回来——N² 量级的 HBM 读写,长 context 时是性能杀手。

FlashAttention 把整个 attention 算法重写为分块版(详见下文),让 S、P 全程留在 SRAM——HBM 读写降到 O(N) 量级。

三、FlashAttention:融合的里程碑

v1(2022)

核心思想:把 attention 的 Q·K·softmax·V 重排为分块计算(tiling),用 SRAM 装下分块,softmax 用 online softmax(数值稳定版),中间矩阵 S、P 全程不写回 HBM。

text
传统:    Q → HBM → S=QK^T → HBM → P=softmax(S) → HBM → O=P·V
                                                    (N² HBM 读写)
Flash:   Q → SRAM → 分块算 S_i, P_i → 立刻乘 V_i → 累加 O
                                                    (N HBM 读写)
  • 收益:长 context(N=4096+)attention 速度提升 2-3×,显存从 O(N²) 降到 O(N);
  • 局限:实现复杂(要手写 CUDA 或 Triton);早期版本 kernel 通用性差。

v2(2023)

v1 的优化:

  • 减少非 matmul 计算(softmax、rescaling)占比;
  • warp 并行改进:v1 跨 warp 同步开销大,v2 让每个 warp 处理一个 query 行;
  • 更高效的 SRAM tiling

v2 在 A100 上比 v1 快 ~2×,比标准 attention 快 ~5-10×(长 context)。

v3(2024,H100 专属)

针对 H100 的特化:

  • WGMMA async 指令:H100 Tensor Core 的 warp-group matmul 异步指令;
  • FP8 Tensor Core 加速:用 FP8 算 Q·K 和 P·V,2× 算力;
  • 流水线并行:异步 HBM 读、SRAM 算、Tensor Core 算 三级流水。

v3 在 H100 上比 v2 快 ~1.5-2×,长 context 时 attention 几乎跑满算力。

用哪个版本

  • A100:FlashAttention v2;
  • H100/H200:FlashAttention v3(需用 flash-attn 3.x);
  • PyTorch 2.x:内置 SDPA(scaled_dot_product_attention)会自动选择最佳 backend;
  • vLLM:内置各版本自动 fallback;
  • Triton 实现flash_attn_triton 适合学习。

四、FlashInfer:2025 的块稀疏 + 可组合格式

FlashInfer(2024-2025 出品)是新一代 attention kernel 库,专攻 LLM 服务场景:

核心特性

  1. 块稀疏 attention:支持 2D 块稀疏 mask(如 sliding window + global tokens 的 Longformer 模式);
  2. 可组合 KV cache 格式:同一个 KV cache 可被不同 attention kernel 调用(vLLM PagedAttention 的 paged KV 直接复用);
  3. append / fork 操作:支持长 context 增量追加(speculative decoding 验证用)、tree decoding(MTP 多 token 预测);
  4. 跨 backend 统一:CUDA、Triton、H100 WGMMA 多 backend 共用 KV。

与 FlashAttention 的关系

FlashAttention 是"密集 attention 优化到极致"——任意位置都能 attend。FlashInfer 更广——支持稀疏 + 分页 + 多种 LLM 服务场景的 attention 变体。vLLM 0.5+ 已经集成 FlashInfer,作为 attention kernel 的首选。

五、手写 kernel 的工具栈

不是所有融合都能靠 PyTorch eager 模式自动完成——很多融合要手写 kernel。主流工具:

工具抽象级难度适用
PyTorch eager + torch.compile自动融合,覆盖度有限
Triton(OpenAI)当前 LLM kernel 首选
CUDA C++极致性能,门槛最高
TVM / Halide中高编译器式,研究友好
CUTLASS(NVIDIA)中高大型 GEMM/epilogue 模板库

Triton:手写 kernel 的事实标准

Triton 是 OpenAI 开发的 Python-like kernel DSL,编译成高效 CUDA:

python
# Triton 简化版向量加 kernel
import triton
import triton.language as tl

@triton.jit
def add_kernel(x_ptr, y_ptr, out_ptr, N, BLOCK_SIZE: tl.constexpr):
    pid = tl.program_id(axis=0)
    offs = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
    mask = offs < N
    x = tl.load(x_ptr + offs, mask=mask)
    y = tl.load(y_ptr + offs, mask=mask)
    tl.store(out_ptr + offs, x + y, mask=mask)

特点:

  • block-level programming:写一个 block 的逻辑,Triton 自动 parallelism across blocks;
  • 不要管 thread/warp/grid——比 CUDA C++ 简单 5-10×;
  • 性能接近手写 CUDA:在 GEMM、attention、layernorm 上常常 80-95% of cuBLAS/cuDNN。

CUDA C++:极致性能的最后一公里

当 Triton 表达不下、或要 H100 的新指令(WGMMA、TMA、FP8 sparse)时,必须用 CUDA C++ + CUTLASS:

  • 优势:完全控制硬件资源(warp、shared memory、寄存器)、可用最新指令;
  • 劣势:代码量 3-5 倍于 Triton、可维护性差、bug 难调试;
  • 用途:Marlin kernel(W4A16 融合)、FlashAttention v3、TensorRT-LLM 的核心 kernel。

六、CUDA 编程模型回顾

写 kernel 必须理解 CUDA 的执行模型:

text
Grid (网格)  ─┬─ Block (0)  ─┬─ Warp (0) ─ Thread (0..31)
              │                ├─ Warp (1) ─ ...
              │                └─ ...
              ├─ Block (1)  ─ ...
              └─ ...
  • Thread:单执行单元,每个 thread 算一个元素;
  • Warp:32 个 thread 组成,单指令多数据(SIMT)——所有 thread 同步执行相同指令;
  • Block:多个 warp 组成,block 内可共享 shared memory + 同步;
  • Grid:多个 block 组成,跨 SM 调度。

关键资源

资源容量(H100 SM)作用
Thread寄存器255/thread最快,单 thread 私有
Warpwarp shuffle32 thread 间 ultrafast 通信
Blockshared memory228 KBblock 内共享,~10 TB/s
SML1 cache256 KB全 SM 共享
GPUL2 cache50 MB全 SM 共享
GPUHBM80 GB~3.35 TB/s,慢

关键优化技巧(与融合有关)

  • shared memory tiling:把数据分块装进 shared memory,反复用——融合的物理基础;
  • warp shuffle:32 thread 间直接传数据,不走 shared memory;
  • async memcpy:HBM 读 + SRAM 算并行流水线(cp.async 指令、H100 TMA);
  • register tiling:把热数据放寄存器,连 shared memory 都不读。 更多详见GPU 体系结构与优化

七、融合的边界

不是所有算子都能融合——判定的关键是算子依赖中间张量大小

1. 算子依赖

text
可融合:    Y = f(g(x))       # g 算完才能算 f,但都在单 pass
不可融合:  Y = f(x) + g(x)   # f 和 g 独立,但相加需要等两个都算完

后者如果 f 和 g 都是大算子(如两个 matmul),相加是 epilogue,可以融合;如果 f 和 g 各自是复杂 kernel,相加要单独 kernel。

2. 中间张量大小

text
可融合:    ReLU(W·X)        # 中间 W·X 在 SRAM 装得下
不可融合:  Softmax(W·X)     # Softmax 要 reduction,W·X 要写回 HBM 跨 block reduction

但 FlashAttention 证明了:重排算法可以让看似不能融合的 attention 也融合——通过分块 reduction。这是 fusion 的极致。

融合不是越多越好

  • 过度融合会让 kernel 巨大,shared memory 不够、寄存器溢出,反而性能下降;
  • 可读性下降——CUDA C++ 融合 5 个算子的 kernel 几千行,难维护;
  • 编译时间长——Triton kernel 编译要几秒到几分钟; 工程上"融合关键瓶颈算子链"比"无脑融合一切"更明智。

八、自动融合:编译器的角色

现代推理引擎尽量自动融合,不靠手写:

引擎/编译器融合能力代表
PyTorch Inductor(torch.compile)自动融合 elementwise + reductionPyTorch 2.x
XLA算子融合 + 布局传播JAX、TF
TVM RelaxLLM 推理融合 + 自动调度Apache TVM
TensorRT多算子融合 + INT8/FP8 epilogueNVIDIA
vLLM内置 Marlin/AWQ/FlashAttention kernel服务级

最常见的是 PyTorch 2.x 的 torch.compile

python
@torch.compile
def forward(x):
    return F.gelu(F.linear(F.layer_norm(x, ...), W, b))

Inductor 会把这个 chain 融合成少量 CUDA kernel,性能经常提升 1.5-3×。但 LLM 推理极致优化还是靠 vLLM/TensorRT-LLM 的专用 kernel,详见计算图优化

九、权衡与取舍

  • 手写 vs 自动融合:能用 torch.compile 解决就别手写;要极致性能(H100 WGMMA、Marlin)必须手写;
  • Triton vs CUDA C++:能用 Triton 就别写 CUDA C++;除非 Triton 表达不下或要新指令;
  • 融合 vs 调度:融合优化单 kernel,调度优化多请求——两者叠加才有最大收益;
  • 融合 vs 量化:量化减字节、融合减读写,先量化再融合叠加效果最好

延伸阅读

参考资料