1. What is it?#
Taking the computational graph and rewriting it into a faster but mathematically equivalent graph, then generating code for it.
Original graph → [analysis] → [rewrite] → [schedule] → [codegen] → executableThe systems that do this: TorchInductor (torch.compile), TensorRT, ONNX Runtime, XLA, TVM,
Apache TVM’s successor projects, and each vendor’s proprietary stack.
2. Why does it exist?#
Because the graph a human writes is optimized for readability, and the graph a GPU wants is optimized for memory traffic, kernel count, and tile shapes. A compiler bridges the two automatically, so you get most of the hand-tuning benefit without hand-tuning.
3. Simple analogy#
A logistics planner rewriting a delivery route. The route you’d naturally describe (“customer A, then B, then C”) is not the route that minimizes miles. A planner reorders, merges stops, drops redundant legs, and picks vehicle sizes — same deliveries, less fuel.
Constant folding is “you always pass the depot, so pick that up in advance.” Fusion is “combine these three stops on the same street.” Layout assignment is “load the truck in delivery order.”
4. Tiny example#
import torch
def f(x, w, b):
scale = torch.tensor(1.0) / torch.sqrt(torch.tensor(64.0)) # constant!
y = x @ w
y = y * scale
y = y + b
y = torch.nn.functional.gelu(y)
z = y + 0.0 # dead
return z
cf = torch.compile(f, mode="max-autotune")A compiler applies:
1. Constant folding: scale = 0.125, computed at compile time
2. Dead code elimination: z = y + 0.0 → z = y
3. Algebraic simplification: (x@w)*0.125 → fold 0.125 into the GEMM's alpha
4. Fusion: bias + gelu into the GEMM epilogue
5. Layout: choose the layout for w that the chosen GEMM kernel prefers
6. Autotuning: benchmark several GEMM configurations, pick the fastest
Result: ONE kernel instead of five, with the scale folded into it.5. Technical explanation#
The optimization catalogue#
GRAPH LEVEL
Constant folding evaluate constant subgraphs at build time
Dead code elimination drop unused nodes
Common subexpression compute shared subgraphs once
Algebraic simplification x*1 → x; (A^T)^T → A; concat/split cancellation
Operator substitution replace with a cheaper equivalent
Fusion file 09
Layout assignment choose formats globally, minimize transposes
Memory planning buffer reuse via live-range analysis
Constant/weight prepacking permute weights offline for the chosen kernel
KERNEL LEVEL
Tiling choose tile sizes for the memory hierarchy
Vectorization use wide loads/stores
Unrolling reduce loop overhead
Software pipelining overlap load and compute
Autotuning benchmark variants, cache the winnerThe three compilation strategies#
AHEAD-OF-TIME (TensorRT, TVM)
Build an engine offline for a fixed model + shapes + hardware.
✓ Best performance; full graph visibility; extensive autotuning
✗ Build takes minutes to hours; must rebuild per GPU arch and shape set
✗ Dynamic shapes need explicit profiles
JUST-IN-TIME (torch.compile)
Compile on first call for each observed shape signature; cache.
✓ Flexible; handles Python; incremental
✗ First-call latency; recompiles on shape changes; graph breaks
RUNTIME LIBRARY (cuBLAS, cuDNN heuristics)
Pick a pre-built kernel per call based on shape.
✓ Zero build time
✗ No cross-operator optimizationProduction LLM engines typically use a hybrid: hand-written kernels (library), CUDA graph
capture (AOT-ish), and optionally torch.compile for the parts not hand-written.
torch.compile mechanics#
TorchDynamo → traces Python bytecode into FX graphs, with "graph breaks"
at anything it can't handle
AOTAutograd → separates forward/backward (irrelevant for inference)
TorchInductor → lowers the FX graph to Triton (GPU) or C++/OpenMP (CPU)
Triton → compiles to PTX → SASSModes:
default balanced compile time and performance
reduce-overhead adds CUDA graphs — good for small batch
max-autotune benchmarks many kernel configs — slow to compile, fastest resultGraph breaks are the main practical issue. Anything Dynamo can’t trace (a .item() call,
data-dependent control flow, an unsupported library call, a print) splits the graph. Each
break means the optimization can’t cross that point.
TORCH_LOGS="graph_breaks,recompiles" python script.pyRun this on your model. A model with 40 graph breaks is getting a fraction of the benefit.
TensorRT-LLM’s approach#
1. Define the network in TensorRT-LLM's Python API (or convert from HF)
2. Specify: max batch, max input len, max output len, precision, parallelism
3. Build → an "engine" file (binary, GPU-arch-specific)
4. Serve the engine via Triton Inference Server or the C++ runtimeBecause everything is fixed at build time, TensorRT can be extremely aggressive: full fusion, optimal kernel selection by measurement, precise memory planning. The cost is inflexibility — change the max batch size and you rebuild.
TensorRT-LLM typically leads on raw throughput for fixed configurations; vLLM/SGLang lead on flexibility and feature velocity. Both are correct answers depending on your constraints.
6. Under the hood#
Memory planning, concretely:
Naive: every intermediate gets its own buffer
peak = sum of all intermediates
Planned: compute live ranges; two tensors whose ranges don't overlap share a buffer
peak = max over time of concurrently-live bytes
Transformer layer: naive peak ≈ 8 × activation size
planned peak ≈ 2-3 × activation sizeFor long-context prefill this is the difference between OOM and working.
7. Performance implications#
Typical speedups over eager PyTorch for LLM inference:
torch.compile (default) 1.1-1.3x
torch.compile (reduce-overhead) 1.3-1.6x at small batch
torch.compile (max-autotune) 1.3-1.8x
Hand-written engine (vLLM) 2-4x (fusion + batching + paged KV)
TensorRT-LLM 2-5x, best on fixed shapesNote that most of vLLM’s advantage is not compilation — it’s continuous batching and paged KV cache, which are scheduling and memory-management wins, not codegen wins. Compilation is a 20-60% lever; scheduling is a 3-10x lever. Prioritize accordingly.
8. Production implications#
- Compile at build time, cache the artifacts.
TORCHINDUCTOR_CACHE_DIRfor PyTorch, a built engine file for TensorRT. Key the cache by (model, GPU arch, precision, shape set, library versions). - Rebuild on GPU architecture change. An engine built for H100 will not run on A100.
- Monitor recompilations.
TORCH_LOGS="recompiles"— in production, unbounded recompilation from varying shapes is a real failure mode (each recompile is seconds of stall). - Set
torch._dynamo.config.cache_size_limitappropriately; when exceeded, Dynamo gives up and falls back to eager permanently for that function. - Validate numerics after compiling. Fusion and different kernel selection change rounding.
9. Common mistakes#
Compiling in the request path. Minutes of stall.
Ignoring graph breaks. Silently losing most of the benefit.
Unbounded shape variation causing recompilation storms. Use dynamic shape support or bucket.
Assuming compilation solves your problem. If you’re at batch 2 because your scheduler is bad, a 30% codegen win is not the fix.
Deploying an engine built on a different GPU.
Not validating outputs after enabling max-autotune.
10. Hands-on exercise#
A. Compile a model. Run a small transformer in eager, torch.compile default, and
max-autotune. Measure compile time and inference time for each at batch 1 and 64.
B. Find graph breaks. Run with TORCH_LOGS="graph_breaks". Fix at least one break and
measure the improvement.
C. Recompilation storm. Run a compiled model with 20 different sequence lengths. Count
recompilations with TORCH_LOGS="recompiles". Then enable dynamic=True and compare.
D. Memory planning. Measure peak memory for a long-context prefill in eager vs compiled. Explain the difference.
E. Numerics. Compare eager and compiled outputs bit by bit for the same input. How large is the difference? Does it change the sampled token?
11. Interview questions#
- Name six graph-level optimizations and give an example of each.
- Compare AOT and JIT compilation for inference. When would you choose each?
- What is a graph break and how do you find them?
- Why must a TensorRT engine be rebuilt per GPU architecture?
- How much of vLLM’s advantage over eager PyTorch comes from compilation vs scheduling?
- What is memory planning and how much does it save?
- How would you manage compilation artifacts in a production deployment?
12. Further reading#
- [REFERENCE] PyTorch 2 paper: “PyTorch 2: Faster Machine Learning Through Dynamic Python Bytecode Transformation and Graph Compilation” (ASPLOS 2024)
- [REFERENCE] TensorRT-LLM documentation, especially the build/engine sections
- [ESTABLISHED] Chen et al., “TVM” (OSDI 2018)
- [REFERENCE] Triton language documentation
- Next: 11 — Static vs dynamic shapes