Write a transform
A transform rewrites a ModelGraph inside optimize(). The runtime ships its own transforms, and you can write more. A transform you write is a function over the graph edits that Edit a graph covers, and optimize() runs it beside the runtime's transforms in one list, at its place, once per iteration. Three more calls build operators inside a transform: add_node adds one operator, and insert and replace add the operators that a function over the public ops records. A rewrite rule is a transform that finds each occurrence of a Pattern and replaces it. The calls are available from C++ and Python, under the same names.
A graph to transform
The model below computes y = Relu(Relu(x)) + Neg(x), a Relu over a Relu and a Neg that the Add reads, and the examples on this page rewrite both. trace_model traces a function over an x of shape [2, 3], the model unless it is given another, and returns the graph as built. Each program prints a graph as the formula it returns. formula names a value by its operator and the values that operator reads, a graph input by its name and a constant as c, and formula_of gives the formula of the graph's output, which the node names and their order do not change. rows lists a report's rows as each transform's name and applications, and in C++ optimize_with runs a list through OptimizeOptions::transforms. drop_repeated_relu is the first transform on this page. It bypasses each Relu that reads a Relu with bypass_node, so that Relu's readers read the one before it.
The program runs remove_redundant_relu, one of the runtime's transforms, in a list of its own, as Optimize and finalize shows. It leaves one Relu where the model has two, and its row counts one application.
- C++
- Python
#include <algorithm>
#include <cstddef>
#include <cstdio>
#include <optional>
#include <string>
#include <utility>
#include <variant>
#include <vector>
#include <ClikaRT/clika_rt.h>
using ClikaRT::Error;
using ClikaRT::Result;
using ClikaRT::Tensor;
using ClikaRT::graph::ModelGraph;
using ClikaRT::graph::Node;
using ClikaRT::graph::NodeKind;
using ClikaRT::graph::OpCode;
using ClikaRT::graph::OptimizeOptions;
using ClikaRT::graph::OptimizeReport;
using ClikaRT::graph::Pattern;
using ClikaRT::graph::RewriteMatch;
using ClikaRT::graph::Transform;
using ClikaRT::graph::TransformReport;
using ClikaRT::graph::Value;
namespace ops = ClikaRT::ops;
namespace transforms = ClikaRT::graph::transforms;
namespace {
// y = Relu(Relu(x)) + Neg(x): a Relu over a Relu, and a Neg that the Add reads.
std::vector<Tensor> model(const std::vector<Tensor>& inputs) {
const Tensor rectified = ops::relu(ops::relu(inputs[0]));
const Tensor negated = ops::neg(inputs[0]);
return {ops::add(rectified, negated)};
}
// The formula a value computes: its operator over the values it reads, a graph input by its name, and a
// constant as c.
std::string formula(const Value& value) {
const std::optional<Node> producer = value.producer();
if (!producer.has_value()) return "c";
if (producer->kind() == NodeKind::Input) return producer->name();
std::string reads;
for (const Value& read : producer->inputs()) reads += (reads.empty() ? "" : ", ") + formula(read);
return std::string(ClikaRT::graph::op_code_name(producer->op_code())) + "(" + reads + ")";
}
// The formula of the value the graph returns.
std::string formula_of(const ModelGraph& graph) {
for (const Node& node : graph.nodes()) {
for (const Value& value : node.outputs()) {
if (value.is_graph_output()) return formula(value);
}
}
return "";
}
// A trace returns the graph as built: every operator as written, nothing optimized or finalized.
ModelGraph trace_model(const ClikaRT::graph::TraceFunction& fn = model) {
const std::vector<ClikaRT::spec::TensorSpec> signature = {{"x", ClikaRT::DataType::Float32, {2, 3}}};
const std::vector<std::string> outputs = {"y"};
return ClikaRT::graph::trace(fn, signature, "transform", outputs);
}
// optimize() over `list`, in order.
OptimizeReport optimize_with(ModelGraph& graph, std::vector<Transform> list) {
OptimizeOptions options;
options.transforms = std::move(list);
return graph.optimize(options);
}
// Each row of an optimize() report, as "name applications", separated by commas.
std::string rows(const OptimizeReport& report) {
std::string out;
for (const TransformReport& row : report.transforms) {
out += (out.empty() ? "" : ", ") + row.transform.name() + " " + std::to_string(row.applications);
}
return out;
}
// A Relu that reads a Relu is bypassed: its readers read the first Relu instead.
Result<void> drop_repeated_relu(ModelGraph& graph) {
for (const Node& relu : graph.find_nodes(OpCode::Relu)) {
const std::optional<Node> producer = relu.input(0)->producer();
if (producer.has_value() && producer->op_code() == OpCode::Relu) graph.bypass_node(relu);
}
return {};
}
} // namespace
int main() {
ModelGraph graph = trace_model();
std::printf("%s\n", formula_of(graph).c_str()); // Add(Relu(Relu(x)), Neg(x))
// optimize() runs a list of transforms. This one of the runtime's removes the repeated Relu.
const OptimizeReport report = optimize_with(graph, {transforms::remove_redundant_relu()});
std::printf("%s | %s\n", formula_of(graph).c_str(), rows(report).c_str());
// Add(Relu(x), Neg(x)) | remove_redundant_relu 1
return 0;
}
from collections.abc import Callable
import clika_runtime as crt
from clika_runtime.graph import NodeKind, OpCode, Pattern, RewriteMatch, Transform, transforms
def model(inputs: list[crt.Tensor]) -> list[crt.Tensor]:
x = inputs[0]
return [crt.relu(crt.relu(x)) + crt.neg(x)] # y = Relu(Relu(x)) + Neg(x)
# The formula a value computes: its operator over the values it reads, a graph input by its name, and a
# constant as c.
def formula(value: crt.graph.Value) -> str:
producer = value.producer()
if producer is None:
return "c"
if producer.kind == NodeKind.Input:
return producer.name
return f"{producer.op_code.name}({', '.join(formula(read) for read in producer.inputs)})"
# The formula of the value the graph returns.
def formula_of(graph: crt.graph.ModelGraph) -> str:
(returned,) = [value for node in graph.nodes() for value in node.outputs if value.is_graph_output()]
return formula(returned)
# A trace returns the graph as built: every operator as written, nothing optimized or finalized.
def trace_model(fn: Callable[[list[crt.Tensor]], list[crt.Tensor]] = model) -> crt.graph.ModelGraph:
return crt.trace(fn, [crt.TensorSpec("x", crt.float32, [2, 3])], output_names=["y"]).graph
# Each row of an optimize() report, as (name, applications).
def rows(report: crt.graph.OptimizeReport) -> list[tuple[str, int]]:
return [(row.transform.name, row.applications) for row in report.transforms]
# A Relu that reads a Relu is bypassed: its readers read the first Relu instead.
def drop_repeated_relu(graph: crt.graph.ModelGraph) -> None:
for relu in graph.find_nodes(OpCode.Relu):
producer = relu.input(0).producer()
if producer is not None and producer.op_code == OpCode.Relu:
graph.bypass_node(relu)
graph = trace_model()
print(formula_of(graph)) # Add(Relu(Relu(x)), Neg(x))
# optimize() runs a list of transforms. This one of the runtime's removes the repeated Relu.
report = graph.optimize([transforms.remove_redundant_relu()])
print(formula_of(graph), rows(report)) # Add(Relu(x), Neg(x)) [('remove_redundant_relu', 1)]
Write a transform when a rewrite should run inside optimize(), beside the runtime's transforms and at every iteration, so it sees what their rewrites expose and they see what it changes.
Each program on this page is complete and runs on its own. From here on, a block shows the part of its program that follows the opening lines the first block shows (the includes or imports, model, the helpers, drop_repeated_relu and, in Python, the trace).
Write a transform
Transform::from_function(name, fn) in C++ and Transform(name, fn) in Python make a transform from a function. fn gets the graph that optimize() is running and edits it with the graph edits. In C++ it returns a Result<void>, and a failed one ends the run. In Python it returns None. The @transforms.transform decorator makes a transform from the function it decorates, named after the function, and @transforms.transform(name, runs_after=...) sets the name and the order. The name is how the transform reads in the report and in errors, so it may be neither empty nor the name of one of the runtime's transforms.
optimize() runs each entry of the list at its place, once per iteration, and stops at the first iteration that changes nothing. Each entry has one row in the report, and its applications counts the iterations in which the transform changed the graph. A transform changes the graph when one of its edits does, so count_nodes below, which only reads the graph, counts none, though it runs in both iterations and records the node count each time. A handle you made equals itself and the copies the report holds.
- C++
- Python
int main() {
ModelGraph graph = trace_model();
// A transform you write: a name, and the function optimize() calls at its place in the list.
const Transform drop = Transform::from_function("drop_repeated_relu", drop_repeated_relu);
// One that reads the graph and changes nothing: it records the node count each time it runs.
std::vector<std::size_t> counts;
const Transform count_nodes = Transform::from_function("count_nodes", [&counts](ModelGraph& g) -> Result<void> {
counts.push_back(g.nodes().size());
return {};
});
const OptimizeReport report = optimize_with(graph, {count_nodes, drop});
std::printf("%s | %s\n", formula_of(graph).c_str(), rows(report).c_str());
// Add(Relu(x), Neg(x)) | count_nodes 0, drop_repeated_relu 1
std::printf("%lld | %zu %zu\n", static_cast<long long>(report.iterations), counts.at(0), counts.at(1));
// 2 | 5 4: count_nodes ran in both iterations
std::printf("%s\n", report.transforms.at(1).transform == drop ? "true" : "false"); // true
return 0;
}
drop = Transform("drop_repeated_relu", drop_repeated_relu) # a name, and the function optimize() calls
counts: list[int] = []
@transforms.transform # a transform named after the function it decorates
def count_nodes(g: crt.graph.ModelGraph) -> None:
counts.append(len(g.nodes())) # it reads the graph and changes nothing
report = graph.optimize([count_nodes, drop])
print(formula_of(graph), rows(report))
# Add(Relu(x), Neg(x)) [('count_nodes', 0), ('drop_repeated_relu', 1)]
print(report.iterations, counts) # 2 [5, 4]: count_nodes ran in both iterations
print(report.transforms[1].transform == drop) # True
Write a transform as a function when the rewrite reads more of the graph than a pattern states, or edits parts of it that a rule does not, such as the graph's inputs and outputs.
Order the transforms
runs_after lists the transforms, the runtime's or yours, that a transform must follow in any list that holds both. It is the third argument of Transform::from_function and the runs_after= keyword of Transform and of the decorator. optimize() checks the list before anything runs. A list that places a transform ahead of one it must follow is refused with the code name INVALID_ARGUMENT, and the message names both entries and the order to use. A transform the list does not hold asks nothing. Here drop_repeated_relu follows remove_double_neg, which can expose a Relu over a Relu.
- C++
- Python
int main() {
ModelGraph graph = trace_model();
// remove_double_neg can expose a Relu over a Relu (Relu(Neg(Neg(Relu(x))))), so drop_repeated_relu runs after it.
const Transform drop =
Transform::from_function("drop_repeated_relu", drop_repeated_relu, {transforms::remove_double_neg()});
try {
optimize_with(graph, {drop, transforms::remove_double_neg()}); // a list is checked before anything runs
} catch (const Error& error) {
std::printf("%s | %s\n", error.code_name().c_str(), error.what());
// INVALID_ARGUMENT | optimize: transforms[0] (drop_repeated_relu) must run after remove_double_neg, which the list places later, at transforms[1]; list remove_double_neg ahead of it
}
const OptimizeReport report = optimize_with(graph, {transforms::remove_double_neg(), drop});
std::printf("%s | %s\n", formula_of(graph).c_str(), rows(report).c_str());
// Add(Relu(x), Neg(x)) | remove_double_neg 0, drop_repeated_relu 1
return 0;
}
# remove_double_neg can expose a Relu over a Relu (Relu(Neg(Neg(Relu(x))))), so drop_repeated_relu runs after it.
drop = Transform("drop_repeated_relu", drop_repeated_relu, runs_after=[transforms.remove_double_neg()])
try:
graph.optimize([drop, transforms.remove_double_neg()]) # a list is checked before anything runs
except crt.InvalidArgumentError as error:
print(error.code_name, "|", error)
# INVALID_ARGUMENT | optimize: transforms[0] (drop_repeated_relu) must run after remove_double_neg, which the list places later, at transforms[1]; list remove_double_neg ahead of it
report = graph.optimize([transforms.remove_double_neg(), drop])
print(formula_of(graph), rows(report))
# Add(Relu(x), Neg(x)) [('remove_double_neg', 0), ('drop_repeated_relu', 1)]
Declare the order when a transform depends on another one's result, so that no list can run them the other way.
What a transform may not do
Inside a transform, its graph serves the graph edits and every query. It refuses optimize(), finalize(), to() and attach_kv_cache() with the code name FAILED_PRECONDITION, since optimize() is still working on it. The graph's inputs and outputs change only through their own edits (add_input, add_output, remove_output, rename_input and rename_output), and every other edit keeps them as they are. The graph also refuses an edit from another thread, and moving from the graph or assigning over it leaves both graphs as they were.
A transform that fails ends the run, and the graph is as optimize() found it, with the changes of every transform in the run undone. In C++ a transform fails by returning a failed Result or by throwing. optimize() raises that failure with its status and its code name, and the message starts with the entry that failed, as in optimize: transforms[1] (no_neg) failed:. In Python, optimize() raises the exception the transform raised, as that same object. The finalize refusal comes back as finalize() raised it, and an exception class of your own comes back as itself.
- C++
- Python
int main() {
ModelGraph graph = trace_model();
// Inside a transform, its own graph refuses finalize(): optimize() is still working on the graph.
const Transform finalizes = Transform::from_function("finalizes", [](ModelGraph& g) -> Result<void> {
g.finalize();
return {};
});
try {
optimize_with(graph, {finalizes});
} catch (const Error& error) {
std::printf("%s | %s\n", error.code_name().c_str(), error.what());
// FAILED_PRECONDITION | optimize: transforms[0] (finalizes) failed: finalize: the transform 'finalizes' is running on this graph inside optimize(); call finalize() after optimize() returns
}
// A failure ends the run: optimize() raises it with its status and code, and every change of the run is undone.
const Transform no_neg = Transform::from_function("no_neg", [](ModelGraph& g) -> Result<void> {
if (g.find_nodes(OpCode::Neg).empty()) return {};
return Result<void>(ClikaRT::Status::Unsupported, "the graph still computes a Neg", "NEG_NOT_SERVED");
});
try {
optimize_with(graph, {Transform::from_function("drop_repeated_relu", drop_repeated_relu), no_neg});
} catch (const Error& error) {
std::printf("%s | %s\n", error.code_name().c_str(), error.what());
// NEG_NOT_SERVED | optimize: transforms[1] (no_neg) failed: the graph still computes a Neg
}
std::printf("%s\n", formula_of(graph).c_str()); // Add(Relu(Relu(x)), Neg(x)): the bypassed Relu is back
return 0;
}
# Inside a transform, its own graph refuses finalize(): optimize() is still working on the graph.
@transforms.transform
def finalizes(g: crt.graph.ModelGraph) -> None:
g.finalize()
try:
graph.optimize([finalizes])
except crt.InvalidArgumentError as error: # the refusal finalize() raised, as it raised it
print(error.code_name, "|", error)
# FAILED_PRECONDITION | finalize: the transform 'finalizes' is running on this graph inside optimize(); call finalize() after optimize() returns
class NegNotServed(Exception):
"""The target the graph is built for runs no Neg."""
# An exception a transform raises ends the run: optimize() raises that same exception, and every change of the
# run is undone.
@transforms.transform
def no_neg(g: crt.graph.ModelGraph) -> None:
if g.find_nodes(OpCode.Neg):
raise NegNotServed("the graph still computes a Neg")
try:
graph.optimize([Transform("drop_repeated_relu", drop_repeated_relu), no_neg])
except NegNotServed as error:
print(type(error).__name__, "|", error) # NegNotServed | the graph still computes a Neg
print(formula_of(graph)) # Add(Relu(Relu(x)), Neg(x)): the Relu that drop_repeated_relu bypassed is back
Fail a transform when it meets a graph it cannot handle: the run ends, and the graph stays as optimize() found it.
Build one operator with add_node
add_node(code, inputs, attributes) adds an operator of code, built the way its public op builds it, so the op's own checks apply. inputs holds its values, one per input port in the op's order, with std::nullopt (None in Python) for an optional operand left out. attributes holds its settings under the names a node's attributes report (Node::attributes() in C++, node.attributes in Python). A setting left out takes the op's default, an integer also serves a number, and an enumerated setting takes its enumerator's name. Nothing reads the new operator's outputs until an edit wires them in, and finalize() drops a node that nothing reads by then. A code that no single public op call builds, such as an operator with several outputs or a quantized one, is refused, and insert builds it from the public ops that compute it.
Here fold_neg_into_sub turns a + Neg(b) into a - b. The Sub takes the Add's own settings, alpha and activation, which the two ops name alike. replace_all_uses_with moves the Add's reads to the Sub, the graph's output among them, and the Add and the Neg go.
- C++
- Python
int main() {
ModelGraph graph = trace_model();
// a + Neg(b) is a - b: a Sub with the Add's own settings takes the Add's place, and the Neg goes once nothing
// reads it.
const Transform fold = Transform::from_function("fold_neg_into_sub", [](ModelGraph& g) -> Result<void> {
for (const Node& add : g.find_nodes(OpCode::Add)) {
const std::optional<Node> neg = add.input(1)->producer();
if (!neg.has_value() || neg->op_code() != OpCode::Neg) continue;
// The Sub reads a and b, with the Add's settings (alpha, activation) under their own names.
const Node sub = g.add_node(OpCode::Sub, {add.input(0), neg->input(0)}, add.attributes());
g.replace_all_uses_with(*add.output(0), *sub.output(0));
g.remove_node(add);
if (neg->out_degree() == 0) g.remove_node(*neg);
}
return {};
});
const OptimizeReport report = optimize_with(graph, {fold});
std::printf("%s | %s\n", formula_of(graph).c_str(), rows(report).c_str());
// Sub(Relu(Relu(x)), x) | fold_neg_into_sub 1
return 0;
}
# a + Neg(b) is a - b: a Sub with the Add's own settings takes the Add's place, and the Neg goes once nothing
# reads it.
@transforms.transform
def fold_neg_into_sub(g: crt.graph.ModelGraph) -> None:
for add in g.find_nodes(OpCode.Add):
neg = add.input(1).producer()
if neg is None or neg.op_code != OpCode.Neg:
continue
sub = g.add_node(OpCode.Sub, [add.input(0), neg.input(0)], add.attributes) # alpha and activation
g.replace_all_uses_with(add.output(0), sub.output(0))
g.remove_node(add)
if neg.out_degree() == 0:
g.remove_node(neg)
report = graph.optimize([fold_neg_into_sub])
print(formula_of(graph), rows(report)) # Sub(Relu(Relu(x)), x) [('fold_neg_into_sub', 1)]
Build with add_node when the new operator is one call of a public op, with settings you read off the graph or choose yourself.
Splice in a function with insert and replace
insert(fn, inputs) traces fn, a function over the public ops, over one tensor per input, each with that value's dtype, device and shape, and adds the operators it records, reading the values given. In C++ fn returns a std::vector<Tensor>, and in Python a Tensor or a sequence of Tensors. insert returns the values fn returns, which nothing reads until an edit wires them in, as add_node's outputs are wired. replace(old_outputs, fn, inputs) is insert followed by the wiring and the removal: every read of old_outputs[i] reads new value i, the graph's outputs included, and the operators that computed the old outputs from the inputs go, with every producer only they read. An operator an input is computed from stays. fn computes values, so it writes no input in place and does not edit the graph.
Here two transforms put one Relu in place of a Relu over a Relu. with_insert wires the new Relu in with replace_all_uses_with and removes the pair itself, and with_replace does all of that in one call. Each rewrites one pair per call and returns, since its edits remove nodes that its loop still holds, and the next iteration finds the next pair.
- C++
- Python
int main() {
// Relu(Relu(v)) is Relu(v): a function over the public ops.
const ClikaRT::graph::TraceFunction relu_of = [](const std::vector<Tensor>& in) {
return std::vector<Tensor>{ops::relu(in[0])};
};
// insert() adds the operators relu_of records over the values given, and returns their values, which nothing
// reads until an edit wires them in.
const Transform with_insert =
Transform::from_function("with_insert", [&relu_of](ModelGraph& g) -> Result<void> {
for (const Node& outer : g.find_nodes(OpCode::Relu)) {
const std::optional<Node> inner = outer.input(0)->producer();
if (!inner.has_value() || inner->op_code() != OpCode::Relu) continue;
const std::vector<Value> made = g.insert(relu_of, {*inner->input(0)});
g.replace_all_uses_with(*outer.output(0), made.at(0));
g.remove_node(outer);
g.remove_node(*inner);
return {}; // one pair per call: the next iteration finds the next one
}
return {};
});
// replace() does all of that in one call: the operators that compute the old output from the inputs go, the
// outer Relu and the inner one only it reads.
const Transform with_replace =
Transform::from_function("with_replace", [&relu_of](ModelGraph& g) -> Result<void> {
for (const Node& outer : g.find_nodes(OpCode::Relu)) {
const std::optional<Node> inner = outer.input(0)->producer();
if (!inner.has_value() || inner->op_code() != OpCode::Relu) continue;
g.replace({*outer.output(0)}, relu_of, {*inner->input(0)});
return {};
}
return {};
});
ModelGraph graph = trace_model();
ModelGraph other = trace_model();
const std::string inserted = rows(optimize_with(graph, {with_insert}));
const std::string replaced = rows(optimize_with(other, {with_replace}));
std::printf("%s | %s\n", formula_of(graph).c_str(), inserted.c_str()); // Add(Relu(x), Neg(x)) | with_insert 1
std::printf("%s | %s\n", formula_of(other).c_str(), replaced.c_str()); // Add(Relu(x), Neg(x)) | with_replace 1
return 0;
}
def relu_of(tensors: list[crt.Tensor]) -> crt.Tensor:
return crt.relu(tensors[0]) # Relu(Relu(v)) is Relu(v): a function over the public ops
# insert() adds the operators relu_of records over the values given, and returns their values, which nothing
# reads until an edit wires them in.
@transforms.transform
def with_insert(g: crt.graph.ModelGraph) -> None:
for outer in g.find_nodes(OpCode.Relu):
inner = outer.input(0).producer()
if inner is not None and inner.op_code == OpCode.Relu:
(made,) = g.insert(relu_of, [inner.input(0)])
g.replace_all_uses_with(outer.output(0), made)
g.remove_node(outer)
g.remove_node(inner)
return # one pair per call: the next iteration finds the next one
# replace() does all of that in one call: the operators that compute the old output from the inputs go, the
# outer Relu and the inner one only it reads.
@transforms.transform
def with_replace(g: crt.graph.ModelGraph) -> None:
for outer in g.find_nodes(OpCode.Relu):
inner = outer.input(0).producer()
if inner is not None and inner.op_code == OpCode.Relu:
g.replace([outer.output(0)], relu_of, [inner.input(0)])
return
other = trace_model()
inserted = rows(graph.optimize([with_insert]))
replaced = rows(other.optimize([with_replace]))
print(formula_of(graph), inserted) # Add(Relu(x), Neg(x)) [('with_insert', 1)]
print(formula_of(other), replaced) # Add(Relu(x), Neg(x)) [('with_replace', 1)]
Use replace when the new values are a computation over values the graph already has, and insert when you wire the result in yourself.
Write a rule
A rule rewrites the occurrences of a Pattern, which names the operators to find, one pattern node each, and the edges between them. Pattern::chain (in Python, Pattern.chain or Pattern([...])) builds a straight line, and add_node and add_edge build any other shape, as the next section does. transforms::rewrite(name, pattern, replacement, condition) in C++ and transforms.rewrite(name, pattern, replacement, condition) in Python make the rule, a transform like any other. Each iteration searches the graph for the pattern once and takes the occurrences in the order the search returns them. The rule replaces each occurrence its condition accepts, every one when there is no condition, with the values replacement(match, tensors) computes from one traced tensor per value the occurrence reads, traced as replace traces its function. Every occurrence is rewritten in the same iteration, so the rule's row counts that iteration once, and the next iteration searches again, so an occurrence a replacement creates is rewritten then. In Python, the @transforms.rule(pattern, condition=..., name=..., runs_after=...) decorator makes a rule from the replacement it decorates, named after the function unless given a name.
The match is a RewriteMatch. nodes holds the matched node for each pattern node, in the order the nodes were added, with std::nullopt (None in Python) for an optional node the occurrence lacks. inputs holds the values the occurrence reads from outside it, each once, and the replacement gets one tensor per entry, in that order. outputs holds the values it makes that a node outside it reads or the graph returns, and the replacement returns one new value per entry. Here the replacement also records what it reads of each occurrence.
- C++
- Python
int main() {
// y = Relu(Relu(x)) + Relu(Relu(Neg(x))): two Relus over a Relu.
ModelGraph pairs = trace_model([](const std::vector<Tensor>& in) {
const Tensor first = ops::relu(ops::relu(in[0]));
const Tensor second = ops::relu(ops::relu(ops::neg(in[0])));
return std::vector<Tensor>{ops::add(first, second)};
});
std::printf("%s\n", formula_of(pairs).c_str()); // Add(Relu(Relu(x)), Relu(Relu(Neg(x))))
std::vector<std::string> seen; // what the replacement reads of each occurrence
const OpCode pair[] = {OpCode::Relu, OpCode::Relu};
const Transform collapse_relu = transforms::rewrite(
"collapse_relu", Pattern::chain(pair), [&seen](const RewriteMatch& match, const std::vector<Tensor>& in) {
seen.push_back(std::to_string(match.nodes.size()) + " nodes, reads " + formula(match.inputs.at(0)) +
", " + std::to_string(match.outputs.size()) + " output");
return std::vector<Tensor>{ops::relu(in[0])}; // one Relu over the value the pair reads
});
const std::string report = rows(optimize_with(pairs, {collapse_relu}));
std::sort(seen.begin(), seen.end());
for (const std::string& occurrence : seen) std::printf("%s\n", occurrence.c_str());
// 2 nodes, reads Neg(x), 1 output
// 2 nodes, reads x, 1 output
std::printf("%s | %s\n", formula_of(pairs).c_str(), report.c_str());
// Add(Relu(x), Relu(Neg(x))) | collapse_relu 1
return 0;
}
# y = Relu(Relu(x)) + Relu(Relu(Neg(x))): two Relus over a Relu.
def two_pairs(inputs: list[crt.Tensor]) -> list[crt.Tensor]:
x = inputs[0]
return [crt.relu(crt.relu(x)) + crt.relu(crt.relu(crt.neg(x)))]
seen: list[tuple[list[str], list[str], int]] = [] # what the replacement reads of each occurrence
def one_relu(match: RewriteMatch, tensors: list[crt.Tensor]) -> crt.Tensor:
nodes = [node.op_code.name for node in match.nodes]
seen.append((nodes, [formula(value) for value in match.inputs], len(match.outputs)))
return crt.relu(tensors[0]) # one Relu over the value the pair reads
collapse_relu = transforms.rewrite("collapse_relu", Pattern.chain([OpCode.Relu, OpCode.Relu]), one_relu)
pairs = trace_model(two_pairs)
print(formula_of(pairs)) # Add(Relu(Relu(x)), Relu(Relu(Neg(x))))
report = pairs.optimize([collapse_relu])
print(sorted(seen)) # [(['Relu', 'Relu'], ['Neg(x)'], 1), (['Relu', 'Relu'], ['x'], 1)]
print(formula_of(pairs), rows(report)) # Add(Relu(x), Relu(Neg(x))) [('collapse_relu', 1)]
Write a rule when the rewrite is local: a fixed arrangement of operators, a condition on the match, and a replacement computed from the values it reads.
The occurrences a rule skips
A rule skips an occurrence, and its condition never sees it, in three cases: an earlier replacement in the same iteration removed one of its nodes, the occurrence is not convex, or nothing outside it reads a value it makes. An occurrence is not convex when a path leaves it and comes back, a path no replacement can keep, and ModelGraph::is_convex answers the same question. Every other occurrence goes to condition(match), which decides whether to replace it, by the truth of its result in Python. With no condition, the rule replaces every occurrence.
Which occurrence an earlier replacement removes depends on the order the search returns them, so no program here shows that case. The program below builds its pattern node by node, the Add as the root and the Neg feeding its second operand through add_edge, and counts the condition's calls. The condition accepts an Add with its default settings whose value is the one the occurrence hands on. On the page's graph the condition is asked once and the fold runs. On n = Neg(x); y = Relu(n) + n, the path from the Neg through the Relu leaves the occurrence and comes back, so the rule skips it and the condition is never asked.
- C++
- Python
int main() {
// Add(a, Neg(b)), built node by node: the Add is the root, and the Neg feeds its second operand.
Pattern pattern;
const Pattern::NodeId add = pattern.add_node(OpCode::Add);
const Pattern::NodeId neg = pattern.add_node(OpCode::Neg);
pattern.add_edge(neg, add, 0, 1); // the Neg's output 0 feeds the Add's input 1
int asked = 0; // how many occurrences the condition is asked about
// a + Neg(b) is a - b for an Add with its default settings, when its value is the one the occurrence hands on.
const auto plain_sum = [&asked](const RewriteMatch& match) {
++asked;
const Node& total = *match.nodes[0];
const std::optional<ClikaRT::graph::Attr> alpha = total.attribute("alpha");
const double* scale = alpha.has_value() ? std::get_if<double>(&alpha->value) : nullptr;
return match.outputs.size() == 1 && scale != nullptr && *scale == 1.0 &&
total.fused_activation() == ops::Activation::Identity;
};
const Transform fold_neg_into_sub = transforms::rewrite(
"fold_neg_into_sub", std::move(pattern),
[](const RewriteMatch&, const std::vector<Tensor>& in) {
return std::vector<Tensor>{ops::sub(in[0], in[1])}; // match.inputs holds a, then b
},
plain_sum);
ModelGraph graph = trace_model();
std::string report = rows(optimize_with(graph, {fold_neg_into_sub}));
std::printf("%s | %s | asked %d\n", formula_of(graph).c_str(), report.c_str(), asked);
// Sub(Relu(Relu(x)), x) | fold_neg_into_sub 1 | asked 1
// n = Neg(x), y = Relu(n) + n: the path from the Neg through the Relu leaves the occurrence and comes back in.
asked = 0;
ModelGraph looped = trace_model([](const std::vector<Tensor>& in) {
const Tensor n = ops::neg(in[0]);
return std::vector<Tensor>{ops::add(ops::relu(n), n)};
});
report = rows(optimize_with(looped, {fold_neg_into_sub}));
std::printf("%s | %s | asked %d\n", formula_of(looped).c_str(), report.c_str(), asked);
// Add(Relu(Neg(x)), Neg(x)) | fold_neg_into_sub 0 | asked 0
return 0;
}
# Add(a, Neg(b)), built node by node: the Add is the root, and the Neg feeds its second operand.
pattern = Pattern()
add = pattern.add_node(OpCode.Add)
neg = pattern.add_node(OpCode.Neg)
pattern.add_edge(neg, add, to_port=1)
asked: list[str] = [] # the Add of each occurrence the condition is asked about
# a + Neg(b) is a - b for an Add with its default settings, when its value is the one the occurrence hands on.
def plain_sum(match: RewriteMatch) -> bool:
total = match.nodes[0]
asked.append(total.name)
return (len(match.outputs) == 1 and total.attribute("alpha") == 1
and total.attribute("activation") == "Identity")
@transforms.rule(pattern, condition=plain_sum)
def fold_neg_into_sub(match: RewriteMatch, tensors: list[crt.Tensor]) -> crt.Tensor:
return tensors[0] - tensors[1] # match.inputs holds a, then b
report = graph.optimize([fold_neg_into_sub])
print(formula_of(graph), rows(report), len(asked)) # Sub(Relu(Relu(x)), x) [('fold_neg_into_sub', 1)] 1
# n = Neg(x), y = Relu(n) + n: the path from the Neg through the Relu leaves the occurrence and comes back in.
def looped(inputs: list[crt.Tensor]) -> list[crt.Tensor]:
n = crt.neg(inputs[0])
return [crt.relu(n) + n]
asked.clear()
loop = trace_model(looped)
report = loop.optimize([fold_neg_into_sub])
print(formula_of(loop), rows(report), len(asked)) # Add(Relu(Neg(x)), Neg(x)) [('fold_neg_into_sub', 0)] 0
Write a condition for each fact the replacement relies on, such as a matched node's settings, since the rule itself checks only the three cases above.
The same API as the runtime's
A rule you write with this API can give the same graph as one of the runtime's own transforms. clamp_to_relu turns a Clamp with a zero floor and no ceiling into a Relu, and the program writes it as a rule over one Clamp node. The rule's condition reads the Clamp's settings: a floor held as the literal zero, no ceiling, and one input, since a bound given as a tensor is an input of its own. Both runs turn the first Clamp into a Relu and keep the one with a ceiling, so the two graphs return the same formula. Their node names differ, since replace names the new Relu afresh where the runtime's transform keeps the Clamp's name.
- C++
- Python
int main() {
// y = Clamp(x, 0) + Clamp(x, 0, 6): a zero floor alone, then a floor and a ceiling.
const ClikaRT::graph::TraceFunction clamps = [](const std::vector<Tensor>& in) {
const Tensor floored = ops::clamp(in[0], 0.0);
const Tensor bounded = ops::clamp(in[0], 0.0, 6.0);
return std::vector<Tensor>{ops::add(floored, bounded)};
};
// A Clamp with a zero floor held as its literal, and no ceiling, is a Relu.
const auto zero_floor = [](const RewriteMatch& match) {
const std::optional<ClikaRT::graph::Attr> lower = match.nodes[0]->attribute("min");
const std::optional<ClikaRT::graph::Attr> upper = match.nodes[0]->attribute("max");
const ClikaRT::Scalar* bound = lower.has_value() ? std::get_if<ClikaRT::Scalar>(&lower->value) : nullptr;
const bool zero = bound != nullptr && ((bound->is_float() && bound->float_value() == 0.0) ||
(bound->is_int() && bound->int_value() == 0));
const bool no_ceiling = !upper.has_value() || std::holds_alternative<std::monostate>(upper->value);
return match.inputs.size() == 1 && zero && no_ceiling;
};
Pattern one_clamp;
one_clamp.add_node(OpCode::Clamp);
const Transform clamp_to_relu_rule = transforms::rewrite(
"clamp_to_relu_rule", std::move(one_clamp),
[](const RewriteMatch&, const std::vector<Tensor>& in) { return std::vector<Tensor>{ops::relu(in[0])}; },
zero_floor);
ModelGraph by_runtime = trace_model(clamps);
ModelGraph by_rule = trace_model(clamps);
const std::string runtime_rows = rows(optimize_with(by_runtime, {transforms::clamp_to_relu()}));
const std::string rule_rows = rows(optimize_with(by_rule, {clamp_to_relu_rule}));
std::printf("%s | %s\n", formula_of(by_runtime).c_str(), runtime_rows.c_str());
// Add(Relu(x), Clamp(x)) | clamp_to_relu 1
std::printf("%s | %s\n", formula_of(by_rule).c_str(), rule_rows.c_str());
// Add(Relu(x), Clamp(x)) | clamp_to_relu_rule 1
std::printf("%s\n", formula_of(by_rule) == formula_of(by_runtime) ? "true" : "false"); // true
return 0;
}
# y = Clamp(x, 0) + Clamp(x, 0, 6): a zero floor alone, then a floor and a ceiling.
def clamps(inputs: list[crt.Tensor]) -> list[crt.Tensor]:
x = inputs[0]
return [crt.clamp(x, 0.0) + crt.clamp(x, 0.0, 6.0)]
# A Clamp with a zero floor held as its literal, and no ceiling, is a Relu.
def zero_floor(match: RewriteMatch) -> bool:
settings = match.nodes[0].attributes
return len(match.inputs) == 1 and settings.get("min") == 0 and settings.get("max") is None
clamp_to_relu_rule = transforms.rewrite("clamp_to_relu_rule", Pattern([OpCode.Clamp]),
lambda match, tensors: crt.relu(tensors[0]), zero_floor)
by_runtime, by_rule = trace_model(clamps), trace_model(clamps)
runtime_rows = rows(by_runtime.optimize([transforms.clamp_to_relu()]))
rule_rows = rows(by_rule.optimize([clamp_to_relu_rule]))
print(formula_of(by_runtime), runtime_rows) # Add(Relu(x), Clamp(x)) [('clamp_to_relu', 1)]
print(formula_of(by_rule), rule_rows) # Add(Relu(x), Clamp(x)) [('clamp_to_relu_rule', 1)]
print(formula_of(by_rule) == formula_of(by_runtime)) # True
Write your own version of one of the runtime's transforms when you need a variant of it, such as a condition of your own, and compare the graphs the two give.
Finalize and run the transformed graph
After the transforms, finalize() makes the graph runnable, and run() computes with it, as for any graph. A list of your own can start from the default list: default_transforms() returns it, and a transform added at its end runs after the runtime's. Each program checks the result against a reference it computes itself, and the Python program also imports numpy as np for that.
- C++
- Python
int main() {
ModelGraph graph = trace_model();
// The runtime's default list, with a transform of your own at its end.
std::vector<Transform> list = graph.default_transforms();
list.push_back(Transform::from_function("drop_repeated_relu", drop_repeated_relu));
optimize_with(graph, list);
graph.finalize(); // the graph runs from here on, and refuses edits
const std::vector<float> x = {1.0F, -2.0F, 3.0F, -4.0F, 5.0F, -6.0F};
const std::vector<Tensor> results = graph.run({Tensor::from_data(x.data(), {2, 3}, ClikaRT::DataType::Float32)});
const std::vector<float> y = results.front().reshape({-1}).item_as_vec<float>();
bool matches = y.size() == x.size();
for (std::size_t i = 0; matches && i < x.size(); ++i) {
matches = y[i] == (x[i] > 0.0F ? x[i] : 0.0F) - x[i]; // the reference, Relu(x) - x, by hand
}
for (std::size_t i = 0; i < y.size(); ++i) std::printf("%s%g", i == 0 ? "" : " ", static_cast<double>(y[i]));
std::printf("\n%s\n", matches ? "true" : "false");
// 0 2 0 4 0 6
// true
return 0;
}
# The runtime's default list, with a transform of your own at its end.
graph.optimize(graph.default_transforms() + [Transform("drop_repeated_relu", drop_repeated_relu)])
graph.finalize() # the graph runs from here on, and refuses edits
x = np.array([[1.0, -2.0, 3.0], [-4.0, 5.0, -6.0]], np.float32) # numpy as the data entry
(y,) = graph.run([crt.tensor(x)])
print(y.numpy().tolist()) # [[0.0, 2.0, 0.0], [4.0, 0.0, 6.0]] (numpy as the data exit)
print(np.array_equal(y.numpy(), np.maximum(x, 0) - x)) # True (numpy states the reference)
Check a graph your transforms rewrote against a reference like this one before you serve it. Optimize and finalize covers the runtime's transforms and their defaults, Edit a graph the edits a transform makes, and Query a graph the reads a transform or a condition can make.