外观
算子融合与自定义核
概念定义:少一次 HBM 读写就快一倍
GPU 性能优化的核心战场是减少 HBM 访问——算力(989 TFLOPS)远超带宽(3.35 TB/s),导致 memory-bound 算子跑不满算力。**算子融合(Operator Fusion)**是最直接的减 HBM 读写手段:把多个本来独立的 kernel 合成一个,让中间张量留在 SRAM/L2,不出 HBM。
理解算子融合的两个关键认知:
- 融合的收益来自少几遍 HBM 读写——3 个 elementwise 算子朴素实现要读写 3 遍中间张量,融合后只读写 1 遍输入输出,HBM 带宽省 3 倍;
- 不是所有算子都能融合——融合受限于算子依赖(下游要上游全算完才能开始)和中间张量大小(中间结果太大塞不进 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 矩阵,要写回 HBMS 是 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 服务场景:
核心特性
- 块稀疏 attention:支持 2D 块稀疏 mask(如 sliding window + global tokens 的 Longformer 模式);
- 可组合 KV cache 格式:同一个 KV cache 可被不同 attention kernel 调用(vLLM PagedAttention 的 paged KV 直接复用);
- append / fork 操作:支持长 context 增量追加(speculative decoding 验证用)、tree decoding(MTP 多 token 预测);
- 跨 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 私有 |
| Warp | warp shuffle | — | 32 thread 间 ultrafast 通信 |
| Block | shared memory | 228 KB | block 内共享,~10 TB/s |
| SM | L1 cache | 256 KB | 全 SM 共享 |
| GPU | L2 cache | 50 MB | 全 SM 共享 |
| GPU | HBM | 80 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 + reduction | PyTorch 2.x |
| XLA | 算子融合 + 布局传播 | JAX、TF |
| TVM Relax | LLM 推理融合 + 自动调度 | Apache TVM |
| TensorRT | 多算子融合 + INT8/FP8 epilogue | NVIDIA |
| 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 量化:量化减字节、融合减读写,先量化再融合叠加效果最好。
延伸阅读
- 计算图优化——融合在图层面的实现
- 显存层次与带宽墙——融合的物理基础(SRAM/HBM)
- GPU 体系结构与优化——CUDA 编程模型与硬件细节
- 权重量化与混合精度——Marlin kernel 的融合实现
- 批处理与调度——堆并发榨干带宽
- vLLM 案例研究——FlashAttention/Marlin 的工业实现
- TensorRT 案例研究——NVIDIA 的算子融合栈
- benchmarking 实践——单 kernel 性能测量
参考资料
- Dao et al. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness(NeurIPS 2022) —— FlashAttention v1 原始论文
- Dao. FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning(2023) —— FlashAttention v2
- Shah et al. FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision(2024) —— FlashAttention v3
- Ye et al. FlashInfer: Efficient and Customizable Attention Engine for LLM Inference Serving(2024) —— FlashInfer 论文
- Triton: An Intermediate Language and Compiler for Tiled Neural Network Computation —— Triton 开源仓库与文档
- NVIDIA CUTLASS Documentation —— CUDA C++ 模板库
- Tillet et al. Triton: an intermediate language and compiler for tiled neural network computation(MAPL 2019) —— Triton 论文
- PyTorch 2.0 torch.compile Documentation —— Inductor 自动融合