PidokuInfra

Compilers: torch.compile and Beyond

Expert Advanced 1h 15m Difficulty 4/5 Topic 12 of 12

Prerequisites IV.10, IV.11, 11


1. What a compiler does for inference#

INPUT:  a model expressed as framework operations
OUTPUT: fused, specialized, hardware-optimized kernels

  fusion              fewer HBM round trips (Section IV.09)
  specialization      constants folded, shapes known
  memory planning     buffer reuse (Section IV.10)
  kernel selection    the best implementation per shape
  layout assignment   minimize transposes
  codegen             Triton (Inductor) or a proprietary backend

Section IV.10 covered the theory. This file covers the practical use of torch.compile for inference and the landscape around it.


2. torch.compile for inference#

Python
model = AutoModelForCausalLM.from_pretrained(...).cuda().eval()

# Basic
model = torch.compile(model)

# For inference, the useful modes:
model = torch.compile(model, mode="reduce-overhead")   # adds CUDA graphs
model = torch.compile(model, mode="max-autotune")      # benchmarks kernels
model = torch.compile(model, dynamic=True)             # symbolic shapes
model = torch.compile(model, fullgraph=True)           # ERROR on graph breaks
MODES
  default          balanced fusion, quick compile
  reduce-overhead  + CUDA graphs → good for small batch (Section VII.08)
  max-autotune     + kernel autotuning → slowest compile, fastest result
  
FLAGS
  dynamic=True     symbolic shapes; avoids recompilation, costs 5-15%
  fullgraph=True   raises on a graph break instead of silently splitting
                   → USE THIS IN TESTING to find breaks

fullgraph=True in testing is the practical tip. It converts silent performance loss into a loud error, so you find and fix graph breaks during development rather than discovering them in a profile six months later.


3. Graph breaks — the thing that ruins it#

Shell
TORCH_LOGS="graph_breaks,recompiles" python your_script.py
COMMON CAUSES IN INFERENCE CODE
  .item()                       ← the classic; also causes a sync
  tensor.cpu(), .numpy(), .tolist()
  print(tensor)
  data-dependent control flow:  if x.sum() > 0: ...
  unsupported library calls
  Python exceptions in the traced region
  in-place mutation of module attributes
  calling into C extensions Dynamo doesn't understand

WHAT A BREAK COSTS
  the graph splits at that point. Fusion cannot cross the boundary.
  A model with 40 breaks gets a fraction of the possible benefit.
Python
# BAD — a graph break AND a sync per step
if logits.argmax().item() == eos_token_id:
    break

# GOOD — keep it on the GPU; check every N steps
finished = (next_token == eos_token_id)      # a tensor
if step % 16 == 0 and finished.all().item():  # sync rarely
    break

Both problems (graph breaks and implicit syncs) usually have the same cause: moving data to the host in the hot loop. Fixing one fixes the other.


4. Recompilation#

Dynamo specializes on shapes by default. A new shape → recompile.

  batch 1  → compile
  batch 2  → compile
  batch 4  → compile
  ...
  after cache_size_limit (default 8) recompilations:
    "torch._dynamo hit config.cache_size_limit"
    → FALLS BACK TO EAGER PERMANENTLY for that function

That fallback is silent and permanent. It’s a real production failure mode: performance degrades and nothing errors.

Python
# Mitigations
torch._dynamo.config.cache_size_limit = 32        # allow more variants

model = torch.compile(model, dynamic=True)        # symbolic shapes,
                                                  # no recompilation

torch._dynamo.mark_dynamic(input_tensor, 0)       # mark the batch dim dynamic,
                                                  # keep others static
FOR LLM INFERENCE
  decode shapes: (batch, 1, d) — only batch varies
    → mark the batch dimension dynamic, or bucket it (as CUDA graph
      capture already does — Section VII.08)
  prefill shapes: (batch, seq, d) — both vary
    → dynamic=True, or accept that prefill runs less optimized

Monitor for recompilation in production:

Shell
TORCH_LOGS="recompiles" python serve.py 2>&1 | grep -c "Recompiling"

5. Compilation cost and caching#

COMPILE TIME (a 7B model)
  default mode:      20-60 s
  reduce-overhead:   40-90 s
  max-autotune:      3-15 minutes

→ this is added to your COLD START (Section II.07)

CACHING
  export TORCHINDUCTOR_CACHE_DIR=/persistent/inductor-cache
  export TRITON_CACHE_DIR=/persistent/triton-cache

  The cache is keyed by: the graph, shapes, GPU architecture,
  PyTorch version, and Triton version.
  → invalidated by any upgrade
  → build it at IMAGE BUILD TIME, or on a persistent volume

Cache the compilation artifacts. Otherwise every pod start pays 60 seconds to 15 minutes. This is the same lesson as TensorRT engines (Section VII.09).


6. The compiler landscape#

TORCH.COMPILE / INDUCTOR          [ESTABLISHED]
  ✓ integrated with PyTorch; works on any model
  ✓ generates Triton → portable to AMD
  ✓ good elementwise and reduction fusion
  ✗ defers GEMM to cuBLAS (except in max-autotune)
  ✗ graph breaks and recompilation are real friction
  → the default choice

