Kernel Fusion at the PyTorch Level¶
Overview¶
Kernel fusion = combine several ops into one kernel: launch once, keep intermediates in registers/SRAM, avoid round-trips to DRAM. It's the single most impactful graph optimization— and torch.compile's Inductor does most of it for you. Knowing what it fuses (and why) lets you write code that compiles into better kernels.
- Elementwise chains (
x * w + b -> relu) fuse into one Triton/CUDA kernel. - Reductions + elementwise (layernorm, softmax+dropout+matmul adjacencies) fuse via "fused" kernels.
- Matmul (GEMM) rarely fuses through— it's a library call (cuBLAS) with its own layout rules.
- Inductor generates Triton kernels with constant-folding, dead-code elimination, and layout specialization.
The classic rule: fuse vertically (elementwise chains), don't fuse across GEMM boundaries— keep GEMMs big.
-
What Fuses Well¶
import torch, torch.nn.functional as F
def fused_example(x, w, b):
return F.relu(x @ w + b) # gemm, add, relu -> add+relu fuse into one epilogue kernel
def bad_example(x, w1, w2):
return (x @ w1) * (x @ w2) # two separate GEMMs w/ separate epilogues
# With torch.compile, the elementwise parts around gemms fuse automatically.
Manual Fusion vs torch.compile¶
| Approach | Effort | Result |
|---|---|---|
| Hand-write fused CUDA/Triton (Ch 9) | high | perfect, portable |
torch.compile + Inductor |
low | near-optimal for elementwise clusters |
| torch.fx passes (Ch 3-03) | med | fuse specific patterns yourself |
Strategy: write clear, idiomatic eager code → compile → profile → hand-fuse only what's left.
-
Inductor's Fusion Playbook (what it actually does)¶
- Elementwise clustering— merges
add/sub/mul/div/exp/sqrt/relu/gelu/...into one kernel. - Epilogue fusion— fold bias+activation into GEMM output (via Triton templates like
MM + bias + relu). - Reduction fusion— layernorm/softmax compute mean/var inline, no separate passes.
- Constant folding—
x*0,1/xwith constant x, dead branches removed. - Layout decisions— pick channels_last / bf16 / fp32 epilogues automatically.
# this compiles to
def layer(x, w, b):
return F.layer_norm(x @ w + b, x.shape[-1:])
-
Reductions & the Memory Wall¶
def naive_softmax(x):
m = x.max(dim=-1, keepdim=True).values # 1st pass over memory
e = (x - m).exp() # 2nd pass
return e / e.sum(dim=-1, keepdim=True) # 3rd pass
def fused_softmax(x):
return F.softmax(x, dim=-1) # single kernel: online softmax
# memory traffic ~3x less in the fused version; torch.compile fuses naive too.
Example: attention with
F.scaled_dot_product_attention(SDPA) is one fused kernel— never hand-roll attention without it (Ch 09).
-
How to Help Inductor¶
# 1. Use functional ops instead of python loops
y = torch.addcmul(x, w, b) # one op vs 3
# 2. Avoid unnecessary.item()/.tolist() inside compiled regions (graph breaks)
# 3. Keep shapes static, use contiguous inputs
# 4. Prefer bf16/fp16 epilogues when memory-bound
torch.compile(model, mode="max-autotune")
-
Fusing Across GEMM Boundaries (advanced)¶
Sometimes you do want cross-GEMM fusion— e.g., fusing the QKV projection and attention reads. Inductor handles parts; for the rest you need custom Triton kernels:
# torch._inductor custom op hooks / triton kernels (Ch 9) let you fuse exactly this
# SDPA is itself the canonical case of a fused attention kernel
attn = F.scaled_dot_product_attention(q, k, v, is_causal=True)
Detection: Did Fusion Happen?¶
import torch
from torch.profiler import profile, ProfilerActivity
def check_kernels(fn, x):
with profile(activities=[ProfilerActivity.CUDA]) as p:
fn(x); torch.cuda.synchronize()
for evt in p.key_averages():
if evt.device_type == torch.autograd.DeviceType.CUDA:
print(evt.key, round(evt.self_device_time_total, 2))
Fewer CUDA kernels for the same logical work = fusion working.
-
Key Takeaways¶
- Fusion kills launch + DRAM traffic, not arithmetic— biggest wins on elementwise/reduction-heavy code.
torch.compile/Inductor fuses elementwise clusters, epilogues, and reductions; GEMM boundaries stay library calls.- Write idiomatic, functional code; fuse by hand (Triton/CUDA) only after profiling.
- SDPA is the golden example: one fused attention kernel, never reimplement.
- Use profiler op/kernel counts to verify fusion took effect.
-
Related Topics¶
- Torch.Compile
- [Profiling](/06-pytorch/04-performance-and-compilation/(04-profiling-benchmarking/)
- Custom Kernels