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

Replaying graph cleanup modifications

Replaying graph cleanup modifications

Replaying graph-rewriting patterns

Replaying graph-rewriting patterns

Gallery generated by Sphinx-Gallery

Example last updated

Date:

2026-10-05