Pattern optimization#
onnx-light rewrites graphs with a pattern-based optimizer built
directly on top of GraphBuilder. The optimizer
recognizes local subgraphs and replaces them with cheaper equivalents,
in the spirit of the Python pattern optimizer it is ported from. The
implementation plan and the pull requests that delivered it are recorded
in Pattern-based optimization in GraphBuilder.
Overview#
Optimization always operates on a
GraphBuilder through
GraphGraph. GraphGraph
wraps a builder with a structural index (successors, predecessors, shape,
type and constant queries) and drives a match/apply loop:
from onnx_light.onnx import TensorProto
from onnx_light.onnx_core.graph_builder import GraphBuilder
from onnx_light.onnx_core.optimization import GraphGraph, standard_patterns
builder = GraphBuilder("optimization")
builder.set_opset_version("", 18)
x = builder.inp("X", TensorProto.FLOAT, [4])
y = builder.op.Cast(x, outputs="Y", to=TensorProto.FLOAT)
builder.out(y, TensorProto.FLOAT, [4])
graph = GraphGraph(
builder,
standard_patterns(["Cast"]),
)
rewrites, report = graph.optimize(report=True)
optimized_model = builder.to_onnx("model")
assert len(rewrites) == 1
assert report.rewrites == 1
assert optimized_model.graph.node[0].op_type == "Identity"
Because the optimizer reuses the builder, it inherits the builder’s shape
and type inference, its constant knowledge and its cleanup passes
(RemoveIdentityNodes, RemoveUnusedNodes,
RemoveDuplicateNodes) instead of duplicating them.
The rewrite invariant#
A pattern must never reuse an existing name: every value it produces is new. This invariant keeps the successor and predecessor maps valid between two rewrites of the same iteration, which is why the builder records every name it hands out and never reuses one.
Pattern registration#
Available registries are merged by the stable PatternOptimization.name;
the patterns selector then chooses which entries to run:
global patterns (
register_pattern()) are available to every newGraphGraph; the standard ONNX patterns are registered globally when the module is imported;builder patterns (
GraphBuilder.register_pattern) override a global pattern for optimizers built over that builder;graph patterns (
GraphGraph(builder, patterns=[...])) form an exclusive selection. Explicit instances override earlier entries with the same name and are retained for that optimizer, including recursive subgraphs.
By default only device-independent patterns are selected. patterns=False
disables patterns without disabling cleanup. A concrete Device selects
independent patterns and patterns targeting that exact device, and sets
builder.device unless it conflicts with an already-defined target.
A regex selects available names using fullmatch. Regexes and lists can
explicitly select device-specific patterns and leave the builder device unchanged.
See onnx_light.onnx_core.optimization for the selector and migration details.
In C++, GraphGraph(builder) likewise selects only device-independent
registered patterns. CreateRegisteredPatterns(device) returns independent
and matching-device patterns; the no-argument factory returns all patterns.
When using a filtered factory result in C++, call optimizer.SetTargetDevice(device)
before optimizing to set a common target and check for conflicts in subgraphs.
Patterns can be written in C++ or in Python; both share the
PatternOptimization
interface, a match step that returns a
MatchResult and an apply
step that produces the replacement nodes.
Custom pattern example#
The following pattern replaces two consecutive Neg nodes with an
Identity. It restricts candidate nodes to Neg, checks the producer of
the candidate’s input, and returns the replacement from apply / Apply.
import onnx_light.onnx.helper as oh
from onnx_light.onnx_core.optimization import GraphGraph, PatternOptimization
class NegNegPattern(PatternOptimization):
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):
previous, node = nodes
return [
oh.make_node("Identity", [previous.input[0]], list(node.output))
]
graph = GraphGraph(builder, [NegNegPattern()])
graph.optimize()
#include "onnx_core/builder/graph_graph.h"
#include "onnx_core/builder/pattern_optimization.h"
namespace builder = onnx_light::core::builder;
class NegNegPattern final : public builder::PatternOptimization {
public:
NegNegPattern() : PatternOptimization(/*priority=*/1, "NegNeg") {}
std::set<std::string> FastOpType() const override { return {"Neg"}; }
builder::MatchResult Match(builder::GraphGraph &graph,
const onnx_light::NodeProto &candidate) const override {
const auto *previous = graph.NodeBefore(candidate.input()[0].value());
if (previous == nullptr || previous->op_type().value() != "Neg") {
return NoMatch(candidate, "the input is not produced by Neg");
}
return builder::MatchResult{this, {previous, &candidate}, &candidate};
}
onnx_light::utils::RepeatedProtoField<onnx_light::NodeProto>
Apply(builder::GraphGraph &,
const std::vector<const onnx_light::NodeProto *> &nodes) const override {
onnx_light::utils::RepeatedProtoField<onnx_light::NodeProto> replacements;
replacements.push_back(onnx_light::MakeNode(
"Identity", {nodes[0]->input()[0].value()}, {nodes[1]->output()[0].value()}));
return replacements;
}
};
std::vector<std::unique_ptr<builder::PatternOptimization>> patterns;
patterns.push_back(std::make_unique<NegNegPattern>());
builder::GraphGraph graph(graph_builder, std::move(patterns));
graph.Optimize();
The pattern is local to this optimizer. See How to add a custom graph-rewriting pattern and set its priority for global and builder registration, diagnostics, and priority selection.
API reference#
Python API and registered pattern list: onnx_light.onnx_core.optimization; the runtime list is available through
standard_pattern_names().C++ API: builder.
Examples#
Optimizing a model with graph-rewriting patterns covers optimization statistics.
Replaying graph-rewriting patterns demonstrates deterministic replay from captured rewrites.
How to add a custom graph-rewriting pattern and set its priority is a Python/C++ how-to on writing a custom pattern and choosing its priority.