Checkpoints, Export and Inference
torch.export and the ExportedProgram
torch.export captures your model as one flat graph of low-level operators with the weights lifted out as inputs, and it refuses to guess where tracing would have quietly guessed wrong.
- 11 min read
- 3 reading levels
- Published
Read these first
On this page 6
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
torch.export turns your model into a single, complete list of low-level steps — and stops with an error rather than guessing.
Think of translating a recipe for a factory. The home version says "fry the onions until golden". The factory version has to say "heat oil to 180 degrees, add onions, stir for 240 seconds". Every vague instruction becomes a precise one, and anything genuinely ambiguous has to be settled before the line can run.
An ExportedProgram is that factory version of your model. One flat list of exact operations, with the weights handed in as ingredients rather than hidden inside the recipe.
Why it replaced the older tools
The previous lesson showed tracing quietly recording one branch of a decision and getting later answers wrong. That silence is the problem the new tool was built to fix.
torch.export handles the same situation by refusing. It stops, names the line of your code, and tells you the branch cannot be captured. An error you can read beats a wrong answer you cannot see.
The other thing it fixes
Tools that record a run also record the sizes they saw. Film a model on batches of 3 and it may only ever accept 3.
With torch.export you say out loud which sizes may vary. For example: the batch may be anything from 1 to 1024, and everything else is fixed. The graph then carries that promise, and the runtime can check it.
How it works
your model (Python)
|
| torch.export.export(model, example_inputs, dynamic_shapes=...)
v
ExportedProgram:
graph: linear -> relu (low-level operators, one flat list)
inputs: weights, biases, x (the weights are handed IN)
promise: x may have any batch from 1 to 1024A real example you have seen
A visa form. Free-text answers are fine for a human reader. The official form has one box per fact and rejects anything ambiguous, because the machine reading it afterwards cannot ask you a follow-up question.
Remember this
- An ExportedProgram is one flat graph of low-level operations, no Python left.
- The weights become inputs to the graph, not hidden state inside it.
- It refuses data-dependent branches instead of silently picking one.
What to learn next
- TorchScript: tracing vs scripting — the older stack this replaces, and why.
- ONNX — the interchange format the exported graph feeds.
- torch.compile — the same capture machinery aimed at speed rather than portability.
Developer — Code and libraries.
Setup
pip install torchRuns on CPU. torch.export is the capture step for the whole modern deployment stack — ExecuTorch for phones, AOTInductor for compiled C++ binaries, and the current ONNX exporter all consume an ExportedProgram.
Exporting, and reading what came out
import torch
import torch.nn as nn
from torch.export import export, Dim
class Net(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(4, 2)
def forward(self, x):
return torch.relu(self.fc(x))
model = Net().eval()
example = (torch.randn(3, 4),)
batch = Dim("batch", min=1, max=1024) # say out loud what may vary
ep = export(model, example, dynamic_shapes={"x": {0: batch}})
print(ep.graph_module.code)
print("signature inputs :", [s.arg.name for s in ep.graph_signature.input_specs])
print("run with batch 3 :", tuple(ep.module()(torch.randn(3, 4)).shape))
print("run with batch 9 :", tuple(ep.module()(torch.randn(9, 4)).shape))def forward(self, p_fc_weight, p_fc_bias, x):
linear = torch.ops.aten.linear.default(x, p_fc_weight, p_fc_bias); x = p_fc_weight = p_fc_bias = None
relu = torch.ops.aten.relu.default(linear); linear = None
return (relu,)
signature inputs : ['p_fc_weight', 'p_fc_bias', 'x']
run with batch 3 : (3, 2)
run with batch 9 : (9, 2)Three things in that graph are worth sitting with.
torch.ops.aten.linear.default — not nn.Linear. Your module hierarchy is gone, flattened into ATen, the low-level operator set every PyTorch backend already implements. A runtime that supports ATen supports your model, with no knowledge of your classes.
forward(self, p_fc_weight, p_fc_bias, x) — the weights arrive as arguments. This is parameter lifting, and it is why the graph is a pure function: same inputs, same outputs, no hidden state. ep.graph_signature records which inputs are parameters, which are buffers and which are real user inputs, so ep.module() can hand the right tensors back in for you.
Batch 3 and batch 9 both work. The example had 3 rows; the Dim said the first dimension may vary. Without that argument, export would have specialised the graph to exactly 3.
Saving it as a file
import torch
import torch.nn as nn
from torch.export import export, save, Dim
class Net(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(4, 2)
def forward(self, x):
return torch.relu(self.fc(x))
ep = export(Net().eval(), (torch.randn(3, 4),),
dynamic_shapes={"x": {0: Dim("batch", min=1, max=1024)}})
save(ep, "net.pt2")
print("saved net.pt2")import torch # the Net class is NOT defined in this file
from torch.export import load
ep = load("net.pt2")
print(tuple(ep.module()(torch.randn(7, 4)).shape))saved net.pt2 (7, 2)
The .pt2 archive holds the graph, the weights and the shape constraints. The consumer needs no model class — the same portability promise TorchScript made, on a foundation that is still being developed.
The branch that tracing got wrong
import torch
import torch.nn as nn
from torch.export import export, Dim
class Clamped(nn.Module):
def forward(self, x):
if x.sum() > 0:
return x * 2
return x * -1
ep = export(Clamped(), (torch.ones(3),), dynamic_shapes={"x": {0: Dim("n", min=1, max=64)}})
print(ep.graph_module.code)torch._dynamo.exc.UserError: Dynamic control flow is not supported at the moment.
Please use functorch.experimental.control_flow.cond to explicitly capture the control
flow. For more information about this error, see:
https://pytorch.org/docs/main/generated/exportdb/index.html#cond-operands
from user code:
File "export_branch.py", line 7, in forward
if x.sum() > 0:Compare this against the same model under torch.jit.trace in the previous lesson, which produced -2.0 where the answer was 1.0 and shipped happily. Here you get a refusal, the exact line, and a link to the fix.
The fix is to state the branch as data instead of as Python:
import torch
import torch.nn as nn
from torch.export import export
class Clamped(nn.Module):
def forward(self, x):
return torch.cond(x.sum() > 0, lambda t: t * 2, lambda t: t * -1, (x,))
model = Clamped()
ep = export(model, (torch.ones(3),))
print("eager, negative :", model(-torch.ones(3)).tolist())
print("exported, negative:", ep.module()(-torch.ones(3)).tolist())
print("exported, positive:", ep.module()(torch.ones(3)).tolist())eager, negative : [1.0, 1.0, 1.0] exported, negative: [1.0, 1.0, 1.0] exported, positive: [2.0, 2.0, 2.0]
torch.cond puts both branches in the graph, with the predicate as a runtime value. The negative input flips the sign, giving 1.0; the positive input doubles, giving 2.0; and the exported program matches eager on both. That is exactly what the traced model could not do.
Choosing your export tool
torch.jit.trace | torch.jit.script | torch.export | |
|---|---|---|---|
| status | maintenance | maintenance | active development |
data-dependent if | silently baked in | preserved | refuses, or cond |
| dynamic shapes | implicit, unchecked | implicit | declared and checked |
| operator level | TorchScript IR | TorchScript IR | ATen |
| weights | inside the module | inside the module | lifted to inputs |
| feeds | LibTorch C++ | LibTorch C++ | ExecuTorch, AOTInductor, ONNX |
Common mistakes
Forgetting dynamic_shapes. The graph is specialised to your example's exact sizes and rejects everything else at runtime.
Exporting in training mode. Call .eval() first, or dropout and BatchNorm's training behaviour go into the graph.
Confusing torch.export with torch.compile. Compile makes this process faster and falls back to Python whenever it needs to. Export produces a portable artefact and is not allowed to fall back — which is why it errors where compile would shrug. See torch.compile.
Data-dependent shapes. x[x > 0] produces a size nobody can know ahead of time and raises a GuardOnDataDependentSymNode error. Restructure to a fixed shape with a mask, or move the filtering outside the model.
Treating the export error as a bug. It is the feature. Every refusal is a case where a recording-based tool would have produced a silently wrong artefact.
Try it yourself
Export Net without dynamic_shapes, then call it with a batch of 9. Read the error carefully — it names the dimension it specialised and suggests the Dim you should have declared.
What to learn next
- TorchScript: tracing vs scripting — the older stack this replaces, and why.
- ONNX — the interchange format the exported graph feeds.
- torch.compile — the same capture machinery aimed at speed rather than portability.
Researcher — Mathematics and papers.
Capture mechanism
torch.export uses TorchDynamo, which analyses CPython bytecode rather than tracing operator dispatch or parsing source. Dynamo runs the function symbolically, building an FX graph of tensor operations and accumulating guards — predicates on input properties that must hold for the graph to be valid. Under torch.compile a failed guard triggers recompilation or a graph break with a fallback to eager Python. Under torch.export neither is permitted: the result must be a single whole-program graph, so anything that would break the graph becomes a UserError. That single design constraint explains every difference in the table above.
The graph is then functionalised — mutations and aliasing are removed so the result is a pure function — and decomposed to the Core ATen operator set, a stable subset of roughly 180 operators that backend authors are expected to implement. run_decompositions() controls how far the lowering goes.
An ExportedProgram holds: the GraphModule, a state_dict of parameters and buffers, a graph_signature classifying every graph input as parameter, buffer, constant or user input, range_constraints over the symbolic sizes, and the module_call_graph that records the original module hierarchy so it can be reconstructed for unflattening.
Symbolic shapes
Dynamic dimensions become SymInt values backed by SymPy expressions. Declaring Dim("batch", min=1, max=1024) introduces a symbol $s_0$ with those bounds; every downstream shape becomes an expression in $s_0$, and every branch on a shape adds a guard. Guards that are implied by the range constraints are discharged; guards that are not become either an added constraint or an error naming the specialisation. This is why export sometimes reports that a dimension you declared dynamic was specialised: some operation in the model compared it against a constant.
Derived dimensions are expressible — Dim("seq") * 2, or 2 * batch — which matters for models where two inputs must agree on a size. Data-dependent sizes, where the value comes from tensor contents rather than shapes (nonzero, boolean masking, unique), produce unbacked symbols with no known range, and any branch on them raises GuardOnDataDependentSymNode unless bounded through torch._check.
Higher-order operators
Control flow that must survive is expressed as higher-order operators whose arguments are subgraphs: cond for branches, while_loop for loops, scan and map for structured iteration, and associative_scan for parallel prefix. cond requires both branches to be traceable, to accept the same operands, and to return the same shapes and dtypes — the last requirement is why cond cannot express a branch that returns different-sized outputs, and why some models need restructuring rather than a mechanical rewrite.
What consumes the result
The .pt2 archive is the interchange point for three distinct backends. AOTInductor compiles an ExportedProgram to a shared library callable from C++ with no Python runtime. ExecuTorch lowers it further to a compact bytecode for mobile and embedded targets, with delegate partitions for vendor accelerators. The dynamo-based ONNX exporter consumes it as well, replacing the older tracing-based path. Standardising capture and letting backends differ afterwards is the architectural bet of PyTorch 2: one front end, many deployment targets, rather than a bespoke tracer per target.
References
- Ansel et al. (2024), PyTorch 2: Faster Machine Learning Through Dynamic Python Bytecode Transformation and Graph Compilation, ASPLOS — Dynamo, guards, and the graph-capture design.
- PyTorch documentation, torch.export —
ExportedProgram,dynamic_shapes,Dim, and the error taxonomy. - PyTorch documentation, ExportDB — a catalogue of supported and unsupported Python patterns with the exact error each produces.
- PyTorch documentation, Core ATen IR — the operator set backends are expected to implement.
What to learn next
- TorchScript: tracing vs scripting — the older stack this replaces, and why.
- ONNX — the interchange format the exported graph feeds.
- torch.compile — the same capture machinery aimed at speed rather than portability.