Note
Go to the end to download the full example code.
Replaying graph cleanup modifications#
Graph cleanup algorithms also produce
LocalRewriting records. These
records replay identity removal, dead-end removal, and initializer
deduplication without running cleanup again.
# sphinx_gallery_thumbnail_path = "_static/gallery_thumbnails/pattern_replay_cleanup.png"
from __future__ import annotations
from onnx_light.onnx import TensorProto
import onnx_light.onnx.helper as oh
from onnx_light.onnx_core.optimization import GraphBuilder, GraphGraph, replay
from onnx_light.tools import pretty_onnx
Create and clean up the source model#
This graph contains an Identity node, a dead-end Neg node, and two
equal initializers used by retained nodes.
model = oh.make_model(
oh.make_graph(
[
oh.make_node("Add", ["x", "weight"], ["summed"]),
oh.make_node("Identity", ["summed"], ["forwarded"]),
oh.make_node("Add", ["forwarded", "duplicate_weight"], ["y"]),
oh.make_node("Neg", ["x"], ["dead_end"]),
],
"cleanup",
[oh.make_tensor_value_info("x", TensorProto.FLOAT, [1])],
[oh.make_tensor_value_info("y", TensorProto.FLOAT, [1])],
initializer=[
oh.make_tensor("weight", TensorProto.FLOAT, [1], [1.0]),
oh.make_tensor("duplicate_weight", TensorProto.FLOAT, [1], [1.0]),
],
),
opset_imports=[oh.make_opsetid("", 18)],
)
builder = GraphBuilder(model)
graph = GraphGraph(builder, patterns=False)
rewrites = list(graph.optimize())
optimized_graph = builder.build_graph()
assert {"RemoveIdentityNodes", "RemoveUnusedNodes", "RemoveDuplicateInitializers"} <= {
rewrite.pattern_name for rewrite in rewrites
}
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='cleanup'
input: float[1] x
init: float[1] weight
init: float[1] duplicate_weight
0: Add(x, weight) -> summed
1: Identity(summed) -> forwarded
2: Add(forwarded, duplicate_weight) -> y
3: Neg(x) -> dead_end
output: float[1] y
Optimized graph:
opset: domain='ai.onnx' version=18
graph: name='cleanup'
input: float[1] x
init: float[1] weight
0: Add(x, weight) -> summed
1: Add(summed, weight) -> y
output: float[1] y
Inspect and replay the cleanup modifications#
Every cleanup operation is captured as a LocalRewriting record. Replay
applies the records to a fresh copy of the source model.
for rewrite in rewrites:
print(rewrite)
replayed_graph = replay(model, rewrites)
assert replayed_graph.SerializeToString() == optimized_graph.SerializeToString()
print("Replayed graph:")
print(pretty_onnx(replayed_graph))
LocalRewriting(pattern=RemoveIdentityNodes, graph_path=<root>, matched_nodes=4, added_nodes=3)
LocalRewriting(pattern=RemoveUnusedNodes, graph_path=<root>, matched_nodes=3, added_nodes=2)
LocalRewriting(pattern=RemoveDuplicateInitializers, graph_path=<root>, matched_nodes=2, added_nodes=2)
Replayed graph:
graph: name='cleanup'
input: float[1] x
init: float[1] weight
0: Add(x, weight) -> summed
1: Add(summed, weight) -> y
output: float[1] 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