Skip to main content

onnx-shape-inference

PyPI - Version PyPI - Python Version codecov Ruff PyPI Downloads

Experimental symbolic shape inference for ONNX models. Built on top of ONNX IR, this library performs shape inference directly on the IR without serialization overhead, using SymPy for symbolic dimension arithmetic.

Features

  • Symbolic shape inference — propagates shapes through the graph using SymPy expressions for symbolic dimensions
  • Shape data propagation — tracks known element values of shape tensors (e.g. through Shape → Slice → Concat → Reshape chains) to resolve concrete output shapes that standard shape inference cannot
  • Broad operator coverage — built-in inference for standard ONNX (ai.onnx) operators plus com.microsoft contrib ops, with version-aware dispatch across opset history
  • Symbolic constraint resolution — reconciles engine-generated dimension names with the symbolic names an author declares on graph outputs / value_info, renaming anonymous dims (including compound expressions like 2*_d0 or past_seq + seq) to the declared names
  • Extensible registry — register custom shape inference functions for custom operators
  • Merge policies — control how newly inferred shapes are merged with existing ones: refine (default), strict, override, and skip

Installation

pip install onnx-shape-inference

Or install from source (main branch):

pip install git+https://github.com/justinchuby/onnx-shape-inference.git

Command line

Run shape inference on a model and see how many new shapes were inferred:

onnx-shape-inference model.onnx

Save the inferred model to a file:

onnx-shape-inference model.onnx -o model_inferred.onnx

Overwrite the input model in place:

onnx-shape-inference model.onnx --in-place

Select a different merge policy:

onnx-shape-inference model.onnx --policy strict

Usage

import onnx_ir as ir
from onnx_shape_inference import infer_symbolic_shapes

# Load a model
model = ir.load("model.onnx")

# Run shape inference
model = infer_symbolic_shapes(model)

# Or with a strict merge policy
model = infer_symbolic_shapes(model, policy="strict")

Use with onnxscript optimizer

You can run symbolic shape inference on the model to help the optimizer discover more optimization opportunities.

import onnx_shape_inference
import onnx_ir as ir
import onnxscript.optimizer

model = ir.load("model.onnx")

# Provide more shape information with infer_symbolic_shapes
model = onnx_shape_inference.infer_symbolic_shapes(model)

# onnxscript optimizer can leverage this information to better optimize the model
onnxscript.optimizer.optimize(model)

ir.save(model, "model_optimized.onnx")

Per-node inference

You can run shape inference on individual nodes by using the ShapeInferenceContext and registry directly. This is useful for debugging, testing, or integrating into custom graph passes.

import onnx_ir as ir
from onnx_shape_inference import ShapeInferenceContext, registry

# Populate the registry with all built-in ops
registry.collect()

# Create a context with the model's opset imports
ctx = ShapeInferenceContext(opset_imports={"": 21})

# Look up the inference function for the op
infer_func = registry.get("", "Relu", version=21)

# Build a node (or get one from an existing graph)
x = ir.Value(name="x", shape=ir.Shape([2, 3]), type=ir.TensorType(ir.DataType.FLOAT))
y = ir.Value(name="y")
node = ir.Node("", "Relu", inputs=[x], outputs=[y])

# Run inference
infer_func(ctx, node)

print(y.shape)  # [2,3]
print(y.dtype)  # FLOAT

Registering custom operators

from onnx_shape_inference import registry

@registry.register("com.custom", "MyOp", since_version=1)
def infer_my_op(ctx, node):
    input_shape = node.inputs[0].shape
    output_shape = ir.Shape([...])
    ctx.set_shape(node.outputs[0], output_shape)

Shape data propagation (pkg.onnx_shape_inference.sym_data)

Shape inference alone cannot resolve output shapes when ops like Reshape consume non-constant shape tensors that were computed at runtime (e.g. Shape → Slice → Concat → Reshape). The sym_data feature bridges this gap by tracking the known element values of 1-D integer tensors as they flow through the graph.

After inference, each value that carries propagated data has a pkg.onnx_shape_inference.sym_data entry in its metadata_props. You can read it directly or use the SYM_DATA_KEY constant:

import json
import numpy as np
import onnx_ir as ir
from onnx_shape_inference import SYM_DATA_KEY, infer_symbolic_shapes

