torch.compile is a just-in-time compiler for PyTorch programs that keeps eager-mode semantics. You wrap a function or module, keep writing ordinary Python, and on the first call PyTorch captures the tensor operations into a graph, generates fused GPU kernels for it and caches the result. Later calls that match the conditions the graph was built under run the compiled code; anything the compiler cannot handle runs eagerly, as before. That fallback is why it is safe to try on almost any model, and also why it can quietly deliver nothing.

Getting the speedup you expect requires knowing what the compiler sees. This article walks the stack from Python bytecode to Triton kernels, explains the two events that decide whether compilation helps (graph breaks and recompilations), and covers modes, dynamic shapes, caching and debugging. Speed figures are deliberately absent: they depend on your model and GPU, so measure them yourself.

Advertisement

The pipeline at a glance

Python functionnn.Module.forwardTorchDynamobytecode to FX graphgraphAOTAutogradjoint fwd + bwd graphTorchInductorfusion, schedulingTriton / C++ kernels+ cuBLAS for matmulsmodeCUDA graph (optional)reduce-overheadguardsGuard set per graphshapes, dtypes, globalsEvery call: evaluate guardshit: run cached code / miss: recompilecached artefactGraph breakunsupported op: split into subgraphs, run gap in eager PythonOn-disk cachesTORCHINDUCTOR_CACHE_DIR, Mega-Cache artefacts
torch.compile stack. Dynamo captures a graph plus the guards that must hold for it to be reused; AOTAutograd adds the backward pass; Inductor emits fused kernels. Unsupported code becomes a graph break and runs eagerly between compiled subgraphs.

Three components do the work, and each can be used on its own. TorchDynamo captures graphs from Python. AOTAutograd turns a forward graph into a joint forward and backward graph ahead of time. TorchInductor, the default backend, compiles the graphs into kernels: Triton on GPUs, C++ with OpenMP on CPUs, with matrix multiplications usually handed to vendor libraries unless autotuning chooses a generated template. The signature is torch.compile(model, *, fullgraph=False, dynamic=None, backend="inductor", mode=None, options=None, ...), and every argument maps onto one of these stages.

Stage 1: Dynamo captures graphs from bytecode

Tracing by running sample inputs silently bakes in whichever branch the sample took, so Dynamo works differently. It hooks CPython's frame evaluation (the PEP 523 API), so when a compiled function is called, Dynamo sees the function's bytecode before it runs. It then interprets that bytecode symbolically: tensor operations are recorded into an FX graph, while ordinary Python such as loops over a module list, attribute lookups and arithmetic on Python ints is executed at compile time and disappears from the graph. A for layer in self.layers loop is unrolled into one straight-line graph.

Because it specialises on what it saw, Dynamo also emits guards: cheap checks that must be true for the graph to be valid again. Typical guards test tensor dtype, device and sizes, requires_grad, and the Python globals, attributes and ints that steered control flow. On every call the guards run first. If they pass, the cached code runs; if not, Dynamo compiles again and stores another entry for that function.

When Dynamo meets something it cannot represent, it produces a graph break. It compiles the graph so far, emits bytecode that runs the unsupported part in normal Python, and starts a new graph afterwards. Common causes are branching on .item() or .tolist() (the value is on the GPU and unknown at compile time), print and logging, and calls into C extensions Dynamo does not understand. The program stays correct, but each break costs a Python round trip, often a device-to-host synchronisation, and it cuts fusion across the break. Passing fullgraph=True turns any break into an error, which is the right setting once a model is clean.

Advertisement

Stage 2: AOTAutograd traces the backward pass ahead of time

Eager autograd builds the backward pass on the fly, which is too late for compilation. AOTAutograd traces the forward graph through autograd once, producing a joint forward and backward graph decomposed into core ATen operators. A partitioner splits it into forward and backward graphs and decides which intermediates to save and which to recompute.

That decision is a real memory trade-off. Saving an activation costs memory for the whole step; recomputing it costs arithmetic. For chains of cheap pointwise operations, recomputing inside a fused backward kernel is often cheaper than writing and re-reading the saved tensor, so compiled training can use less activation memory than eager. Measure peak memory as well as step time.

Stage 3: Inductor generates fused kernels

Inductor lowers the graph into a loop-level intermediate representation and schedules it. Its biggest lever is fusion: pointwise operations and reductions that consume each other's outputs are merged into one kernel, so intermediates stay in registers instead of travelling through GPU memory. The theory is covered in GPU kernel fusion and the memory levels involved in the GPU memory hierarchy.

A worked example shows why this matters. Take y = gelu(x @ w + b) * s where the matmul output is 8192 by 4096 in bfloat16, which is 64 MiB. Eagerly, the bias add, GELU and scale are three kernels, and each reads 64 MiB and writes 64 MiB, so the epilogue moves about 384 MiB. Fused, one kernel reads the matmul output once and writes the result once, about 128 MiB, a third of the traffic for the same arithmetic. Because these operations are memory-bound, runtime tracks bytes moved, so the saving is roughly proportional. Under max-autotune, Inductor can also benchmark Triton matmul templates and fuse the epilogue into the matmul itself, removing the round trip altogether. The generated kernels are ordinary Triton you can read; the Triton deep dive explains the programming model they use.

