Note
Go to the end to download the full example code.
Replaying graph-rewriting patterns#
Every successful pattern optimization produces a
LocalRewriting record. These
records can reconstruct the optimized graph from the original model without
running pattern matching again.
# sphinx_gallery_thumbnail_path = "_static/gallery_thumbnails/pattern_replay.png"
from __future__ import annotations
from onnx_light.onnx_lib import parser
from onnx_light.onnx_core.optimization import GraphBuilder, GraphGraph, replay, standard_patterns
from onnx_light.tools import pretty_onnx
Create and optimize the source model#
The source graph contains two type-preserving Cast nodes. The optimizer
rewrites them and returns the corresponding modification records.
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"
" negated = Neg(casted)\n"
" y = Cast <to=1> (negated)\n"
"}\n"
)
builder = GraphBuilder(model)
graph = GraphGraph(builder, standard_patterns(["Cast"]))
rewrites = graph.optimize()
optimized_graph = builder.build_graph()
print("Original graph:")
print(pretty_onnx(model))
print("Optimized graph:")
print(pretty_onnx(builder.to_onnx("model")))
Original graph:
opset: domain='' version=18
graph: name='agraph'
input: float[4] x
0: Cast(x) -> casted
1: Neg(casted) -> negated
2: Cast(negated) -> y
output: float[4] y
Optimized graph:
opset: domain='ai.onnx' version=18
graph: name='agraph'
input: float[4] x
0: Neg(x) -> negated
1: Identity(negated) -> y
output: float[4] y
Inspect the captured modifications#
Each record has a concise one-line display. Use to_detailed_string to
inspect the pattern, matched nodes, inserted nodes, and optimization iteration
needed to reproduce one modification.
for rewrite in rewrites:
print(rewrite)
print(rewrite.to_detailed_string())
LocalRewriting(pattern=Cast, graph_path=<root>, matched_nodes=1, added_nodes=1)
LocalRewriting:
pattern: Cast
graph_path: <root>
iteration: 0
matched_nodes:
positions: [0]
added_nodes:
nodes: [Identity(outputs=[casted])]
positions: [0]
initializers:
added: []
added_positions: []
removed: []
value_renames: []
timings:
match_time_ns: 2695
apply_time_ns: 3897
LocalRewriting(pattern=Cast, graph_path=<root>, matched_nodes=1, added_nodes=1)
LocalRewriting:
pattern: Cast
graph_path: <root>
iteration: 0
matched_nodes:
positions: [2]
added_nodes:
nodes: [Identity(outputs=[y])]
positions: [2]
initializers:
added: []
added_positions: []
removed: []
value_renames: []
timings:
match_time_ns: 231
apply_time_ns: 821
LocalRewriting(pattern=RemoveIdentityNodes, graph_path=<root>, matched_nodes=3, added_nodes=2)
LocalRewriting:
pattern: RemoveIdentityNodes
graph_path: <root>
iteration: 1
matched_nodes:
positions: [0, 1, 2]
added_nodes:
nodes: [Neg(outputs=[negated]), Identity(outputs=[y])]
positions: [0, 1]
initializers:
added: []
added_positions: []
removed: []
value_renames: [casted->x]
timings:
match_time_ns: 0
apply_time_ns: 0
Replay without matching patterns#
replay() applies the captured
records to a fresh copy of the source model. The reconstructed graph is
byte-for-byte identical to the graph produced by the optimizer.
replayed_graph = replay(model, rewrites)
assert replayed_graph.SerializeToString() == optimized_graph.SerializeToString()
print("Replayed graph:")
print(pretty_onnx(replayed_graph))
Replayed graph:
graph: name='agraph'
input: float[4] x
0: Neg(x) -> negated
1: Identity(negated) -> y
output: float[4] y
Total running time of the script: (0 minutes 0.004 seconds)
Related examples
Gallery generated by Sphinx-Gallery
Example last updated
- Date:
2026-10-05