---
title: Write a transform
description: "Write graph transforms of your own: functions over the ModelGraph edits that optimize() runs beside the runtime's transforms, operators built with add_node, insert and replace, and rewrite rules that replace each occurrence of a Pattern."
---

{/* Every block is a program under examples/<language>/howto/write_a_transform/: the first
     block of each tab is get_a_graph whole, and every later block is its program's docs
     region. tools/tutorial_check.py runs each program against its recorded output. */}

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](edit-a-graph.mdx) 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](optimize-and-finalize.mdx#run-a-list-of-your-own) shows. It leaves one Relu where the model has two, and its row counts one application.

<LangTabs>

<LangTab value="cpp">

```cpp title="get_a_graph.cpp"
#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;
}
```

</LangTab>

<LangTab value="python">

```python title="get_a_graph.py"
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)]
```

</LangTab>

</LangTabs>

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.

<LangTabs>

<LangTab value="cpp">

```cpp title="a_transform.cpp"
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;
}
```

</LangTab>

<LangTab value="python">

```python title="a_transform.py"
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
```

</LangTab>

</LangTabs>

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.

<LangTabs>

<LangTab value="cpp">

```cpp title="order.cpp"
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;
}
```

</LangTab>

<LangTab value="python">

```python title="order.py"
# 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)]
```

</LangTab>

</LangTabs>

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.

<LangTabs>

<LangTab value="cpp">

```cpp title="refusals.cpp"
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;
}
```

</LangTab>

<LangTab value="python">

```python title="refusals.py"
# 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
```

</LangTab>

</LangTabs>

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.

<LangTabs>

<LangTab value="cpp">

```cpp title="add_node.cpp"
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;
}
```

</LangTab>

<LangTab value="python">

```python title="add_node.py"
# 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)]
```

</LangTab>

</LangTabs>

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.

<LangTabs>

<LangTab value="cpp">

```cpp title="insert_replace.cpp"
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;
}
```

</LangTab>

<LangTab value="python">

```python title="insert_replace.py"
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)]
```

</LangTab>

</LangTabs>

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.

<LangTabs>

<LangTab value="cpp">

```cpp title="rules.cpp"
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;
}
```

</LangTab>

<LangTab value="python">

```python title="rules.py"
# 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)]
```

</LangTab>

</LangTabs>

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.

<LangTabs>

<LangTab value="cpp">

```cpp title="skipped_occurrences.cpp"
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;
}
```

</LangTab>

<LangTab value="python">

```python title="skipped_occurrences.py"
# 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
```

</LangTab>

</LangTabs>

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.

<LangTabs>

<LangTab value="cpp">

```cpp title="same_api.cpp"
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;
}
```

</LangTab>

<LangTab value="python">

```python title="same_api.py"
# 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
```

</LangTab>

</LangTabs>

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.

<LangTabs>

<LangTab value="cpp">

```cpp title="after_the_transforms.cpp"
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;
}
```

</LangTab>

<LangTab value="python">

```python title="after_the_transforms.py"
# 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)
```

</LangTab>

</LangTabs>

Check a graph your transforms rewrote against a reference like this one before you serve it. [Optimize and finalize](optimize-and-finalize.mdx) covers the runtime's transforms and their defaults, [Edit a graph](edit-a-graph.mdx) the edits a transform makes, and [Query a graph](query-a-graph.mdx) the reads a transform or a condition can make.
