PidokuInfra

Kernel and Operator Fusion

Intermediate Advanced 1h Difficulty 3/5 Topic 07 of 14

Prerequisites IV.09, VI.03

Section IV.09 covered fusion mechanically. This file covers it as an optimization decision: what to fuse in an LLM, what it’s worth, and when it isn’t.


1. Problem → Why → Optimization#

PROBLEM   A transformer layer in eager PyTorch runs ~15 kernels, most of them
          memory-bound elementwise or reduction ops, each doing a full
          round trip to HBM.
WHY       Framework granularity (one op = one kernel) is much finer than the
          granularity at which HBM round trips are amortized.
OPTIMIZE  Combine adjacent operations into single kernels so intermediates
          stay in registers/SRAM.

TRADE-OFFS
  ✓ 2-4x on the fused ops; 10-25% end-to-end
  ✓ fewer launches (compounds with the CUDA graph benefit)
  ✗ more register pressure → possibly lower occupancy
  ✗ changed numerics (different rounding order)
  ✗ engineering effort for custom architectures

WHEN TO USE   Always, for the standard fusions. Your engine already does them.
WHEN NOT TO   When the ops are already compute-bound (fusing two big GEMMs'
              glue saves nothing), or when fusion causes spilling.

2. The fusion catalogue for LLMs#

The specific fusions that matter, in order of value:

1. add_rms_norm            residual add + RMSNorm          6 passes → 2      ~12% of decode
2. silu_and_mul            SiLU(gate) × up                 3 passes → 1      ~6%
3. QKV projection          3 GEMMs → 1                     3 launches → 1    ~4%
4. gate_up projection      2 GEMMs → 1                     2 → 1             ~3%
5. GEMM + bias + activation  epilogue fusion               2 passes → 1      ~4%
6. rope + cache write      RoPE then KV store              2 → 1             ~2%
7. FlashAttention          the whole attention block       ~6 kernels → 1    ~10%+
8. dequant + GEMM          prologue fusion (W4A16)         essential         —
9. sampling pipeline       temp/top-k/top-p/sample         5 → 1-2           ~2%

Together these account for roughly half of the gap between naive PyTorch and a production engine. (The other half is batching and paged KV.)


3. Simple analogy#

A production line versus separate workshops.

Separate workshops: each stage receives a crate, does one operation, packs it back into a crate, and ships it to the next workshop. Half the effort is packing and shipping.

Production line: the item moves along a belt through five operations without ever being packed.

The crate is HBM. The belt is registers.


4. Tiny example — the highest-value fusion#

add_rms_norm, which appears twice per layer:

Python
# UNFUSED (what eager PyTorch does)
def unfused(x, residual, weight, eps=1e-6):
    x = x + residual                      # read x, read residual, write t1
    var = x.pow(2).mean(-1, keepdim=True) # read t1, write t2 (via t1²)
    x_n = x * torch.rsqrt(var + eps)      # read t1, read t2, write t3
    return x_n * weight, x                # read t3, write out; also return t1
# ~6-8 full traversals of the tensor

# FUSED (one kernel)
#   load x tile and residual tile into registers
#   add
#   compute sum of squares via warp reduction
#   rsqrt, multiply by weight
#   store BOTH the normalized output and the new residual
# 2 reads + 2 writes

Measured on a (64, 8192) BF16 tensor:

unfused:  142 us    (1.05 GB effective traffic)
fused:     41 us    (0.27 GB)
speedup:  3.5x

At 80 layers × 2 norms × 5,000 tokens/sec, that’s a substantial share of your step time.


5. Technical explanation#

Why add_rms_norm returns two things#

The residual is needed by the next layer’s residual add:

x_new = x + attn_out          ← needed later
h     = rmsnorm(x_new)        ← fed to the FFN

Both outputs come from the same kernel. Writing both costs one extra store
but saves a full re-read.

This “return the intermediate too” pattern is common in fused kernels and is why their signatures look odd.

Epilogue fusion in GEMM#

CUTLASS/cuBLASLt let you attach an epilogue:

  D = activation(alpha × A×B + beta × C + bias)

executed in registers, before the accumulator is written to HBM.

Savings: one full read + write of the output tensor.
For a (2048, 14336) BF16 output: 118 MB of traffic avoided, per matrix, per layer.

Prologue fusion (for quantized weights)#

  load INT4 weights from HBM
  unpack + dequantize IN REGISTERS
  feed tensor cores

Without this, you'd dequantize to a full FP16 weight tensor in HBM first —
completely defeating the purpose.

