Deciding When to Use torch.compile in PyTorch Training Loops
Learn how to decide when to wrap a PyTorch model with torch.compile, see a concrete training‑loop example, and understand the trade‑offs like compilation overhead, dynamic shapes, and graph breaks.
17 Jun 2026, 23:05 UTC

The problem: slow training loops
When you run a training loop on a GPU, each iteration spends time launching kernels, executing Python bytecode, and synchronizing streams. For models that run many iterations (e.g., full‑epoch training or repeated inference), these overheads can become a bottleneck even though the model architecture itself is unchanged.
How torch.compile works under the hood
torch.compile (available since PyTorch 2.0) wraps a model or a Python function and runs it through TorchDynamo, which traces the Python bytecode into an FX graph. The graph is then lowered by TorchInductor to optimized kernels, optionally using CUDA graphs in the reduce-overhead mode. The result is a callable that behaves like the original Python code but executes the compiled graph on subsequent calls.
Worked example: compiling a transformer training step
Below is a minimal, self‑contained snippet that shows how to apply torch.compile to a training step without changing the optimizer or loss code. The example assumes you have a recent PyTorch 2.x install and a CUDA‑capable GPU.
import torch
import torch.nn as nn
from torch.utils.data import DataLoader, TensorDataset
# Dummy model and data
model = nn.Transformer(d_model=512, nhead=8, num_encoder_layers=6)
model = model.cuda()
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3)
loss_fn = nn.CrossEntropyLoss()
# Synthetic batch: (seq_len, batch_size, vocab_size)
data = torch.randint(0, 1000, (32, 16)).cuda()
targets = torch.randint(0, 1000, (32, 16)).cuda()
dataset = TensorDataset(data, targets)
loader = DataLoader(dataset, batch_size=16)
# Wrap the model with torch.compile
# Choose a mode that matches your workload:
# - default: balances compile time and speedup
# - reduce-overhead: better for small models, uses CUDA graphs
# - max-autotune: longer compile, searches for optimal kernels
compiled_model = torch.compile(model, mode="reduce-overhead")
# Training loop
for epoch in range(3):
for xb, yb in loader:
optimizer.zero_grad()
# Forward pass through the compiled model
logits = compiled_model(xb) # shape: (seq_len, batch, vocab)
loss = loss_fn(logits.view(-1, logits.size(-1)), yb.view(-1))
loss.backward()
optimizer.step()
To see what TorchDynamo is doing, enable logging before the loop:
import os
os.environ["TORCH_LOGS"] = "graph_breaks,recompiles"
# re‑run the loop; messages will appear on stdout
The log lines tell you whether the tracer had to break the graph (falling back to eager execution) or if it recompiled because input shapes changed. This diagnostic helps you decide whether to use dynamic=True or to redesign data‑dependent control flow.
Trade‑offs and practical checks
- Compilation overhead: The first few iterations are slower because TorchDynamo traces and TorchInductor generates kernels. The amortized benefit appears after a warm‑up period, so torch.compile pays off in long‑running jobs.
- Dynamic shapes: If your batch size or sequence length varies per iteration, Dynamo will trigger a recompilation each time a new shape is seen. Setting
torch.compile(model, dynamic=True)lets the generated graph handle a range of shapes, but it may still recompile when the shape falls outside the guarded range. - Graph breaks: Certain Python constructs (e.g., data‑dependent loops, native I/O, or calls to unsupported C extensions) cause Dynamo to break the graph. The model remains correct, but you lose the speedup for the broken portions. Refactoring such code or using
torch._dynamo.config.suppress_errors = True(for debugging) can help identify the source. - Numerical differences: Kernel fusion and altered reduction order can produce results that differ slightly from eager mode. Verify that your validation metrics stay within an acceptable tolerance (e.g.,
1e‑4) after enabling compile. - Compatibility with other features: torch.compose works with
torch.autocastfor mixed precision and with distributed wrappers like DDP or FSDP, but you should check the release notes for your PyTorch version because some combinations have version‑specific caveats.
When to enable it and next steps
If you are training a model for many epochs or running inference repeatedly on the same architecture, try the following steps:
- Wrap your model (or the training step function) with
torch.compile, starting with thereduce-overheadmode for small models or the default mode otherwise. - Run a short benchmark (e.g., 100–200 iterations) and compare the average step time against the eager baseline, excluding the first few warm‑up iterations.
- Check
TORCH_LOGSoutput for excessive recompilations or graph breaks; adjustdynamic=Trueor refactor problematic Python code if needed. - Validate that evaluation metrics remain within tolerance.
- If the speedup is satisfactory, keep the compiled wrapper in your training script; otherwise, revert to the eager model.
By treating compilation as an engineering decision—measuring warm‑up cost, monitoring recompilations, and verifying numerical stability—you can adopt torch.compile where it genuinely reduces runtime without introducing hidden regressions.
0 replies
A thoughtful contribution can make all the difference. Be the first to share one.