model = infer_symbolic_shapes(model)

for node in model.graph:
    for value in node.inputs:
        if SYM_DATA_KEY in value.metadata_props:
            text = value.metadata_props[SYM_DATA_KEY]  # e.g. '["N",3,768]'
            elements = json.loads(text)                # ["N", 3, 768]

            # You can create an ir.Shape from it
            shape = ir.Shape(elements)

            # Then you can replace this input with a constant value

When all elements are concrete integers the value is also stored as a constant tensor, so downstream consumers that read constants directly can access it without parsing metadata_props.

Adopting declared symbolic names (constraint resolution)

Per-operator inference names data-dependent dimensions with anonymous symbols (_d0, _d1, …). Model authors, however, usually declare meaningful symbolic names on graph outputs and value_info (e.g. Y: [batch, seq]). After the main inference pass, a constraint-resolution pass records equalities between inferred and declared shapes and renames the anonymous symbols to the author's declared names — including compound occurrences such as 2*_d0 → 2*batch and past_seq + seq → total_seq. Only anonymous _dN symbols are renamed; declared names are treated as authoritative.

This is enabled by default. Pass adopt_declared_symbols=False to keep the raw engine-generated symbols instead:

model = infer_symbolic_shapes(model, adopt_declared_symbols=False)

Development

pip install -r requirements/ci/requirements.txt
pip install -e .
pytest

Lint and format with lintrunner:

pip install -r requirements/lintrunner/requirements.txt
lintrunner -a

Fuzzing

A deterministic, seeded fuzzer stress-tests shape inference against ONNX reference inference and ONNX Runtime. A fast tier runs as part of the normal test suite; replay a failing seed with:

FUZZ_SEED=<seed> python3 -m pytest tests/shape_inference_fuzz_test.py

See docs/fuzzing.md for how to turn a fuzzer finding into a regression test, and docs/fuzzing-design.md for the fuzzer's design (generator, oracles, harness, and shrinking).

License

MIT

Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

onnx_shape_inference-0.3.2.tar.gz (611.0 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

onnx_shape_inference-0.3.2-py3-none-any.whl (248.5 kB view details)

Uploaded Python 3

File details

Details for the file onnx_shape_inference-0.3.2.tar.gz.

File metadata

  • Download URL: onnx_shape_inference-0.3.2.tar.gz
  • Upload date:
  • Size: 611.0 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/7.0.0 CPython/3.13.14

File hashes

Hashes for onnx_shape_inference-0.3.2.tar.gz
Algorithm Hash digest
SHA256 0deffa9639986c9326435edf4920c59c2fde97a9a6a200c20ff3b08a7adf21ca
MD5 d8389da5623a374f16d6bf76a70af6a2
BLAKE2b-256 cb49d2785d775f2e233bc4e78a84424e999e08a7e15274dac7681dc61a849f88

See more details on using hashes here.

Provenance

The following attestation bundles were made for onnx_shape_inference-0.3.2.tar.gz:

Publisher: main.yml on justinchuby/onnx-shape-inference

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

File details

Details for the file onnx_shape_inference-0.3.2-py3-none-any.whl.

File metadata

File hashes

Hashes for onnx_shape_inference-0.3.2-py3-none-any.whl
Algorithm Hash digest
SHA256 0500a46d3ff45c0380feff5bf5b5a5dbb667720dfc06fbc63aa02d67530d51c8
MD5 b8007f4d89aa743257393b08ca24f1f3
BLAKE2b-256 988a41f459c946b1880d6b7df72705abcf56366a15c1449aec11f0dc97255768

See more details on using hashes here.

Provenance

The following attestation bundles were made for onnx_shape_inference-0.3.2-py3-none-any.whl:

Publisher: main.yml on justinchuby/onnx-shape-inference

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

Release history Release notifications | RSS feed

This release

0.3.2 This release

2 files

0.3.1

2 files

0.3.0

2 files

0.2.0

2 files

0.1.9

2 files

0.1.7

2 files

0.1.6

2 files

0.1.5

2 files

0.1.4

2 files

0.1.3

2 files

0.1.2

2 files

0.1.1

2 files

0.0.1

2 files

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page