TorchONNX: A Compiler for ONNX-to-PyTorch Conversion
ONNX is already an explicit graph, so compiling it does not make it readable -- it makes it directly consumable by the PyTorch tooling that can analyze it. Why the compilation is worth doing, and what the intermediate representation is actually for.
For inference, a model only has to produce an output. Anything that computes the right number will do, and the fastest thing that does is usually the right choice.
For verification the requirement is different in kind. A verifier needs to know how the computation unfolds, how shapes propagate, and what each step means — and a black-box call gives it none of that. It cannot tighten a bound it cannot reason about, and it cannot reason about a program it has never seen.
This is the whole argument for compiling an ONNX model to PyTorch source rather than wrapping it in an interpreter. The value is not a faster execution format. It is that the structure which was implicit inside someone else’s dispatch table becomes explicit, named and editable.
Three questions, three representations
The reason a compiler has an intermediate representation is that a conversion has three questions in it, and they have different answers.
What is the graph? Nodes, edges, attributes, tensor shapes. This is structural and it is answerable without knowing anything about PyTorch.
What does each operation mean? Conv in ONNX is not Conv2d in PyTorch: the layout, the padding convention and the weight arrangement all differ, and translating them is where the real work is. This is the algorithmic step, and it is the only one that needs to know about the target.
How should the target program be written? Naming, structure, __init__ versus forward, how the state dict is loaded. This is a question about taste and convention, and it is independent of both of the above.
Answering them in one pass makes each of them harder. The same code that maps an operator also decides the name of the resulting layer; the same pass that discovers shapes also emits the __init__. Each decision becomes entangled with two others, and changing one means re-deriving the rest.
So the conversion is staged, and each stage has a representation it can be tested against:
ONNX Model
↓ normalize: validate, convert opset, infer shapes
Validated ONNX
↓ build: extract topology (nodes, edges, shapes, initializers)
Structural IR (ModelIR, NodeIR)
↓ analyze: classify tensors, map ops to PyTorch types
Semantic IR (SemanticModelIR)
↓ generate: emit __init__, forward(), state dict
Python source + state dict
↓ simplify: strip default args, drop unused buffers, format, add header
Final: model.py + model.pthThe value is not the diagram. It is that the structural IR carries no PyTorch in it at all — a node name, an operator type, raw attributes, tensor names and shapes, initializers still unclassified — so it can be inspected without also asking whether the mapping is right. analyze then attaches typing to that structure without either pass knowing anything about code emission: every tensor becomes a runtime variable, a parameter or a constant, every node acquires a PyTorch type and a class, and every literal argument is recorded against its PyTorch default. What arrives at generate is a fully resolved model with no ONNX types left in it, and code generation is its only consumer — which is why a new operator is a mapping entry plus a handler rather than an edit to a traversal that also does layout.
The implementation numbers five stages rather than four transformations, because simplifying the generated text, formatting it and writing the file header are a separate pass over the output rather than a transformation of the IR. Keeping them apart is what lets the code-generation templates stay about semantics.
The batch dimension is a design constraint
Verification evaluates many inputs at once — a whole batch of candidates per solver iteration — so how the converted program handles a batch is not a convenience. It decides what a verifier can do at all.
There is a specific reason a runtime wrapper resists this, and it is worth stating precisely because the intuitive version is wrong. torch.vmap runs the Python body once with a batched tensor, and relies on every operation inside having a batching rule. Two consequences follow, and they pull in opposite directions:
- A Python
forloop is fine. Its trip count must be knowable without reading a tensor value, but a static count works, as does one derived from.shape. A dict of tensor intermediates works too. The loop and the dict are not the obstacle, and a converter that contorts correct code on that assumption is solving a problem nobody has. - Reading a tensor’s value into Python is not fine.
.item(),.tolist(), anifon a tensor, or using a tensor as a shape all fail, because undervmapthe value in question is per-lane and there is no single value to read.
Measured on torch 2.12.1, eager, CPU:
| Body | Result | In the recorded run? |
|---|---|---|
| loop over a dict of tensors, static trip count | works, matches the per-input loop exactly | yes |
loop whose trip count is x.shape[1] | works | no |
y = x.clone(); y.add_(1) | works | yes |
x[torch.tensor(1).item():] — .item() on a tensor the body created | works | no |
x[:, s.item():] where s arrived with the batch | RuntimeError: We don't support vmap over calling .item() | yes, as x + x[0].item() |
x + 1 if x.sum() > 0 else x - 1 | RuntimeError: data-dependent control flow | no |
torch.add(x, x, out=torch.empty_like(x)) | RuntimeError: Batching rule not implemented for aten::add.out | yes |
The third column is there because scripts/verify/walkthroughs/vmap-report.json records five of these and not the rest, and a table that says “measured” without saying which rows are backed is the same overstatement in a smaller font. The recorded five are the loop-and-dict case, in-place on a clone, the batched .item() rejection, the out= rejection, and the fact that the body runs once.
The fourth and fifth rows are the pair that matters, and the fourth is the one the recorded run does not cover. What carries that claim is the mechanism stated above it rather than the row: a tensor the body constructed is not per-lane, so there is a single value to read, while s arrived with the batch and its value is per-lane. The row is an illustration of the mechanism, and .item() is not forbidden — reading a batched value is what fails. Compatibility is therefore per-operation rather than a property of a program being “functional”: in-place mutation on a clone is accepted while out= on the same arithmetic is not.
So the real obstacle for a wrapper is narrow: ONNX permits operations whose output shape or indices are data, and a dispatcher that forwards those into Python cannot batch them.
Value-dependent shapes
The hard case is an operation whose output shape depends on a runtime tensor value. The straightforward transliteration of a Slice with runtime bounds reads the bounds into Python:
start = start_tensor.item()
end = ends_tensor.item()
output = data[..., start:end]Under vmap this fails, and the failure is the right one: the direct form and a loop that hoists the same value both raise. Rewriting it as a gather does not rescue it — with a static index that is the case where the plain slice already worked, and when the bounds vary per lane the index varies per lane, so gather either reproduces the problem or yields per-lane shapes that disagree. A Slice whose length varies changes extent, which no gather selection expresses.
The honest response is to pin the extent and move the variability somewhere that can batch. The generated code keeps each sliced axis a static length, reads the start per lane as a tensor, and emits a validity flag alongside the result so a slice that fell out of bounds becomes zeroed data rather than a Python branch. Pad takes the other exit: with vmap_mode=True and runtime pads, code generation raises rather than emitting the .tolist() it knows will break torch.vmap.
What is checked, and what is not
A converter can be wrong in ways that are invisible from the outside: a handler that maps an operator to something plausible but wrong, or one that guesses when it should report. The properties worth asserting are about behaviour rather than appearance, and each is something a reader can check for themselves:
| Property | How to check it |
|---|---|
| The output is source you can read | Open it. It is a module, not a wrapper. |
| Weights round-trip | load_state_dict(..., strict=True) reports no missing or unexpected keys. |
| An unregistered operator fails loudly | This is the one worth caring about most. A converter that guesses produces a model that loads and computes the wrong answer. |
The module is vmap-traceable | torch.vmap(model) on a batch. Fails loudly, not silently. |
There is no conversion success rate quoted here, no runtime overhead figure and no batched speedup factor. Those are empirical claims about a hardware and software configuration, and without the configuration they are not claims a reader can check. What the tool does claim about coverage is narrower and testable: an ONNX operator outside the mapping tables is reported at analysis, and a mapped type with no registered handler is reported at generation — reported, never approximated.
A walkthrough was run. A constructed four-op graph over 40 inputs — eight hand-picked to sit on the kinks and at the boundaries, 32 from a fixed seed — on torch 2.12.1, numpy 2.4.6, onnx 1.16.0, onnxruntime 1.22.0, torchonnx 2026.8.5. Three of those five versions are read from __version__ because the packages ship no distribution metadata, so the report records what the modules claim to be rather than what was resolved by an installer; and the report carries an expected converter commit rather than a verified one, because it records no direct URL for the build it ran. The comparison result does not depend on either caveat, and the caveats are here because “at commit” would have been the wrong way to write it.
Three comparisons at rtol=1e-5, atol=1e-6, each with maximum absolute error 0.0: onnxruntime against a NumPy formula, the generated module against onnxruntime, and vmap against a per-input loop. That error is exactly zero and it is the weakest passing outcome rather than the strongest — two matmuls and a ReLU dispatch to the same kernels on both sides, and a model large enough for kernel selection and accumulation order to differ is the case this does not reach. It establishes the generated module computes the same function as the graph over these inputs and is vmap-traceable. It does not establish that the converter is correct. The script and the full report are in scripts/verify/walkthroughs/.
Implementation notes
Immutable IR. The intermediate representations are frozen dataclasses, which prevents an accidental mutation in one pass from being observed by the next and makes an IR inspectable at a breakpoint.
A registry of handlers. Mapping is a table keyed by ONNX operator type and split three ways — layer, operation or operator — and the code generators are registered separately, under the PyTorch type each one emits. An operator therefore needs an entry in both, and the pipeline says so rather than guessing: unmapped raises at analysis, mapped but unhandled raises at generation.
Generated code is meant to be read. Layer names are semantic rather than positional — self.conv1, not self.onnx_op_0 — each module opens with a header naming the ONNX file it came from, constructor arguments are written out rather than left at their PyTorch defaults, and the output is formatted to Black’s conventions. This is a design commitment rather than a nicety: the whole argument for compilation is that a reader will be able to look at it, and generated code nobody would have written by hand undermines that.
Known limitations
- Control flow:
If,LoopandScanare unmapped, so a graph containing one is reported at analysis rather than interpreted - Custom operators: must be implemented as PyTorch extensions first
- Training mode: inference only, which is the case verification cares about
- Dynamic shapes: limited. The batch dimension is the well-handled case, since it is genuinely uniform across a batch; spatial dimensions varying per sample are not
- Value-dependent shapes: an operation whose output extent depends on a runtime value has to pin that extent to batch — a
Slicekeeps a static length per axis and flags the slice that fell out of bounds, and aPadwith runtime pads is refused at code generation
Conclusion
A model that runs and a computation that can be read are different requirements, and ONNX only supplies the first. Compiling to PyTorch source does not make the model run faster so much as it makes the model legible: the layer names, the operator sequence and the shapes are all inspectable, so a surprising bound has somewhere to be explained, and in a form the existing PyTorch analysis stack can consume directly. A runtime wrapper moves the model’s semantics into a dependency’s dispatch table, where the analysis and the artefact it analyses become different objects.
Repository: https://github.com/ZhongkuiMa/torchonnx
Related projects:
- ShapeONNX: Shape inference with static shape resolution
- SlimONNX: ONNX optimizer for verification
- PropDAG: Graph traversal framework for bound propagation
Where to go next
Related software
- torchonnx
The tool this article walks through.
Continue reading
- ShapeONNX: Solving ONNX's Dynamic Shape Problem
The shape information a rewrite depends on.
- Neural Network Decomposition: The Unary-Binary Framework
The graph form these methods assume.