ShapeONNX: Solving ONNX's Dynamic Shape Problem
A dual-track shape inference tool that resolves ONNX's dynamic shapes to concrete static values for neural network verification workflows.
Introduction
The Fundamental Problem:
ONNX was designed to answer: “What shape does this tensor have?” (metadata)
But verification tools need: “What values does this shape tensor contain?” (semantics)
Example:
# Verification tool receives Vision Transformer with:
x = torch.randn(1, 3, 224, 224)
patches = x.reshape(1, 196, 768) # 196 patches = (224/16)^2By default, shape inference records the rank of the Reshape target and stops
there: a 1D tensor with 3 elements, annotated [3]. What a verification tool
needs is the other half — that those 3 elements contain [1, 196, 768].
The tempting conclusion is that this is a bug in ONNX, or at least that ONNX has
no answer. Both are wrong, and it is worth being precise about what ONNX
actually offers rather than leaving the stronger claim standing here to be
corrected thirty lines later. onnx.shape_inference.infer_shapes takes a
data_prop argument, and the
official documentation
describes it as enabling “data propagation for limited operators to perform
shape computation”. Two things follow from that wording. It defaults to
False, so the behaviour above is what a caller gets without asking for
anything. And it is limited to a subset of operators, so a network that
computes its reshape target through a chain the propagator does not cover
still ends up annotated [3].
So the gap is narrower than “ONNX cannot do this”: data propagation is available, opt-in, and partial, and a verification pipeline that needs every shape resolved on every model cannot rely on it. ShapeONNX computes the shape-tensor values itself, for the operators a verification pipeline actually meets — which is what “a different problem domain” was gesturing at, without the overstatement.
Solution: Dual-track representation. Ask two questions separately:
data_shapes["x"]: “What shape?” →[1, 196, 768](metadata)explicit_shapes["shape_tensor"]: “What values?” →[1, 196, 768](semantics)
Result: a single topological pass per graph. That is a structural property and it is what the implementation does; a per-model timing is not claimed here, for the reason given under Part 3.
Part 1: The Dual-Track Representation
Core idea: Two dictionaries instead of one.
# Traditional ONNX: only data_shapes
shapes = {
"input": [1, 48, 2, 2],
"shape_output": [4], # "This is a 1D tensor"
"reshape_output": [-1, -1], # "Dynamic, don't know"
}
# ShapeONNX: data_shapes + explicit_shapes
data_shapes = {
"input": [1, 48, 2, 2],
"shape_output": [4],
"reshape_output": [2, 2]
}
explicit_shapes = {
"shape_output": [1, 48, 2, 2], # "Contains these values"
"gather_output": [2, 2] # "Sliced to these values"
}Decision Table - When to use which dictionary:
| Operator | Input (data/explicit) | Output (data/explicit) | Logic |
|---|---|---|---|
| Shape | data only | (None, data) | Extract dimension values |
| Gather on shapes | explicit | (None, sliced) | Slice shape values |
| Slice on shapes | explicit | (None, sliced) | Slice shape values |
| Concat shapes | explicit | (None, concatenated) | Concatenate values |
| ConstantOfShape | explicit | (explicit, None) | Output shape = input values |
| Reshape | data + explicit | (output, None) | Target shape from explicit |
Retrieval pattern (used by all operators):
def get_shape(name, data_shapes, explicit_shapes):
# Priority: explicit values first (for shape tensors)
if name in explicit_shapes:
return explicit_shapes[name], True # "I know the values"
elif name in data_shapes:
return data_shapes[name], False # "I know the dimensions"
else:
raise RuntimeError(f"Shape unknown: {name}")Part 2: How It Works
Vision Transformer patch embedding (step-by-step):
-
Input:
data_shapes["input"] = [1, 3, 224, 224] -
Shape operator: Extract dimension values
explicit_shapes["shape_vec"] = [1, 3, 224, 224] data_shapes["shape_vec"] = [4] (it's a 1D tensor) -
Gather(indices=[2,3]): Slice the shape values
explicit_shapes["spatial_dims"] = [224, 224] data_shapes["spatial_dims"] = [2] (it's a 1D tensor) -
Div(16): Compute patch count
explicit_shapes["patch_count"] = [14, 14] (224/16 = 14) -
Mul: Flatten to scalar
explicit_shapes["n_patches"] = 196 -
Concat([1, 196, -1]): Build target shape
explicit_shapes["reshape_target"] = [1, 196, -1] -
Reshape: Use explicit shape to infer -1
data_shapes["output"] = [1, 196, 768] (inferred: total 1*3*224*224 / (1*196) = 768)
ONNX reports by default: [-1, -1, -1] (all dynamic). With
data_prop=True it can resolve chains its propagator covers, and stops at
[3] on the ones it does not.
ShapeONNX reports: [1, 196, 768] (fully static)
Part 3: Key Design Decisions
Single-pass inference: ONNX graphs are topologically sorted. Process each node once in order. No multi-pass analysis needed.
Performance: Pure Python, no C extensions (the package declares requires-python = ">=3.11"). Shape inference is a single linear pass over the graph, so its cost is in the number of nodes rather than in the depth of the network the way bound propagation is: O(nodes) against O(nodes x layers) for per-layer bounding.
That is a complexity claim, and it is where the claim stops. A wall-clock figure would need a configuration to be meaningful, and this repository does not offer one: it carries no timing code, and its README quotes no benchmark, so there is nothing here to reproduce a number against. The complexity is the part that survives the tool being changed.
Part 4: Common Questions
Q: Can I use only explicit_shapes?
No. data_shapes tracks regular tensor operations. explicit_shapes tracks shape tensor operations. Both are needed.
Q: What if a shape tensor’s value isn’t computed statically?
Falls back to data_shapes. ShapeONNX resolves what it can statically determine, and what it cannot is bounded by what the single pass over a topologically sorted graph can reach:
- Asymmetric padding: not supported. Rare in verification models, which use symmetric padding.
- Control flow (
If/Loop): not supported. The pass follows a topological order, and a graph with control flow has no single one to follow — which is the same property that makes the pass cheap also being what excludes them. Most verification models use static graphs, so this costs little in practice. - Dynamic input shapes: assumes static inputs, or
batch_size=1. - Truly dynamic shapes (processing arbitrary image sizes, say) require symbolic analysis, which this is not. That is a different method rather than a slower version of this one.
Q: How does this integrate with verification tools?
Pass the resolved data_shapes dict to your bound propagator. Static dimensions enable:
- Bound allocation for each neuron
- Constraint generation with known dimensions
- Memory allocation without surprises
Q: How fast is it? One linear pass per graph, over the nodes, with the cost in the number of nodes rather than the depth of the network. The repository carries no timing code and its current README quotes no benchmark, so there is nothing here to reproduce a figure against.
Part 5: Verification Integration
Typical workflow:
PyTorch/TensorFlow model
↓ export
ONNX model (contains dynamic shapes, shape operators)
↓ ShapeONNX (resolve shape values)
Static shape dictionary (e.g., output = [1, 1000])
↓ SlimONNX (optional optimization)
Simplified graph (redundant ops removed)
↓ Verification tool (CROWN, DeepPoly)
Safety proof or counterexampleWhat ShapeONNX does: Computes shape tensor values (the explicit_shapes dict). Verification tool uses this to allocate bounds for each neuron.
Real impact: Vision Transformer models, whose patch embedding is where dynamic shapes come from, stop being blocked at the shape step. Whether that unblocks the verification depends on what comes after it.
For the fixed input domain that verification normally assumes, static shapes are sufficient.
Try It Yourself
git clone https://github.com/ZhongkuiMa/shapeonnx.git
cd shapeonnx
pip install -e ".[dev]"The tool is not on PyPI, so the clone is the install path rather than a
convenience. Its pyproject.toml pins onnx==1.16.0 and numpy==1.26.4, and
has pinned those two in every release — an environment that satisfies this and
the SlimONNX note’s pins cannot satisfy both.
import onnx
from shapeonnx import infer_onnx_shape
from shapeonnx.utils import get_initializers, get_input_nodes, get_output_nodes
# Load model
model = onnx.load("your_model.onnx")
model = onnx.version_converter.convert_version(model, target_version=21)
# Infer shapes
initializers = get_initializers(model)
input_nodes = get_input_nodes(model, initializers, has_batch_dim=True)
output_nodes = get_output_nodes(model, has_batch_dim=True)
shapes = infer_onnx_shape(
input_nodes, output_nodes,
list(model.graph.node), initializers,
has_batch_dim=True
)
# Use static shapes in verification
print(f"Output shape: {shapes['output']}") # [1, 1000] not [-1, -1]Conclusion
ShapeONNX is not a criticism of ONNX—ONNX’s design for dynamic shapes is correct for deployment. ShapeONNX is a complementary tool for verification workflows that require static dimensions.
Key insight: Track shape tensor values separately from regular tensor metadata.
Result: Resolve the shapes that default inference, and partial data propagation, leave dynamic — to concrete static values, enabling verification on models whose shapes are computed rather than annotated.
If you’re building verification tools, ShapeONNX bridges the gap between ONNX’s flexible dynamic shapes and verification’s need for concrete dimensions.
Repository: https://github.com/ZhongkuiMa/shapeonnx
Where to go next
Related software
- shapeonnx
The tool this article walks through.
Continue reading
- SlimONNX: A Story of Optimizing Neural Networks for Verification
The next rewrite in the same ONNX pipeline.
- TorchONNX: A Compiler for ONNX-to-PyTorch Conversion
The conversion step that consumes a simplified graph.