Using it in a training loop

import torch

model = build_model().cuda()    # model.layers: nn.ModuleList of 12 identical blocks

# Option A: whole-model compilation. One large graph, longest compile, most fusion.
# compiled = torch.compile(model)

# Option B: regional compilation. Every Block has identical code, so Dynamo compiles it
# once and reuses the artefact for the other 11 layers: far shorter cold start.
for layer in model.layers:
    layer.compile()

opt = torch.optim.AdamW(model.parameters(), lr=3e-4)
for step, (x, y) in enumerate(loader):          # loader yields fixed-shape CUDA batches
    with torch.autocast("cuda", dtype=torch.bfloat16):
        loss = torch.nn.functional.mse_loss(model(x), y)
    loss.backward()                              # backward graph was compiled by AOTAutograd
    opt.step()
    opt.zero_grad(set_to_none=True)

Two choices in that loop matter. Whole-model versus regional compilation. Compiling the whole model gives Inductor the most freedom but compiles one large graph, and cold start grows with depth. Compiling each repeated block instead means Dynamo compiles the block's code once and reuses the artefact for every instance that passes the same guards, which cuts cold-start time sharply for deep transformer stacks at the cost of no fusion across block boundaries. nn.Module.compile() compiles a module in place, keeping its state-dict keys unchanged, whereas torch.compile(module) returns a wrapper whose state-dict keys gain an _orig_mod. prefix, a classic source of checkpoint-loading errors.

Where compilation starts. Logging stays outside the compiled region deliberately: a loss.item() inside it would create a graph break and a synchronisation every step.

Modes and CUDA graphs

modeWhat it addsCostsUse when
defaultFusion and code generation with a balance between runtime and compile timeModerate compile timeThe starting point for training and inference
reduce-overheadCaptures the compiled code into CUDA graphs, removing per-kernel launch overhead from Python and the driverExtra memory for static buffers; needs stable shapesSmall batches or many tiny kernels, where the CPU cannot launch fast enough
max-autotuneBenchmarks Triton or template matmul variants and picks the fastest; CUDA graphs on by defaultMuch longer compile timeLong-running jobs and serving, where compile cost amortises
max-autotune-no-cudagraphsAutotuning without CUDA graphsLong compile timeYou want tuned kernels but CUDA graphs conflict with your memory or input patterns

A CUDA graph records a sequence of kernel launches and replays it with one call. That removes launch overhead but freezes memory addresses: inputs are copied into static buffers, outputs live in memory the next replay overwrites, and a new shape means a new graph. CUDA streams and graphs covers the mechanism at the CUDA level.

Dynamic shapes and recompilation

By default Dynamo specialises on the exact tensor sizes of the first call. With dynamic=None (the default), a second call with a different size in some dimension triggers a recompilation in which that dimension is marked dynamic and represented symbolically, so a third size usually reuses the second graph. dynamic=False always specialises, and dynamic=True asks for dynamic kernels from the start. Dynamic kernels avoid recompiles but can be slightly slower, since the compiler cannot use the concrete size, for example to drop bounds checks or pick a tile size.

Every recompilation adds an entry to the function's cache. When torch._dynamo.config.cache_size_limit is reached (8 by default), Dynamo stops compiling new variants, so any call that misses the cache runs eagerly. This is the most common way for a compiled model to become slow without an error: something the guards depend on keeps changing, such as a mutable global or module attribute updated every step, or input shapes run with dynamic=False so that every new length is a new specialisation, and the entries pile up until the limit is hit. Fix the cause rather than raising the limit. Pass changing scalars as tensors, pad inputs to a small set of shape buckets, or mark a dimension up front with torch._dynamo.mark_dynamic(tensor, dim). Serving systems generally bucket sequence lengths, for example to multiples of 64, and warm up each bucket.

Compile time and caching

The first call pays for tracing, code generation and, in autotuning modes, benchmarking; for a large model that can be minutes, repeated in every fresh process unless cached. Inductor keeps local caches under TORCHINDUCTOR_CACHE_DIR, which defaults to a per-user directory under /tmp. In containers that directory disappears with the container, so either mount it on persistent storage or use the end-to-end cache, known as Mega-Cache.

Mega-Cache has two calls: torch.compiler.save_cache_artifacts() returns (artifact_bytes, cache_info) after you have compiled and run the model, and torch.compiler.load_cache_artifacts(artifact_bytes) pre-populates the caches in a new process, possibly on another machine. The artefacts are only valid with the same PyTorch and Triton versions and, for CUDA, the same GPU, so key your stored blob on all three.

import torch

