Course Content
Model Deployment for AI Engineers
4 sections · 10 lessons
Model Export Formats — ONNX, TorchScript, and TensorFlow SavedModel
A team trains a defect classifier on a factory line. It hits 96.4% accuracy in the notebook. They hand it to the platform team as model.pt — a file produced by torch.save(model, "model.pt"). The platform team runs torch.load("model.pt") and gets:
AttributeError: Can't get attribute 'DefectNet' on <module '__main__'>The file contains a pickled reference to a Python class called DefectNet that lives in a notebook cell on someone's laptop. Without that class definition — the exact source file, the exact attribute names, a compatible PyTorch version — the file is 240 MB of unusable bytes. The weights are in there. The model is not.
That gap is what export formats exist to close. Training produces a Python object entangled with your source tree; serving needs a self-contained artefact that a C++ runtime on an ARM chip can load six months from now without importing anything you wrote. This lesson is about how to produce that artefact, in three formats, and how to prove it still computes the same thing.
What a trained model actually is
Split it into three parts, because export formats treat them very differently.
| Part | What it is | Where it lives in PyTorch |
|---|---|---|
| Weights | The learned numbers — matrices, bias vectors, batch-norm running statistics | model.state_dict(), a dict of tensors |
| Architecture | Which operations run, in what order, with what shapes | The forward() method — arbitrary Python |
| Pre/post-processing | Resize, normalise, tokenise, argmax, label lookup | Usually scattered outside the model entirely |
state_dict() gives you the weights cleanly, and saving it is genuinely portable — it is just tensors. But it does not tell anyone what to do with them. torch.save(model) tries to capture the architecture too, by pickling, which is why it drags your source tree along as a hidden dependency.
An export format's real job is to serialise the architecture as data rather than as code, so that loading a model never requires running your code.
Every format on this page does that same trick: it converts forward() from Python control flow into a computation graph — a list of nodes, each naming an operation, its inputs, its outputs, and its constants. A graph is data. Data survives the loss of your repository.
ONNX: one graph, many runtimes
ONNX (Open Neural Network Exchange) is a file format plus a fixed vocabulary of operators. An ONNX file is a protobuf containing a directed acyclic graph. Each node says something like "run Conv with these attributes on inputs x and W, produce y". The weights ride along as initialiser tensors inside the same file.
Why the vocabulary matters more than the file format
The critical piece is the opset version — the numbered snapshot of what operators exist and what each one means. Opset 11 defines Resize one way; opset 13 refines it; opset 17 adds LayerNormalization as a single fused operator instead of a dozen primitive nodes. When you export you pick an opset, and that number is a contract between you and whatever runtime loads the file.
Get it wrong in either direction and you get a specific, recognisable failure:
| Choice | Symptom | Why |
|---|---|---|
| Opset too high for the runtime | Unsupported model IR version or No opset import for domain '' at load time | The runtime's operator table stops below your number |
| Opset too low for your model | Export itself fails: Exporting the operator X is not supported | The op you used has no equivalent in that vocabulary yet |
| Opset matched to runtime | Loads and runs | Both sides agree on operator semantics |
The practical rule: check the ONNX Runtime version on your serving machine, look up the highest opset it supports, and export at or a little below that. Pin both versions in requirements.txt. A model that exported fine on your laptop and fails in the container is nearly always this. One floor applies too: the current PyTorch exporter emits opset 18 or newer. Ask it for 17 and it exports at 18, then tries to convert the file down, which can fail and leave you with an opset-18 file anyway.
Exporting from PyTorch
1import torch2import torch.nn as nn34class DefectNet(nn.Module):5 def __init__(self, n_classes=4):6 super().__init__()7 self.features = nn.Sequential(8 nn.Conv2d(3, 32, 3, padding=1), nn.ReLU(), nn.MaxPool2d(2),9 nn.Conv2d(32, 64, 3, padding=1), nn.ReLU(), nn.AdaptiveAvgPool2d(1),10 )11 self.head = nn.Linear(64, n_classes)1213 def forward(self, x):14 z = self.features(x).flatten(1)15 return self.head(z)1617model = DefectNet()18model.load_state_dict(torch.load("weights.pt", map_location="cpu"))19model.eval() # non-negotiable, see below2021dummy = torch.randn(2, 3, 224, 224) # batch of 2: a batch of 1 is treated as a constant22batch = torch.export.Dim("batch")2324torch.onnx.export(25 model, (dummy,), "defectnet.onnx",26 input_names=["pixels"],27 output_names=["logits"],28 dynamic_shapes={"x": {0: batch}}, # "x" is forward()'s argument name29 opset_version=18,30 dynamo=True, # the default since PyTorch 2.931)This is the torch.export-based exporter, the default since PyTorch 2.9 (it needs PyTorch 2.5 or later and the onnxscript package). Older tutorials pass dynamic_axes= and do_constant_folding= to the legacy TorchScript-based exporter; those arguments still work for now, with a deprecation warning, and dynamo=False brings the old exporter back if an old model will not export any other way.
Three things about that call deserve real explanation.
Why a dummy input at all
The exporter does not read your source code as source. It traces: it feeds the dummy tensor through forward() and records every tensor operation that actually executes. Whatever the trace touches becomes the graph. Whatever it skips does not exist in the exported model. That is true of the legacy exporter, of torch.jit.trace, and — for ordinary Python branches — of the current torch.export-based exporter too.
This is the single most important thing to understand about ONNX and TorchScript tracing, and it produces a failure mode that is completely silent. Consider:
1def forward(self, x):2 if x.shape[2] > 256: # recorded as a fixed decision, not a branch3 x = self.downsample(x)4 return self.head(x)Trace with a 224×224 image and the condition is False, so downsample is never recorded. The exported graph unconditionally skips it. Feed the deployed model a 512×512 image and it silently produces garbage — no error, just wrong predictions. Under the legacy tracer, Python-level if, for over a variable count, and .item() calls all collapse into constants. The current exporter is stricter in one respect: a branch on a tensor's values (the .item() case) makes the export fail loudly instead of baking in one answer. A branch on a shape is still frozen silently — export the snippet above with height and width marked dynamic and the graph still contains no downsample at all.
Tracing records one execution path and calls it the model; any branch your dummy input did not take has been deleted, not preserved.
Why model.eval() before exporting
Dropout and batch normalisation behave differently in training and evaluation mode, and tracing captures whichever behaviour was active. Export in training mode and you bake in dropout — the deployed model will randomly zero 20% of its activations on every request, and accuracy drops by a few points for no visible reason. Batch norm is worse: in training mode it normalises using the statistics of the current batch, so a single-image request normalises the image against itself, which destroys the signal entirely.
Why dynamic axes
Without dynamic_shapes, the traced shapes are frozen. Your dummy was (2, 3, 224, 224), so the graph demands exactly batch size 2 and errors on anything else. Marking axis 0 with a Dim tells the exporter to leave that dimension symbolic; the output's batch axis follows automatically. (Use a dummy batch of at least 2: torch.export treats a dimension of size 1 as a constant and will refuse to make it dynamic.) If you also serve variable-resolution images, mark axes 2 and 3 too:
1dynamic_shapes = {"x": {2 0: torch.export.Dim("batch"),3 2: torch.export.Dim("height"),4 3: torch.export.Dim("width"),5}}There is a cost. Fixed shapes let the runtime pre-compute memory layouts and pick specialised kernels; fully dynamic shapes typically cost 5–15% throughput. Make dynamic only what genuinely varies. Batch almost always varies. Image size often does not, because your preprocessing resizes anyway.
Verifying the export
Never trust an export you have not numerically checked. The check is cheap and it catches nearly everything.
1import numpy as np, onnx, onnxruntime as ort23onnx.checker.check_model(onnx.load("defectnet.onnx")) # structural validity only45sess = ort.InferenceSession("defectnet.onnx", providers=["CPUExecutionProvider"])67x = torch.randn(5, 3, 224, 224) # batch 5, not 2: tests dynamic axis8with torch.no_grad():9 torch_out = model(x).numpy()10onnx_out = sess.run(["logits"], {"pixels": x.numpy()})[0]1112print("max abs diff:", np.abs(torch_out - onnx_out).max())13np.testing.assert_allclose(torch_out, onnx_out, rtol=1e-3, atol=1e-5)What counts as "the same"? Floating-point arithmetic is not associative, and the runtimes fuse operations differently, so exact equality is the wrong bar. Realistic tolerances:
| Observed max absolute difference | Verdict |
|---|---|
| Up to about 1e-5 on FP32 logits | Normal. Kernel fusion and summation order. |
| 1e-5 to 1e-3 | Suspicious. Often an accumulated difference in a normalisation layer. Check predicted labels still match on a few hundred real samples. |
| Above 1e-2, or labels flipping | A real bug. Usually forgotten eval(), a traced-away branch, or a preprocessing mismatch. |
Note that onnx.checker.check_model only validates structure — that the protobuf is well-formed and the nodes reference real operators. It will happily pass a model whose predictions are wrong. The numerical comparison is the test that matters. Run it on a batch size your dummy did not use, so the dynamic axis gets exercised too.
TorchScript: staying inside PyTorch, leaving Python
TorchScript serialises a model into a form the PyTorch C++ runtime (libtorch) can execute with no Python interpreter present. You stay in the PyTorch ecosystem, but you can now load the model from a C++ service, or — historically — an iOS or Android app.
It offers two conversion routes, and choosing between them is the whole skill.
Tracing
1traced = torch.jit.trace(model, dummy)2traced.save("defectnet_traced.pt")34loaded = torch.jit.load("defectnet_traced.pt") # no DefectNet class needed anywhereSame mechanism as ONNX export, same limitation: one recorded path, control flow flattened. Fast, works on essentially any model built from standard layers.
Scripting
scripted = torch.jit.script(model)scripted.save("defectnet_scripted.pt")Scripting actually compiles your forward(). It parses the Python source into a typed intermediate representation, preserving if, for, and while as genuine control flow. The price is that TorchScript is a strict, statically typed subset of Python. It will reject things ordinary Python allows.
1class Decoder(nn.Module):2 def forward(self, x, max_steps: int):3 out = [] # must be annotated for scripting4 for _ in range(max_steps): # a real loop, preserved5 x = self.step(x)6 out.append(x)7 return torch.stack(out)To script that, you need out: List[torch.Tensor] = [], imported from typing. Unannotated empty lists default to List[Tensor] in recent versions but the error messages when it guesses wrong are cryptic, so annotate. Other common rejections: dictionaries with mixed value types, *args, calling into third-party libraries like NumPy or PIL, and returning different types from different branches.
You can mix the two. Script the outer module that owns the control flow, and let it call traced submodules — torch.jit.script will keep an already-traced child as-is.
torch.export: the replacement
torch.export captures the same kind of whole-model graph, as an ExportedProgram you can save and reload without the class definition:
1ep = torch.export.export(model, (dummy,),2 dynamic_shapes={"x": {0: torch.export.Dim("batch")}})3torch.export.save(ep, "defectnet.pt2")45loaded = torch.export.load("defectnet.pt2").module() # no DefectNet class neededIt behaves like tracing rather than scripting: a Python if on a shape is fixed at export time, and a branch on tensor values fails the export unless you rewrite it with torch.cond, which is recorded as real control flow. The .pt2 file is the starting point for the Python-free runtimes that replace TorchScript: AOTInductor compiles it into a shared library a C++ server can load, and ExecuTorch runs it on phones and embedded devices. It is also what the ONNX exporter above uses internally.
TensorFlow SavedModel: a directory with signatures
SavedModel is TensorFlow's native format, and it is a directory rather than a single file:
saved_model/├── saved_model.pb # graph structure + signature definitions├── variables/│ ├── variables.data-00000-of-00001 # the weights│ └── variables.index└── assets/ # vocabulary files, label maps, anything elseIts distinguishing feature is signatures: named entry points with declared input and output tensor specs. One artefact can expose several callable functions, which is how you ship preprocessing inside the model.
1import tensorflow as tf23model = tf.keras.models.load_model("defect_keras.keras")45@tf.function(input_signature=[tf.TensorSpec([None, 224, 224, 3], tf.float32, name="pixels")])6def serve_logits(pixels):7 return {"logits": model(pixels, training=False)}89@tf.function(input_signature=[tf.TensorSpec([None], tf.string, name="jpeg_bytes")])10def serve_from_jpeg(jpeg_bytes):11 def decode(b):12 img = tf.io.decode_jpeg(b, channels=3)13 img = tf.image.resize(img, [224, 224])14 return tf.cast(img, tf.float32) / 255.015 batch = tf.map_fn(decode, jpeg_bytes, fn_output_signature=tf.float32)16 probs = tf.nn.softmax(model(batch, training=False))17 return {"probabilities": probs, "predicted": tf.argmax(probs, axis=1)}1819tf.saved_model.save(model, "saved_model", signatures={20 "serving_default": serve_logits,21 "from_jpeg": serve_from_jpeg,22})The from_jpeg signature is the interesting one. Resizing and the divide-by-255 now live inside the graph. That kills an entire class of production bug: the client that resizes with a different interpolation mode, or normalises with ImageNet means when training used plain 0–1 scaling. Accuracy quietly falls by several points and nobody can find the cause, because the model file is identical in both cases.
Every preprocessing step you leave outside the exported artefact is a step someone will eventually implement differently in the serving path.
Inspect what you actually shipped from the command line — worth doing before every deploy:
saved_model_cli show --dir saved_model --tag_set serve --signature_def serving_defaultThe None in [None, 224, 224, 3] is TensorFlow's equivalent of a dynamic axis: batch is variable, spatial dimensions are fixed.
Choosing between the three
| ONNX | TorchScript | SavedModel | |
|---|---|---|---|
| Source framework | PyTorch, TF, sklearn, XGBoost | PyTorch only | TensorFlow/Keras only |
| Artefact | Single .onnx file | Single .pt file | Directory |
| Runtimes | ONNX Runtime, TensorRT, OpenVINO, CoreML, browsers via WASM | libtorch (C++); deprecated in favour of torch.export | TF Serving, LiteRT (formerly TFLite), TF.js |
| Preserves control flow | Partly (via torch.cond with the current exporter) | Yes, with scripting | Yes, in tf.function |
| Preprocessing in artefact | Only as ONNX ops | Yes, if TorchScript-compatible | Yes, full TF ops including JPEG decode |
| Typical CPU speedup vs. eager Python | 1.5–3× | 1.1–1.5× | 1.2–2× |
| Main risk | Unsupported op, opset mismatch | Scripting rejects your Python | Lock-in to TF serving stack |
Reasonable defaults: PyTorch model going to a CPU or GPU microservice, or to specialised hardware — ONNX. PyTorch model going into a C++ or mobile application — torch.export, then AOTInductor for a C++ server or ExecuTorch for a device; keep TorchScript for maintaining deployments that already use it. TensorFlow model of any kind — SavedModel, because converting TF to ONNX adds a failure surface for no benefit.
Export failures you will actually hit
Unsupported operator
torch.onnx._internal.exporter._errors.DispatchError:No ONNX function found for <OpOverload(op='mylib.my_custom_op', overload='default')>Three fixes, in order of how often they work. First, upgrade PyTorch and onnxscript and raise the opset — coverage improves with every release and this resolves maybe half of these. Second, rewrite the layer using supported primitives; a custom activation can nearly always be expressed with existing ops. Third, give the exporter a translation that spells your op in standard ONNX operators:
1from onnxscript import opset18 as op23def my_op_onnx(x, alpha: float):4 # express the op with standard ONNX nodes5 return op.Mul(x, op.Sigmoid(op.Mul(x, alpha)))67torch.onnx.export(8 model, (dummy,), "defectnet.onnx",9 custom_translation_table={torch.ops.mylib.my_custom_op.default: my_op_onnx},10)The legacy exporter used register_custom_op_symbolic for the same job; it is deprecated and has no effect on the current exporter.
Data-dependent control flow
Beam search, early-exit networks, and anything whose loop count depends on tensor values cannot be traced correctly. A branch on tensor values can be rewritten with torch.cond, which the current exporter turns into an ONNX If node. For a data-dependent loop, the usual answer is to move the loop out of the model and into your serving code, exporting only the single-step function. That is simpler and easier to debug than any in-graph loop.
Quantising before exporting
PyTorch's quantised modules use a different set of kernels, and export support for them lags. The reliable order is: export the FP32 model to ONNX first, then quantise the ONNX file with ONNX Runtime's own tooling.
1from onnxruntime.quantization import quantize_dynamic, QuantType2from onnxruntime.quantization.shape_inference import quant_pre_process34quant_pre_process("defectnet.onnx", "defectnet_prep.onnx") # shape inference + cleanup5quantize_dynamic("defectnet_prep.onnx", "defectnet_int8.onnx", weight_type=QuantType.QInt8)Do not skip the pre-processing step: on files from the current exporter, quantising directly can fail with a shape-inference error.
On a 240 MB FP32 model this typically lands around 62 MB — the weights go from four bytes to one, with the graph structure and a few unquantised layers accounting for the rest.
The version-drift failure
The exported file is fine; the environment moved. Someone moves onnxruntime in an unpinned Dockerfile — down to an older build, or up past a change in behaviour — and a model that loaded yesterday stops loading. Pin the exporting framework version, the opset, and the runtime version together, and record all three next to the artefact:
1{2 "artefact": "defectnet.onnx",3 "sha256": "9f2c1a...",4 "exported_from": "torch==2.14.0",5 "opset_version": 18,6 "verified_against_runtime": "onnxruntime==1.30.0",7 "max_abs_diff_vs_pytorch": 3.1e-06,8 "input_signature": {"pixels": ["batch", 3, 224, 224]}9}What this changes about how you finish a model
Treat export as part of training, not as a handover step afterwards. Concretely, the last cell of a training run should export the artefact, run the numerical comparison against the eager model on a held-out batch, and fail loudly if the difference exceeds your tolerance. That way a model that cannot be exported never gets recorded as a successful run in the first place.
Two habits follow from that. Keep a small, fixed set of golden inputs — a dozen real samples with their expected outputs, stored as a .npz — and check every artefact against them, at every stage: after export, after quantisation, after the container build. When a prediction changes, you learn which stage changed it instead of bisecting a pipeline.
And decide deliberately where the preprocessing boundary sits, then document it in the same file as the checksum. Whether resize-and-normalise lives inside the artefact or in the serving code is a legitimate choice with real trade-offs. Leaving it undecided is not — that is how the same model scores 96.4% offline and 91% in production, with nothing in the logs to explain the gap.