1. What is it?#
Combining several operations into one kernel so intermediate results stay in registers or shared memory instead of round-tripping through HBM.
UNFUSED: FUSED:
kernel 1: read a, write b kernel 1: read a
kernel 2: read b, write c compute all three steps in registers
kernel 3: read c, write d write d
→ 6 HBM traversals → 2 HBM traversalsFor memory-bound operations — which is most of an LLM outside the big GEMMs — this is a 3x reduction in the thing that limits you.
2. Why does it exist?#
Because frameworks express models as many small operations, and each operation naively means a separate trip to memory. The math doesn’t require those trips; the API granularity does.
Fusion recovers the performance the abstraction cost you.
3. Simple analogy#
Cooking with a single pan versus washing up between steps.
Unfused: sauté onions, transfer to a bowl, wash the pan; sauté garlic, transfer, wash; combine. Every transfer is a round trip to the sink (HBM).
Fused: onions in the pan, add garlic to the same pan, add tomatoes to the same pan. One pan, one wash. The intermediate states never left the cooking surface (registers).
4. Tiny example#
// fusion.go — four elementwise ops as four passes over memory, then as one.
package main
import (
"fmt"
"time"
)
func bench(f func()) time.Duration {
f()
t0 := time.Now()
for i := 0; i < 3; i++ {
f()
}
return time.Since(t0) / 3
}
func main() {
const n = 1 << 25 // 32M elements, 128 MB per array
x, w := make([]float32, n), make([]float32, n)
a, b, c, out := make([]float32, n), make([]float32, n), make([]float32, n), make([]float32, n)
for i := range x {
x[i], w[i] = float32(i%13)-6, 0.5
}
unfused := bench(func() {
for i, v := range x {
a[i] = v * 2 // pass 1: read x, write a
}
for i, v := range a {
b[i] = v + 1 // pass 2: read a, write b
}
for i, v := range b {
c[i] = max(v, 0) // pass 3: read b, write c (ReLU)
}
for i, v := range c {
out[i] = v * w[i] // pass 4: read c, read w, write out
}
})
fused := bench(func() {
for i, v := range x {
out[i] = max(v*2+1, 0) * w[i] // one pass: read x, read w, write out
}
})
const bytes = n * 4
fmt.Printf("unfused %v (9 array passes = %d MB of traffic)\n", unfused.Round(time.Millisecond), 9*bytes>>20)
fmt.Printf("fused %v (3 array passes = %d MB of traffic)\n", fused.Round(time.Millisecond), 3*bytes>>20)
fmt.Printf("speedup %.2fx\n", float64(unfused)/float64(fused))
}Typical result on a laptop CPU: about 2.5x. Nine passes over memory become three, and the arithmetic did not change at all.
On a GPU the same experiment (a framework’s four separate elementwise kernels versus one fused kernel on a 134 MB FP16 tensor) gives 3.5-4x: unfused does 4 reads + 4 writes = 1.07 GB of HBM traffic, fused does 1 read + 1 write = 268 MB. The ratio is 4x and you measure ~3.5x, the difference being launch overhead and imperfect fusion.
5. Technical explanation#
The taxonomy of fusions#
1. ELEMENTWISE CHAINS (vertical fusion)
add → mul → silu → mul → one kernel
Easy, automatic, big wins. torch.compile does this well.
2. REDUCTION + ELEMENTWISE
rmsnorm = square → mean → rsqrt → mul → mul → one kernel
Needs a warp/block reduction inside. Standard in hand-written kernels.
3. GEMM EPILOGUE FUSION
matmul → bias → activation → (quantize) → store
The activation happens in registers before the accumulator is written.
CUTLASS/cuBLASLt support this natively. Huge win: avoids a full
read+write of the output tensor.
4. GEMM PROLOGUE FUSION
dequantize weights → matmul
Essential for W4A16: dequantization happens as weights are loaded into
registers, so the INT4 form is what crosses HBM.
5. HORIZONTAL FUSION
three independent matmuls on the same input → one bigger matmul
q_proj, k_proj, v_proj → one qkv_proj
gate_proj, up_proj → one gate_up_proj
Fewer launches, better tiles, one read of the input instead of three.
6. ALGORITHMIC FUSION
FlashAttention: matmul → scale → mask → softmax → matmul, all in SRAM.
Not a mechanical fusion — it required rethinking the algorithm.What can and cannot be fused#
CAN fuse:
elementwise ops with matching shapes
reductions with their consumers (with care)
a GEMM with anything elementwise on its output
a GEMM with anything elementwise on its inputs (prologue)
CANNOT easily fuse:
two GEMMs back to back (the intermediate is needed in full;
though "fused MLP" kernels exist for specific small shapes)
ops with a data-dependent shape between them
anything crossing a synchronization boundary
ops where the intermediate is used elsewhere (multiple consumers)That last one matters: a residual connection means the pre-norm value is used twice, which constrains fusion. Engines handle this by writing both outputs from one kernel.
The fusion boundary problem#
x = residual + attn_out ← elementwise
h = rmsnorm(x) ← reduction over d
gate, up = h @ W_gate_up ← GEMM
act = silu(gate) * up ← elementwise
out = act @ W_down ← GEMM
new_residual = x + out ← elementwise, uses x from the first line!The natural kernel boundaries are the GEMMs. So a well-fused layer looks like:
kernel 1: add_rmsnorm(residual, attn_out) → (x, h) [2 outputs]
kernel 2: GEMM(h, W_gate_up) with silu*up epilogue → act
kernel 3: GEMM(act, W_down) with residual-add epilogue → new_residualThree kernels for what PyTorch eager expresses as ~10.
Automatic vs manual fusion#
torch.compile / TorchInductor: excellent at elementwise + reduction fusion,
generates Triton kernels. Struggles with GEMM epilogues
(defers to cuBLAS) unless using max-autotune.
TensorRT: very good at the full range, including GEMM epilogues,
because it owns the whole graph and builds engines.
Hand-written: what vLLM/SGLang do for the hot paths
(add_rms_norm, silu_and_mul, rope, paged attention).
CUTLASS: lets you compose GEMM + arbitrary epilogue.In practice a production engine uses all four: hand-written kernels for the 10 hottest ops, CUTLASS/cuBLASLt for GEMMs with epilogues, and a compiler for the rest.
6. Under the hood#
What TorchInductor generates for the example in section 4:
@triton.jit
def fused_kernel(in_ptr0, in_ptr1, out_ptr, xnumel, XBLOCK: tl.constexpr):
xoffset = tl.program_id(0) * XBLOCK
xindex = xoffset + tl.arange(0, XBLOCK)[:]
xmask = xindex < xnumel
x = tl.load(in_ptr0 + xindex, xmask) # ONE load
w = tl.load(in_ptr1 + (xindex % 8192), xmask)
t = x * 2.0
t = t + 1.0
t = t * tl.sigmoid(t) # silu
t = t * w
tl.store(out_ptr + xindex, t, xmask) # ONE storeYou can see the generated code:
TORCH_COMPILE_DEBUG=1 python script.py
# writes generated Triton to torch_compile_debug/*/output_code.pyReading the generated code is the fastest way to learn what the compiler actually did. Do it at least once.
7. Performance implications#
Measured contributions on a 70B decode step (approximate, engine-dependent):
Baseline (eager PyTorch, no fusion) 100%
+ fused add_rms_norm -12%
+ fused silu_and_mul -6%
+ fused QKV and gate_up (horizontal) -8%
+ GEMM epilogue fusion (bias/activation) -4%
+ FlashAttention -10% (much more at long context)
+ CUDA graphs -15%
─────
~45% of baselineA well-fused engine is roughly 2x a naive PyTorch implementation before you touch precision or batching. That’s why “just use vLLM/TensorRT-LLM” is genuinely good advice.
8. Production implications#
- Use an engine that already has the fusions. Do not reimplement.
- Verify fusion is active. Profile and look at the kernel list. Seeing
aten::mul,aten::add,aten::rsqrtas separate kernels means you’re running unfused. torch.compilein production requires cache management. Compilation takes minutes; persist the cache (TORCHINDUCTOR_CACHE_DIR) and key it correctly.- Custom model architectures lose fusions. If you add a novel operation, the hand-written fused kernels won’t cover it and you fall back to eager for that part.
- Fusion changes numerics. Different operation order, different rounding. Validate.
9. Common mistakes#
Assuming PyTorch fuses automatically in eager mode. It doesn’t. Eager means one kernel per op.
Fusing without measuring. Some fusions increase register pressure and reduce occupancy, making things slower.
Expecting fusion to help compute-bound ops. Fusing two large GEMMs’ elementwise glue saves little when the GEMMs dominate.
Ignoring the multiple-consumer constraint. If an intermediate is used twice, naive fusion recomputes it.
Compiling at request time. Cold-start disaster.
Not validating numerics after enabling fusion.
10. Hands-on exercise#
A. Measure the chain. Run the section 4 benchmark. Then extend the chain to 8 operations and re-measure. Does the speedup grow linearly with chain length? Why or why not?
B. Read generated code. Run with TORCH_COMPILE_DEBUG=1 and read the generated Triton
kernel. Identify the loads and stores. Confirm there are only two.
C. Horizontal fusion. Implement q/k/v projections as three GEMMs and as one concatenated GEMM. Benchmark at M=1 and M=2048. Explain the difference at each.
D. Find unfused ops. Profile a real model and list every kernel by time. Which memory-bound kernels could be fused with a neighbor? Estimate the win.
E. Fusion regression. Construct a case where fusion makes things slower (hint: high register pressure, or a fused op whose intermediate is needed elsewhere). Measure it.
11. Interview questions#
- What is operator fusion and why does it help memory-bound operations so much?
- Name six kinds of fusion and give an example of each from a transformer.
- What is GEMM epilogue fusion and why is it particularly valuable?
- Why can’t you generally fuse two consecutive GEMMs?
- How would you verify that fusion is actually happening in production?
- When can fusion make things slower?
- How much of a well-tuned engine’s advantage over naive PyTorch comes from fusion?
12. Further reading#
- [REFERENCE] PyTorch
torch.compile/ TorchInductor design documents - [REFERENCE] CUTLASS epilogue documentation
- [REFERENCE] vLLM
csrc/layernorm_kernels.cu,csrc/activation_kernels.cu— small, readable, production fused kernels - [ESTABLISHED] Dao et al., FlashAttention — algorithmic fusion
- Next: 10 — Graph optimization and compilers