compiled = torch.compile(model, mode="max-autotune")
for shape in serving_shapes:                     # warm up every shape bucket you will serve
    compiled(torch.randn(*shape, device="cuda"))

artifact_bytes, cache_info = torch.compiler.save_cache_artifacts()
blob_store.put("compile-cache/model-v7", artifact_bytes)

# In CI or a canary: after warm-up, any recompilation is a bug, so make it an error.
with torch.compiler.set_stance("fail_on_recompile"):
    compiled(torch.randn(*serving_shapes[0], device="cuda"))

# On a fresh replica, same PyTorch, Triton and GPU model:
torch.compiler.load_cache_artifacts(blob_store.get("compile-cache/model-v7"))

torch.compiler.set_stance controls behaviour without editing call sites. Its stances are default, force_eager (ignore all compile directives, useful for A/B comparisons), eager_on_recompile (use cached code where valid, otherwise run eagerly rather than compile) and fail_on_recompile (raise instead of recompiling). The last turns silent performance regressions into test failures.

A debugging workflow

import torch

def f(x):
    y = torch.sin(x) * 2
    if y.sum().item() > 0:        # .item() pulls a value to Python: a graph break
        return y.cos()
    return y.tanh()

x = torch.randn(4096, device="cuda")
report = torch._dynamo.explain(f)(x)
print(report.graph_count, report.graph_break_count)
for reason in report.break_reasons:
    print(reason)

#   TORCH_LOGS="graph_breaks,recompiles" python train.py
#   TORCH_TRACE=/tmp/tc_trace python train.py && tlparse /tmp/tc_trace

torch._dynamo.explain runs the function once and reports how many graphs and breaks it produced and why. On a real job, set TORCH_LOGS="graph_breaks,recompiles" to log each break with its source line and each recompilation with the guard that failed. For a complete picture, TORCH_TRACE writes structured logs that the separate tlparse tool renders as an HTML report with every graph, guard, recompile and the generated code.

The break above has a standard fix: keep the decision on the device.

def f(x):
    y = torch.sin(x) * 2
    # Both branches are computed on the GPU and selected elementwise: no host sync, one graph.
    return torch.where(y.sum() > 0, y.cos(), y.tanh())

When compilation crashes or changes results, bisect the stack with the backend argument. backend="eager" runs Dynamo capture only and executes the captured graph eagerly; backend="aot_eager" adds AOTAutograd; the default adds Inductor. Whichever step first reproduces the problem owns it. Fusion reorders floating-point operations, so compare against eager with a dtype-appropriate tolerance, not exact equality. To profile, capture a trace as in GPU profiling: gaps between kernels point at graph breaks or host synchronisation.

Failure modes

  • Silent eager fallback. The recompile limit is hit, or so many breaks fragment the graph that nothing fuses, and the model runs at eager speed. Detect it with TORCH_LOGS=recompiles and fail_on_recompile in tests.
  • Checkpoint key mismatch. State dicts saved from a torch.compile wrapper carry an _orig_mod. prefix. Save from the original module or use module.compile().
  • CUDA-graph output overwritten. In reduce-overhead mode, a tensor returned by one step is overwritten by the next replay. Clone outputs you keep.
  • Compile storms at startup. Every replica compiling on boot delays readiness. Ship cache artefacts and warm up before accepting traffic.
  • Distributed desynchronisation. If ranks see different shapes and one recompiles while others do not, collectives can stall. Keep shapes identical across ranks and warm up on all of them.

When not to use it

Compilation pays off for code that runs many times with stable shapes. It is a poor fit for one-off scripts where compile time exceeds runtime, for models dominated by large matmuls that vendor libraries already run near peak, and for control flow driven by tensor values at every step.

What to do next

  1. Measure eager step time and peak memory, then wrap the model with torch.compile and measure again, excluding the first few steps.
  2. Run torch._dynamo.explain on one step and remove the breaks it reports, then set fullgraph=True to keep them out.
  3. Run a few hundred steps with TORCH_LOGS=recompiles, and fix any recompilation by bucketing shapes, passing scalars as tensors or marking dimensions dynamic.
  4. For deep stacks of identical blocks, try regional compilation and compare cold-start time against whole-model compilation.
  5. Try reduce-overhead for small batches and max-autotune for long jobs, and keep whichever wins on your own measurements.
  6. For serving, warm up every shape bucket, save Mega-Cache artefacts keyed on PyTorch, Triton and GPU, and gate deployments with fail_on_recompile.
Key takeaway: torch.compile captures Python into graphs with Dynamo, adds a compiled backward with AOTAutograd and generates fused kernels with Inductor, so memory-bound chains of operations stop round-tripping through GPU memory. Its results depend on two events you can observe: graph breaks, which cut graphs and force synchronisation, and recompilations, which silently fall back to eager after the cache limit. Remove breaks, stabilise shapes, pick the mode by measurement, cache compiled artefacts for fast starts, and use set_stance to make regressions fail loudly.