Prologue fusion is not an optimization for W4A16; it’s a requirement.

Horizontal fusion (concatenated GEMMs)#

q = x @ Wq   (4096 → 4096)
k = x @ Wk   (4096 → 1024)      →   qkv = x @ [Wq|Wk|Wv]   (4096 → 6144)
v = x @ Wv   (4096 → 1024)           then split the output

Benefits:
  - x is read once instead of three times
  - one launch instead of three
  - one larger, better-shaped GEMM (better tiles, higher occupancy)

Done at model load time by concatenating the weight matrices. Gain: 5-15% at small batch, less at large.

When fusion doesn’t help#

1. The ops are compute-bound.
   Fusing an activation into a huge GEMM saves the output round trip —
   worthwhile — but fusing two GEMMs saves nothing (you need the full
   intermediate).

2. Register pressure causes spilling.
   A fused kernel needing 200 registers/thread may drop occupancy below
   the point where latency is hidden. Measure.

3. The intermediate has multiple consumers.
   Fusion either recomputes it or writes it anyway.

4. The ops are already tiny relative to everything else.
   Amdahl (file 01).

6-9. Under the hood, performance, production, mistakes#

Under the hood — a fused RMSNorm kernel structure:

one block per token (row of length d)
  each thread loads d/blockDim elements as float4/half8 (vectorized)
  add residual in registers
  partial sum of squares
  warp reduction (__shfl_down_sync) → block reduction (shared memory)
  broadcast rsqrt via shared memory
  each thread scales its elements and stores both outputs

Performance — cumulative effect on a 70B decode step:

Eager PyTorch, no fusion                      100%
+ add_rms_norm fused                           -12%
+ silu_and_mul fused                            -6%
+ QKV and gate_up horizontally fused            -7%
+ GEMM epilogue fusion                          -4%
+ FlashAttention                               -10%
+ fused RoPE + cache write                      -2%
                                              ──────
                                               ~59%   (1.7x)
+ CUDA graphs (compounds — fewer kernels to launch)  -12%
                                               ~47%   (2.1x)

Production:

  • Your engine already does these. Verify by profiling: you should see fused_add_rms_norm_kernel, silu_and_mul_kernel, flash_attn_*, not aten::add/aten::mul/aten::rsqrt.
  • Custom architectures lose fusions. A novel activation or norm variant falls back to eager. Budget for writing kernels or using torch.compile.
  • torch.compile covers the elementwise and reduction fusions automatically. It does not do GEMM epilogue fusion as well as TensorRT.
  • Validate numerics after enabling fusion. Different rounding.

Mistakes:

  • Assuming eager PyTorch fuses. It doesn’t.
  • Fusing without measuring. Register pressure can make it slower.
  • Fusing compute-bound ops and expecting a gain.
  • Writing custom kernels before enabling torch.compile.
  • Not checking that the engine’s fusions apply to your model’s architecture.

10. Hands-on exercise#

A. Measure the top fusion. Implement add_rms_norm unfused, with torch.compile, and with an existing fused kernel (from vllm or flash_attn). Benchmark all three. Report GB/s for each.

B. Horizontal fusion. Implement QKV as three GEMMs and as one. Measure at batch 1 and batch 2048. Explain why the gain differs.

C. Find unfused ops. Profile a model in eager mode. List every kernel by time. Identify which memory-bound kernels could be fused with a neighbor. Estimate the total win with Amdahl.

D. Register pressure. Write a deliberately over-fused kernel (fuse 10 elementwise ops) and measure register usage and occupancy. Does it get slower at some point?

E. Epilogue fusion. Using cublasLt or CUTLASS, run a GEMM with and without a fused bias+ activation epilogue. Measure the difference.


11. Interview questions#

  1. List the six most valuable fusions in an LLM and their approximate contributions.
  2. Why does add_rms_norm return two tensors?
  3. What is prologue fusion and why is it mandatory for W4A16?
  4. What is horizontal fusion and why does its benefit depend on batch size?
  5. When does fusion make things slower?
  6. How would you verify that fusion is active in a production engine?
  7. How much of a production engine’s advantage over eager PyTorch comes from fusion?

12. Further reading#

  • [REFERENCE] vLLM csrc/layernorm_kernels.cu, csrc/activation_kernels.cu
  • [REFERENCE] CUTLASS epilogue documentation
  • [REFERENCE] TorchInductor design notes
  • Next: 08 — CUDA graphs in serving

↑↓ navigate↵ openesc close