GPU Memory and Speed

torch.compile

torch.compile watches your Python model once, fuses its operations into optimised kernels, and then runs the fused version — one line for a real speedup, paid for by a slow first call.

On this page 5
  1. Why it exists
  2. How it works
  3. A real example you have seen
  4. Remember this
  5. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

torch.compile reads your model once, rewrites it into a faster fused version, and runs that version from then on.

Think of a new cook following a recipe. The first evening, he reads every line aloud, walks to the shelf for each spice separately, and the dish takes ages. After that night he knows the whole dish. He grabs all the spices in one trip and cooks without pausing to read.

Normal PyTorch is the first evening, every evening — each operation is a separate trip. Compiling is learning the dish.

Why it exists

PyTorch's charm is that it runs like ordinary Python, one line at a time. That is wonderful for experimenting and debugging. It is wasteful for speed.

Every small operation sends a separate work order to the GPU. Each order has fixed overhead, and each one reads its numbers from memory and writes them back. Ten small steps means ten round trips through memory for the same data.

Fusing — combining several steps into one — lets the data make one trip while all ten steps happen to it. That is where most of the speed comes from.

How it works

your Python model
      |  first call: watch what it does, record the recipe
      v
recorded graph of operations
      |  rewrite: fuse steps, generate fast kernels
      v
compiled version  -> every later call runs this

The first call is slow — sometimes very slow, tens of seconds — because the watching and rewriting happen then. Every call after that gets the fast version.

A real example you have seen

This is the same idea behind why a website loads slowly on first visit and instantly afterwards — expensive preparation once, cheap reuse forever. Training jobs run millions of identical steps, so one slow step to speed up all the rest is a wonderful trade.

Remember this

  • One line — wrap the model — and PyTorch rewrites it into fused kernels.
  • The first call pays for compilation; every later call collects the winnings.
  • It shines on long training runs; it is pointless for code that runs once.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch

Compilation targets Linux; it is standard on Colab and every cloud GPU box. Native Windows support is still limited — use WSL there. Outputs captured with torch 2.5.1 on an NVIDIA RTX A6000; timings vary a lot by machine, but the shape of the result — slow first call, fast ever after — is universal.

The one-liner and its price

compile_speed.py
import torch
import time

device = "cuda" if torch.cuda.is_available() else "cpu"

def gelu_chain(x):
    for _ in range(20):
        x = torch.nn.functional.gelu(x * 1.01 + 0.01)
    return x

fast = torch.compile(gelu_chain)
x = torch.randn(4096, 4096, device=device)

def timed(fn):
    if device == "cuda":
        torch.cuda.synchronize()
    start = time.perf_counter()
    fn(x)
    if device == "cuda":
        torch.cuda.synchronize()
    return time.perf_counter() - start

print(f"eager, warm:          {timed(gelu_chain)*1000:8.1f} ms")
print(f"compiled, first call: {timed(fast)*1000:8.1f} ms   <- includes compilation")
print(f"compiled, second:     {timed(fast)*1000:8.1f} ms")
print(f"compiled, third:      {timed(fast)*1000:8.1f} ms")
Output
eager, warm:              43.6 ms
compiled, first call:   2869.5 ms   <- includes compilation
compiled, second:          1.1 ms
compiled, third:           1.2 ms

The chain of 60 elementwise operations (multiply, add, gelu, twenty times) fused into a handful of kernels. Instead of 60 memory round trips over a 64 MB tensor, the data makes very few — a 40x win here. This example is deliberately fusion-friendly; a full real model typically lands between 1.2x and 2x. The synchronize() calls are why these timings are trustworthy — see timing GPU code correctly.

The same line works on a model: model = torch.compile(model). Training loops compile too — backward passes are generated from the same graphs.

When it cannot compile: graph breaks

The recorder follows your Python. Code whose path depends on tensor values cannot become one fixed recipe:

compile_break.py
import torch

@torch.compile(fullgraph=True)
def choosy(x):
    if x.sum() > 0:          # a decision that depends on the data itself
        return x * 2
    return x - 1