TENSORRT / TENSORRT-LLM           [ESTABLISHED]
  ✓ the most aggressive optimization; measured kernel selection
  ✓ excellent FP8 support
  ✗ AOT: rebuild per GPU arch, per shape range (Section VII.09)
  ✗ NVIDIA only
  → for stable, high-volume, NVIDIA deployments

ONNX RUNTIME                       [ESTABLISHED]
  ✓ cross-platform, many execution providers
  ✗ export lags for new architectures; limited LLM-specific features
  → for small models, CPU, edge (Section VIII.13)

XLA (JAX / PyTorch-XLA)            [ESTABLISHED]
  ✓ excellent on TPU; strong whole-graph optimization
  ✗ requires static shapes; different programming model
  → if you're on TPU

TVM / MLIR-based stacks            [EMERGING]
  ✓ research vehicle; hardware portability
  ✗ less mature for LLM inference specifically

ENGINE-INTERNAL COMPILATION        [ESTABLISHED]
  vLLM, SGLang increasingly use torch.compile internally,
  plus hand-written kernels for the hot paths
  → you often get compilation without invoking it yourself

7. What compilation is and isn’t worth#

TYPICAL GAINS 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
  TensorRT-LLM                     2-5x vs eager (but vs vLLM: 1.05-1.25x)

CONTEXT: a production engine (vLLM/SGLang) already gives 3-10x over
         eager PyTorch, and MOST of that is continuous batching and
         paged KV — NOT compilation (Section IV.10).

→ compilation is a 20-60% lever.
→ scheduling is a 300-1000% lever.
→ do them in that order (Section VII.01).

This is the honest framing. Compilation is worthwhile and you should use it, but it is not where the big wins are, and teams that focus on it before fixing their batching are optimizing the wrong thing.


8. Numerics#

COMPILATION CHANGES RESULTS.

  fusion changes the order of operations
  different kernel selection changes reduction order
  max-autotune may select a numerically different algorithm

  → outputs differ in the last bits (Section IV.12)
  → after sampling, that can change the generated text

VALIDATE after enabling compilation:
  □ logit difference vs eager (max, mean)
  □ top-1 agreement
  □ generation comparison on a corpus
  
This is the same discipline as validating a quantization change.

9. Production implications#

  • Use compilation. It’s a real 20-60%.
  • Cache the artifacts. TORCHINDUCTOR_CACHE_DIR and TRITON_CACHE_DIR on persistent storage, or build at image build time.
  • Use fullgraph=True in testing to find graph breaks.
  • Monitor for recompilation in production. The silent eager fallback is a real failure mode.
  • Set cache_size_limit appropriately, or use dynamic=True.
  • Validate numerics after enabling it.
  • Fix scheduling first. Compilation is not where the big wins are.
  • You may already have it: modern engines use torch.compile internally.

10. Common mistakes#

Compiling in the request path. Minutes of stall.

Not caching artifacts. Every cold start recompiles.

Ignoring graph breaks. Silently losing most of the benefit.

Hitting cache_size_limit. Silent permanent fallback to eager.

.item() in the decode loop. Graph break and sync.

Not validating numerics.

Expecting compilation to fix a scheduling problem.

Deploying a TensorRT engine built for a different GPU architecture.


11. Hands-on exercise#

A. Measure the gain. For a model you can run, measure throughput and latency in eager, torch.compile default, reduce-overhead, and max-autotune. Record compile time for each.

B. Find graph breaks. Run with fullgraph=True and fix the errors one at a time. Measure the improvement after each fix.

C. Trigger the recompilation cliff. Call a compiled model with 12 different shapes. Watch it hit cache_size_limit and fall back. Then set dynamic=True and confirm it doesn’t.

D. Measure the caching benefit. Time a cold start with and without a warm TORCHINDUCTOR_CACHE_DIR. How much does caching save?

E. Read the generated code. With TORCH_COMPILE_DEBUG=1, find the generated Triton for a fused RMSNorm. Compare to your hand-written version from Section XIII.11. Which is better?

F. Validate numerics. Compare eager and compiled outputs on 200 prompts at temperature 0. How often do they differ? At what token index do generations diverge?

G. Put it in context. Measure: eager PyTorch, eager + continuous batching, compiled + continuous batching. Which contributed more?


12. Interview questions#

  1. What does a compiler do for inference, and which optimizations require graph visibility?
  2. What is a graph break and how do you find them?
  3. What happens when you hit cache_size_limit, and why is it dangerous?
  4. Why must you cache compilation artifacts?
  5. How much does compilation contribute relative to scheduling improvements?
  6. Why does compilation change numerical results, and what do you do about it?
  7. When would you choose TensorRT-LLM over torch.compile?

13. Further reading#

  • [REFERENCE] PyTorch 2 paper (ASPLOS 2024) and torch.compile documentation
  • [REFERENCE] TorchInductor design notes
  • [REFERENCE] TensorRT-LLM documentation (Section VII.09)
  • [ESTABLISHED] Chen et al., “TVM” (OSDI 2018) — for the general approach
  • Next: Section XIV — Research and Frontier

↑↓ navigate↵ openesc close