Note
Go to the end to download the full example code.
Optimizing a model with graph-rewriting patterns#
onnx-light optimizes a model by repeatedly matching small
PatternOptimization subgraphs and
replacing them with a simplified equivalent. This example walks through the
whole workflow on a tiny model:
Build a model containing redundant nodes (two useless
Castand two consecutiveNeg).Run the standard patterns and inspect the optimization statistics (
OptimizationReport).Write a custom pattern.
Inspect a candidate rejected by that pattern and the recorded failure reason.
Apply the custom pattern together with the standard ones.
See How to add a custom graph-rewriting pattern and set its priority for a companion how-to that focuses on writing and registering a pattern (including how priorities order patterns), in both Python and C++. Replaying graph-rewriting patterns demonstrates how to capture and replay the applied modifications separately.
# sphinx_gallery_thumbnail_path = "_static/gallery_thumbnails/pattern_optimization.png"
from __future__ import annotations
import onnx_light.onnx.helper as oh
from onnx_light.onnx_lib import parser
from onnx_light.onnx_core.optimization import (
GraphBuilder,
GraphGraph,
PatternOptimization,
standard_patterns,
)
from onnx_light.tools import pretty_onnx
Build a model with redundant nodes#
x is cast to float (a no-op, since it is already float) and then
negated twice before being cast back. The standard Cast pattern removes
the two Cast nodes; the custom pattern added further below removes the
two Neg nodes.
model = parser.parse_model(
'<ir_version: 10, opset_import: ["" : 18]>\n'
"agraph (float[4] x) => (float[4] y) {\n"
" casted = Cast <to=1> (x)\n"
" middle = Neg(casted)\n"
" negated = Neg(middle)\n"
" y = Cast <to=1> (negated)\n"
"}\n"
)
print(pretty_onnx(model))
opset: domain='' version=18
graph: name='agraph'
input: float[4] x
0: Cast(x) -> casted
1: Neg(casted) -> middle
2: Neg(middle) -> negated
3: Cast(negated) -> y
output: float[4] y
Run the standard patterns and read the statistics#
optimize() returns the
list of applied rewrites; passing report=True additionally returns an
OptimizationReport with timing
and match/no-match counters for every pattern that was tried.
builder = GraphBuilder(model)
graph = GraphGraph(builder, standard_patterns(["Cast"]))
rewrites, report = graph.optimize(report=True)
print(pretty_onnx(builder.to_onnx("model")))
print(report)
for pattern_stats in report.patterns:
print(
f"{pattern_stats.pattern_name}: {pattern_stats.matches} match(es) over "
f"{pattern_stats.attempts} attempt(s)"
)
opset: domain='ai.onnx' version=18
graph: name='agraph'
input: float[4] x
0: Neg(x) -> middle
1: Neg(middle) -> negated
2: Identity(negated) -> y
output: float[4] y
OptimizationReport(iterations=2, rewrites=3, total_time_ns=111137, phases={matching: 5790, rewriting: 28994, cleanup: 42439, constant_folding: 33914, subgraph_optimization: 0}, patterns=[Cast(attempts=2, matches=2, match_time_ns=3036, apply_time_ns=7404)], subgraphs=[])
Cast: 2 match(es) over 2 attempt(s)
Add a custom pattern#
A pattern derives from
PatternOptimization and
implements fast_op_type (the operator types it may start from), match
(returns self.result(...) on success or self.no_match(...) with a
diagnostic otherwise) and apply (builds the replacement nodes). Passing
it to GraphGraph alongside the standard patterns runs both together, in
ascending priority order.
class NegNegPattern(PatternOptimization):
"""Replaces two consecutive Neg nodes with Identity."""
def __init__(self):
super().__init__(priority=1, name="NegNeg")
def fast_op_type(self):
return {"Neg"}
def match(self, graph, node):
previous = graph.node_before(node.input[0])
if previous is None or previous.op_type != "Neg":
return self.no_match(node, "the input is not produced by Neg")
return self.result([previous, node], insert_at=node)
def apply(self, graph, nodes):
del graph
previous, node = nodes
return [oh.make_node("Identity", [previous.input[0]], list(node.output))]
Inspect a failed pattern candidate#
The same custom pattern rejects a single Neg because its input is not
produced by another Neg. Returning self.no_match(candidate, reason)
keeps that rejection out of normal output, but the optional report stores the
aggregated reason so the pattern author can understand why nothing changed.
failure_model = parser.parse_model(
'<ir_version: 10, opset_import: ["" : 18]>\n'
"agraph (float[4] x) => (float[4] y) {\n"
" y = Neg(x)\n"
"}\n"
)
builder = GraphBuilder(failure_model)
graph = GraphGraph(builder, [NegNegPattern()])
failed_rewrites, failed_report = graph.optimize(report=True)
print(f"failed rewrite count: {len(failed_rewrites)}")
for pattern_stats in failed_report.patterns:
for no_match in pattern_stats.no_matches:
print(
f"{pattern_stats.pattern_name} rejected "
f"{no_match.occurrences} candidate(s): {no_match.reason}"
)
assert not failed_rewrites
failed_no_match_reasons = {
no_match.reason
for pattern_stats in failed_report.patterns
if pattern_stats.pattern_name == "NegNeg"
for no_match in pattern_stats.no_matches
}
assert "the input is not produced by Neg" in failed_no_match_reasons
failed rewrite count: 0
NegNeg rejected 1 candidate(s): the input is not produced by Neg
Apply the custom pattern with the standard patterns#
builder = GraphBuilder(model)
graph = GraphGraph(builder, [*standard_patterns(["Cast"]), NegNegPattern()])
rewrites = graph.optimize()
fully_optimized = builder.to_onnx("model")
print(pretty_onnx(fully_optimized))
print([rewrite.pattern_name for rewrite in rewrites])
assert [node.op_type for node in fully_optimized.graph.node] == ["Identity"]
opset: domain='ai.onnx' version=18
graph: name='agraph'
input: float[4] x
0: Identity(x) -> y
output: float[4] y
['Cast', 'Cast', 'RemoveIdentityNodes', 'NegNeg', 'RemoveIdentityNodes']
Select patterns by name, regex or target device#
Lists are exclusive; they do not implicitly add registered defaults.
Regexes use fullmatch. None selects all device-independent defaults,
and False disables patterns while retaining cleanup passes.
builder = GraphBuilder(model)
graph = GraphGraph(builder, patterns=r"Cast.*")
assert graph.patterns
assert all(pattern.name.startswith("Cast") for pattern in graph.patterns)
graph = GraphGraph(builder, patterns=False)
assert not graph.patterns
# A device selector adds only patterns targeting that exact device to the
# independent defaults and sets the builder target. Existing standard patterns
# are independent, even when their matching heuristics inspect the target.
from onnx_light.onnx_core.shape_inference import Device
graph = GraphGraph(builder, patterns=Device.kCPU)
assert builder.device == Device.kCPU
assert all(pattern.device in (Device.kUndefined, Device.kCPU) for pattern in graph.patterns)
Total running time of the script: (0 minutes 0.008 seconds)
Related examples
Gallery generated by Sphinx-Gallery
Example last updated
- Date:
2026-10-05