Appearance
Kernel Fusion and Custom Kernels
Concept Definition: Every HBM Round Trip You Skip Is a Win
The central battleground of GPU performance optimization is reducing HBM accesses — compute (989 TFLOPS) far exceeds bandwidth (3.35 TB/s), leaving memory-bound operators unable to saturate compute. Kernel fusion is the most direct way to cut HBM traffic: merge multiple originally independent kernels into one, keeping intermediate tensors in SRAM/L2 instead of spilling them to HBM.
Two key insights for understanding kernel fusion:
- The gain comes from eliminating HBM round trips — three naive elementwise kernels read and write the intermediate tensor three times; fused, they touch only the input and output once, saving 3× of HBM bandwidth;
- Not every operator can be fused — fusion is constrained by operator dependencies (a downstream operator may need the upstream one fully computed first) and intermediate tensor size (an intermediate too large for SRAM must be written back to HBM).
FlashAttention is the ultimate expression of the fusion idea — it keeps the softmax intermediates of attention in SRAM, never writing them back to HBM, reducing attention's HBM traffic from O(N²) to O(N).
1. Why Fusion Speeds Things Up: Do the Math
Consider a typical fragment: Y = ReLU(W·X + b). Naive implementation:
text
1. GEMM: C = W · X → read X, W; write C to HBM
2. Bias: D = C + b → read C, b; write D to HBM
3. ReLU: Y = max(0, D) → read D; write Y to HBMEach step reads HBM once and writes HBM once. Assume X, C, D, Y are all N×M (FP16) — three kernels make 6 HBM round trips in total (2N·M bytes each).
After fusion: at the end of the GEMM kernel, C is computed, immediately biased and passed through ReLU, producing Y. HBM traffic is reduced to: read X, W once; write Y once. Four N·M-byte HBM round trips saved — a 3× bandwidth saving.
text
Naive: X,W → HBM → C → HBM → D → HBM → Y (6 HBM round trips)
Fused: X,W → HBM → [GEMM+bias+ReLU inside SRAM] → Y → HBM (2 HBM round trips)Fusion vs. Quantization: Which Pays More
- Quantization: cuts bytes (INT4 uses 4× fewer bytes than FP16) — the gain is directly bandwidth ×4;
- Fusion: cuts the number of round trips (fusing 3 operators saves 3 passes) — the gain is in HBM trips;
- The two stack multiplicatively: INT4 quantization + kernel fusion = 4× fewer bytes + 3× fewer trips = 12× lower bandwidth pressure.
2. Classic Fusion Patterns
1. Elementwise + Elementwise (The Simplest)
text
ReLU(scale(x) + bias) → one kernelAny chain of elementwise operators fuses (add, mul, scale, bias, ReLU, GELU, SiLU, Tanh, etc.).
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. · γ + β # elementwiseFive kernels naive, one kernel fused (the first two reduction steps stay in shared memory, and the last three elementwise steps run continuously in SRAM).
3. GEMM + Bias + Activation
text
Y = GELU(W · X + b) # the standard Transformer FFNcuBLASLt's epilogue fusion supports this natively: matmul completes internally → bias + GELU applied in the epilogue → written back to HBM once.
4. LayerNorm + Linear
text
y = LayerNorm(x) → Linear(y)Fuse LayerNorm with the following Q/K/V projection: LayerNorm's output stays in SRAM and feeds the Linear directly — eliminating the write-back of the LayerNorm output to HBM.
5. Intra-Attention Fusion (the FlashAttention Revolution)
Traditional attention:
text
Q·K → S → softmax(S) → P → P·V → O
↑ ↑
intermediate S, P are N×N matrices that must be written back to HBMS is N×N (the square of sequence length): N=2048 → 4M elements × 4 bytes = 16MB per attention head. Every head writes back to HBM and reads it back again — N²-scale HBM traffic, a performance killer at long context.
FlashAttention rewrites the whole attention algorithm as a blocked version (details below), keeping S and P in SRAM throughout — HBM traffic drops to O(N).
3. FlashAttention: The Milestone of Fusion
v1 (2022)
Core idea: restructure attention's Q·K·softmax·V into blocked computation (tiling), fit the blocks in SRAM, use online softmax (the numerically stable version), and never write the intermediate matrices S, P back to HBM.
text
Traditional: Q → HBM → S=QK^T → HBM → P=softmax(S) → HBM → O=P·V
(N² HBM round trips)
Flash: Q → SRAM → compute S_i, P_i block by block → multiply V_i immediately → accumulate O
(N HBM round trips)- Gain: 2-3× attention speedup at long context (N=4096+), memory from O(N²) down to O(N);
- Limitation: complex to implement (hand-written CUDA or Triton required); early kernels lacked generality.
v2 (2023)
Optimizations over v1:
- Reduce the share of non-matmul work (softmax, rescaling);
- Better warp parallelism: v1 had heavy cross-warp synchronization; v2 gives each warp its own query rows;
- More efficient SRAM tiling.
v2 is ~2× faster than v1 on A100 and ~5-10× faster than standard attention (long context).
v3 (2024, H100-Exclusive)
Specialized for H100:
- WGMMA async instructions: H100 Tensor Core's warp-group matmul asynchronous instructions;
- FP8 Tensor Core acceleration: compute Q·K and P·V in FP8 for 2× compute;
- Three-stage pipelining: asynchronous HBM reads, SRAM compute, and Tensor Core compute overlapped.
v3 is ~1.5-2× faster than v2 on H100 — attention nearly saturates compute at long context.
Which Version to Use
- A100: FlashAttention v2;
- H100/H200: FlashAttention v3 (requires flash-attn 3.x);
- PyTorch 2.x: the built-in SDPA (scaled_dot_product_attention) automatically picks the best backend;
- vLLM: built-in automatic fallback across versions;
- Triton implementation:
flash_attn_tritonis good for learning.
4. FlashInfer: Block Sparsity + Composable Formats in 2025
FlashInfer (2024-2025) is a new-generation attention kernel library targeting LLM serving:
Core Features
- Block-sparse attention: supports 2D block-sparse masks (e.g. the Longformer pattern of sliding window + global tokens);
- Composable KV cache formats: the same KV cache can be consumed by different attention kernels (vLLM PagedAttention's paged KV is reused directly);
- Append / fork operations: supports incremental appending for long contexts (for speculative decoding verification), tree decoding (MTP multi-token prediction);
- Cross-backend unification: CUDA, Triton, and H100 WGMMA backends share the same KV.
Relationship to FlashAttention
FlashAttention is "dense attention optimized to the extreme" — every position can attend. FlashInfer goes wider — sparse + paged + attention variants for many LLM serving scenarios. vLLM 0.5+ has integrated FlashInfer as the preferred attention kernel.
5. The Toolchain for Hand-Written Kernels
Not every fusion can be automated by PyTorch eager mode — many fusions must be hand-written. Mainstream tools:
| Tool | Abstraction Level | Difficulty | Best For |
|---|---|---|---|
| PyTorch eager + torch.compile | High | Low | Automatic fusion, limited coverage |
| Triton (OpenAI) | Medium | Medium | The current default for LLM kernels |
| CUDA C++ | Low | High | Ultimate performance, highest bar |
| TVM / Halide | Medium | Medium-high | Compiler-style, research-friendly |
| CUTLASS (NVIDIA) | Medium | Medium-high | Large GEMM/epilogue template library |
Triton: The De Facto Standard for Hand-Written Kernels
Triton is a Python-like kernel DSL developed by OpenAI that compiles into efficient CUDA:
python
# Simplified Triton vector-add 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)Characteristics:
- Block-level programming: write the logic for one block, and Triton handles parallelism across blocks automatically;
- No thread/warp/grid bookkeeping — 5-10× simpler than CUDA C++;
- Performance close to hand-written CUDA: often 80-95% of cuBLAS/cuDNN on GEMM, attention, and LayerNorm.
CUDA C++: The Last Mile for Ultimate Performance
When Triton can't express what you need, or you need H100's new instructions (WGMMA, TMA, FP8 sparse), you must use CUDA C++ + CUTLASS:
- Strengths: full control over hardware resources (warps, shared memory, registers); access to the latest instructions;
- Weaknesses: 3-5× the code volume of Triton, poor maintainability, hard-to-debug bugs;
- Used by: the Marlin kernel (W4A16 fusion), FlashAttention v3, and the core kernels of TensorRT-LLM.
6. The CUDA Programming Model, Revisited
Writing kernels requires understanding CUDA's execution model:
text
Grid ─┬─ Block (0) ─┬─ Warp (0) ─ Thread (0..31)
│ ├─ Warp (1) ─ ...
│ └─ ...
├─ Block (1) ─ ...
└─ ...- Thread: a single execution unit; each thread computes one element;
- Warp: 32 threads — single instruction, multiple data (SIMT) — all threads execute the same instruction in lockstep;
- Block: multiple warps; threads within a block share shared memory + synchronization;
- Grid: multiple blocks, scheduled across SMs.
Key Resources
| Level | Resource | Capacity (H100 SM) | Role |
|---|---|---|---|
| Thread | Registers | 255/thread | Fastest, private per thread |
| Warp | Warp shuffle | — | Ultra-fast communication among 32 threads |
| Block | Shared memory | 228 KB | Shared within a block, ~10 TB/s |
| SM | L1 cache | 256 KB | Shared across the SM |
| GPU | L2 cache | 50 MB | Shared across all SMs |
| GPU | HBM | 80 GB | ~3.35 TB/s, slow |
Key Optimization Techniques (Related to Fusion)
- Shared-memory tiling: load data into shared memory in blocks and reuse it — the physical basis of fusion;
- Warp shuffle: pass data directly among 32 threads without going through shared memory;
- Async memcpy: overlap HBM reads with SRAM compute (the cp.async instruction, H100 TMA);
- Register tiling: keep hot data in registers — don't even read shared memory. More in GPU Architecture and Optimization.
7. The Boundaries of Fusion
Not every operator can be fused — the key criteria are operator dependencies and intermediate tensor size:
1. Operator Dependencies
text
Fusable: Y = f(g(x)) # f can't start until g finishes, but both fit in a single pass
Not fusible: Y = f(x) + g(x) # f and g are independent, but the addition must wait for bothIn the latter case, if f and g are both large operators (e.g. two matmuls), the addition is an epilogue and can fuse; if f and g are each complex kernels, the addition needs its own kernel.
2. Intermediate Tensor Size
text
Fusable: ReLU(W·X) # the intermediate W·X fits in SRAM
Not fusible: Softmax(W·X) # Softmax needs a reduction, so W·X must be written back to HBM for cross-block reductionBut FlashAttention proved the point: rearranging the algorithm can make even seemingly unfusable attention fuse — through blocked reduction. That is fusion taken to the extreme.
More Fusion Is Not Always Better
- Over-fusion makes kernels huge — shared memory runs out, registers spill, and performance actually drops;
- Readability suffers — a CUDA C++ kernel fusing 5 operators runs to thousands of lines and is hard to maintain;
- Compilation takes long — Triton kernels take seconds to minutes to compile; In production engineering, "fusing the critical bottleneck operator chain" is wiser than "fusing everything blindly."
8. Automatic Fusion: The Role of Compilers
Modern inference engines automate fusion as much as possible instead of relying on hand-written kernels:
| Engine/Compiler | Fusion Capability | Representative |
|---|---|---|
| PyTorch Inductor (torch.compile) | Automatic elementwise + reduction fusion | PyTorch 2.x |
| XLA | Operator fusion + layout propagation | JAX, TF |
| TVM Relax | LLM inference fusion + auto-scheduling | Apache TVM |
| TensorRT | Multi-operator fusion + INT8/FP8 epilogue | NVIDIA |
| vLLM | Built-in Marlin/AWQ/FlashAttention kernels | Serving-level |
The most common case is PyTorch 2.x's torch.compile:
python
@torch.compile
def forward(x):
return F.gelu(F.linear(F.layer_norm(x, ...), W, b))Inductor fuses this chain into a handful of CUDA kernels, often gaining 1.5-3×. But ultimate LLM inference optimization still relies on the dedicated kernels of vLLM/TensorRT-LLM — see Computation Graph Optimization.
9. Trade-offs
- Hand-written vs. automatic fusion: if torch.compile solves it, don't hand-write; ultimate performance (H100 WGMMA, Marlin) requires hand-written kernels;
- Triton vs. CUDA C++: if Triton expresses it, don't write CUDA C++ — unless Triton falls short or you need new instructions;
- Fusion vs. scheduling: fusion optimizes a single kernel; Batching and Request Scheduling optimizes across requests — stacking both yields the most;
- Fusion vs. quantization: quantization cuts bytes, fusion cuts trips — quantize first, then fuse for the best stacked effect.
Further Reading
- Computation Graph Optimization — fusion implemented at the graph level
- The GPU Memory Hierarchy and the Bandwidth Wall — the physical basis of fusion (SRAM/HBM)
- GPU Architecture and Optimization — the CUDA programming model and hardware details
- Weight-Only Quantization and Mixed Precision — the fused implementation of the Marlin kernel
- Batching and Request Scheduling — piling on concurrency to drain bandwidth
- vLLM and PagedAttention — the industrial implementation of FlashAttention/Marlin
- TensorRT and GPU Inference — NVIDIA's operator-fusion stack
- Inference Benchmarking in Practice — measuring single-kernel performance
References
- Dao et al. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness (NeurIPS 2022) — the original FlashAttention v1 paper
- 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) — the FlashInfer paper
- Triton: An Intermediate Language and Compiler for Tiled Neural Network Computation — the Triton open-source repository and documentation
- NVIDIA CUTLASS Documentation — the CUDA C++ template library
- Tillet et al. Triton: an intermediate language and compiler for tiled neural network computation (MAPL 2019) — the Triton paper
- PyTorch 2.0 torch.compile Documentation — Inductor automatic fusion