Skip to content

Kernel Fusion and Custom Kernels

At a glance Kernel fusion merges multiple elementwise/reduction operators into a single kernel, avoiding HBM round trips for intermediate tensors — the core technique of LLM inference optimization. This article covers classic fusion patterns, FlashAttention v1/v2/v3's SRAM tiling, FlashInfer's block sparsity, the Triton/CUDA programming models, and how to judge fusion boundaries.

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:

  1. 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;
  2. 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 HBM

Each 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 kernel

Any 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. · γ + β                  # elementwise

Five 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 FFN

cuBLASLt'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 HBM

S 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_triton is 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 ​

  1. Block-sparse attention: supports 2D block-sparse masks (e.g. the Longformer pattern of sliding window + global tokens);
  2. Composable KV cache formats: the same KV cache can be consumed by different attention kernels (vLLM PagedAttention's paged KV is reused directly);
  3. Append / fork operations: supports incremental appending for long contexts (for speculative decoding verification), tree decoding (MTP multi-token prediction);
  4. 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:

ToolAbstraction LevelDifficultyBest For
PyTorch eager + torch.compileHighLowAutomatic fusion, limited coverage
Triton (OpenAI)MediumMediumThe current default for LLM kernels
CUDA C++LowHighUltimate performance, highest bar
TVM / HalideMediumMedium-highCompiler-style, research-friendly
CUTLASS (NVIDIA)MediumMedium-highLarge 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 ​

LevelResourceCapacity (H100 SM)Role
ThreadRegisters255/threadFastest, private per thread
WarpWarp shuffle—Ultra-fast communication among 32 threads
BlockShared memory228 KBShared within a block, ~10 TB/s
SML1 cache256 KBShared across the SM
GPUL2 cache50 MBShared across all SMs
GPUHBM80 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 both

In 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 reduction

But 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/CompilerFusion CapabilityRepresentative
PyTorch Inductor (torch.compile)Automatic elementwise + reduction fusionPyTorch 2.x
XLAOperator fusion + layout propagationJAX, TF
TVM RelaxLLM inference fusion + auto-schedulingApache TVM
TensorRTMulti-operator fusion + INT8/FP8 epilogueNVIDIA
vLLMBuilt-in Marlin/AWQ/FlashAttention kernelsServing-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 ​

References ​