try:
    choosy(torch.ones(4))
except Exception as e:
    print(type(e).__name__)
    for line in str(e).splitlines()[:3]:
        print(line)
Output
Unsupported
Data-dependent branching
  Explanation: Detected data-dependent branching (e.g. `if my_tensor.sum() > 0:`). Dynamo does not support tracing dynamic control flow.
  Hint: This graph break is fundamental - it is unlikely that Dynamo will ever be able to trace through your code. Consider finding a workaround.

Without fullgraph=True there is no error: the compiler splits the function at the branch — a graph break — compiles the pieces, and runs the branch in plain Python. Correct, but each break costs speed. fullgraph=True turns silent breaks into loud errors, which is exactly what you want while optimising.

Common mistakes

Benchmarking the first call. The compile cost lands there. Warm up, then measure — and measure with synchronisation.

Changing shapes every call. A new input shape can trigger recompilation. A run that sees many different sequence lengths recompiles repeatedly; dynamic=True asks for shape-flexible kernels, or bucket your inputs to a few fixed shapes.

Prints and .item() inside the hot path. Each is a graph break — the recipe stops, Python runs, a new recipe starts. Move logging outside the compiled region, or accept the cost knowingly.

Compiling throwaway code. A script that runs a model a handful of times will never repay a multi-second compile. This tool is for loops that run thousands of times.

Try it yourself

Add print("hello") inside gelu_chain, remove fullgraph=True-style strictness by compiling plainly, and rerun. Then run with TORCH_LOGS=graph_breaks python compile_speed.py and watch the compiler confess where it split.

What to learn next

Researcher — Mathematics and papers.

The three-layer stack

torch.compile is TorchDynamo + AOTAutograd + TorchInductor:

  • TorchDynamo hooks CPython frame evaluation (PEP 523), symbolically evaluating bytecode to extract FX graphs of tensor operations, with guards — predicates on shapes, dtypes, globals — that decide when a cached graph is reusable. Guard failure triggers recompilation; unsupported constructs produce graph breaks and resume in the interpreter.
  • AOTAutograd traces joint forward+backward graphs ahead of time, so the backward is compiled with the same fusion opportunities, and handles functionalisation of mutations.
  • TorchInductor lowers graphs to generated code: Triton kernels on GPU, C++/OpenMP on CPU — performing fusion, tiling and scheduling decisions. mode="max-autotune" benchmarks candidate tilings per shape at compile time.

Why fusion is the win: a roofline argument

An elementwise op on $n$ float32 elements moves $8n$ bytes (read + write) to perform $n$ FLOPs — arithmetic intensity $\frac{1}{8}$ FLOP/byte, hopelessly memory-bound on hardware whose ridge point sits near 100 FLOP/byte. Fusing $k$ such ops multiplies intensity by $k$ without touching the compute. The demo above is the argument made flesh: 60 pointwise passes collapse to a few, and time drops ~40x. Matmul-heavy code, already compute-bound and served by cuBLAS, gains far less — hence the modest end-to-end speedups on real transformers relative to pointwise-heavy chains.

Shape polymorphism uses symbolic shapes (SymInt): dimensions marked dynamic become symbols, guards become inequalities, and one kernel serves a family of shapes at a small performance cost relative to shape-specialised kernels.

Relation to the export path

torch.compile is a JIT that tolerates graph breaks by design; torch.export reuses the same tracing machinery but demands a single whole graph and produces a serialisable artifact. Compile for speed in-process; export for deployment.

References

  • Ansel et al. (2024), PyTorch 2: Faster Machine Learning Through Dynamic Python Bytecode Transformation and Graph Compilation, ASPLOS — the paper on Dynamo and Inductor, with the guard mechanism and benchmark suite.
  • Tillet, Kung, Cox (2019), Triton: an intermediate language and compiler for tiled neural network computations — the kernel language Inductor emits.
  • Williams, Waterman, Patterson (2009), Roofline: an insightful visual performance model — the memory-bound argument above.

What to learn next