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 new GraphGraph; 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#

Examples#