---
title: Merge and split
description: "Join ModelGraphs into one graph that runs them all, side by side, with the outputs of one feeding the inputs of another, or across two devices; cut a graph into parts at its cheapest cut or by a function, extract the operators between values, and merge the parts back."
---

{/* Every block is a program under examples/<language>/howto/merge_and_split/: the first
     block of each tab is get_the_graphs whole, and every later block is its program's docs
     region. tools/tutorial_check.py runs each program against its recorded output, and it
     reports across_devices as blocked on a host with no accelerator. */}

`merge` joins model graphs into one graph that runs them all. The graphs sit side by side, or the outputs of one feed the inputs of another, on one device or across two. `split` cuts one graph into parts that `merge` joins back, at the cheapest cut through its values or by a function that names each operator's part, and `extract` takes out the operators that compute some values from others. Each call consumes the graphs it is given and returns new ones. The calls are available from C++ and Python, under the same names.

## Graphs to merge and split

A trace returns the graph as built, with every operator as written and nothing optimized or finalized, so it is ready to merge or split. The examples on this page use three models over an input of shape [2, 3]. `rectify` computes `h = Relu(x)`, `twice` computes `y = x + x`, and `pool` negates the sum of each row of `Relu(x)`, which pools the [2, 3] input to [2, 1]. `trace_graph` traces a model under the input and output names a section needs. `label` names a node by its input name, or by its operator's code after the part its name starts with, so a merged graph reads `encoder/Relu`. `labels` lists a graph's nodes, and `io` its input names and then its output names.

This program runs each graph once. A finalized graph runs, and merge and split refuse it, so the other programs merge and split their graphs before `finalize()`.

<LangTabs>

<LangTab value="cpp">

```cpp title="get_the_graphs.cpp"
#include <cstdio>
#include <string>
#include <vector>

#include <ClikaRT/clika_rt.h>

using ClikaRT::DataType;
using ClikaRT::Error;
using ClikaRT::Tensor;
using ClikaRT::graph::Connection;
using ClikaRT::graph::GraphPart;
using ClikaRT::graph::ModelGraph;
using ClikaRT::graph::NameKind;
using ClikaRT::graph::Node;
using ClikaRT::graph::NodeKind;
using ClikaRT::graph::Rename;
namespace ops = ClikaRT::ops;

namespace {

// h = Relu(x).
std::vector<Tensor> rectify(const std::vector<Tensor>& inputs) { return {ops::relu(inputs[0])}; }

// y = x + x.
std::vector<Tensor> twice(const std::vector<Tensor>& inputs) { return {ops::add(inputs[0], inputs[0])}; }

// y = -(the sum of each row of Relu(x)): the [2, 3] input pools to [2, 1].
std::vector<Tensor> pool(const std::vector<Tensor>& inputs) {
    return {ops::neg(ops::sum(ops::relu(inputs[0]), {1}, true))};
}

// A trace of `model` over one float32 [2, 3] input, named `input`, with its output named `output`.
// A trace returns the graph as built: every operator as written, nothing optimized or finalized.
ModelGraph trace_graph(const ClikaRT::graph::TraceFunction& model, const std::string& input,
                       const std::string& output) {
    const std::vector<ClikaRT::spec::TensorSpec> signature = {{input, DataType::Float32, {2, 3}}};
    const std::vector<std::string> outputs = {output};
    return ClikaRT::graph::trace(model, signature, "merge_and_split", outputs);
}

// A node's label: an input's name, or an operator's code after the part its name starts with.
std::string label(const Node& node) {
    if (node.kind() == NodeKind::Input) return node.name();
    const std::string name = node.name();
    const std::size_t slash = name.find('/');
    const std::string part = slash == std::string::npos ? std::string() : name.substr(0, slash + 1);
    return part + std::string(ClikaRT::graph::op_code_name(node.op_code()));
}

// The labels of a graph's nodes, in order, separated by spaces.
std::string labels(const ModelGraph& graph) {
    std::string out;
    for (const Node& node : graph.nodes()) out += (out.empty() ? "" : " ") + label(node);
    return out;
}

// Names separated by spaces.
std::string names(const std::vector<std::string>& list) {
    std::string out;
    for (const std::string& name : list) out += (out.empty() ? "" : " ") + name;
    return out;
}

// A graph's input names, then its output names.
std::string io(const ModelGraph& graph) {
    return names(graph.input_names()) + " -> " + names(graph.output_names());
}

// What a rename names.
const char* kind_name(NameKind kind) {
    switch (kind) {
        case NameKind::Node: return "node";
        case NameKind::Input: return "input";
        case NameKind::Output: return "output";
    }
    return "node";
}

// The input every program runs, [2, 3].
Tensor input() {
    const std::vector<float> x = {1.0F, -2.0F, 3.0F, -4.0F, 5.0F, -6.0F};
    return Tensor::from_data(x.data(), {2, 3}, DataType::Float32);
}

// A tensor's values, flattened and separated by spaces.
std::string values(const Tensor& tensor) {
    std::string out;
    for (const float value : tensor.reshape({-1}).item_as_vec<float>()) {
        char text[32];
        std::snprintf(text, sizeof(text), "%g", static_cast<double>(value));
        out += (out.empty() ? "" : " ") + std::string(text);
    }
    return out;
}

}  // namespace

int main() {
    ModelGraph encoder = trace_graph(rectify, "x", "h");
    ModelGraph head = trace_graph(twice, "h", "y");
    ModelGraph pooled = trace_graph(pool, "x", "y");
    for (ModelGraph* graph : {&encoder, &head, &pooled}) {
        std::printf("%s | %s\n", labels(*graph).c_str(), io(*graph).c_str());
        graph->finalize();   // a finalized graph runs, and merge and split refuse it
        std::printf("%s\n", values(graph->run({input()}).front()).c_str());
    }
    // x Relu | x -> h
    // 1 0 3 0 5 0
    // h Add | h -> y
    // 2 -4 6 -8 10 -12
    // x Relu Sum Neg | x -> y
    // -4 -5
    return 0;
}
```

</LangTab>

<LangTab value="python">

```python title="get_the_graphs.py"
from collections.abc import Callable

import clika_runtime as crt
from clika_runtime.graph import ModelGraph, NameKind, Node, NodeKind

Model = Callable[[list[crt.Tensor]], list[crt.Tensor]]


def rectify(inputs: list[crt.Tensor]) -> list[crt.Tensor]:
    return [crt.relu(inputs[0])]   # h = Relu(x)


def twice(inputs: list[crt.Tensor]) -> list[crt.Tensor]:
    return [crt.add(inputs[0], inputs[0])]   # y = x + x


def pool(inputs: list[crt.Tensor]) -> list[crt.Tensor]:
    # y = -(the sum of each row of Relu(x)): the [2, 3] input pools to [2, 1].
    return [crt.neg(crt.sum(crt.relu(inputs[0]), [1], True))]


def trace_graph(model: Model, input_name: str, output_name: str) -> ModelGraph:
    # A trace of `model` over one float32 [2, 3] input, named `input_name`, with its output named `output_name`.
    # A trace returns the graph as built: every operator as written, nothing optimized or finalized.
    signature = [crt.TensorSpec(input_name, crt.float32, [2, 3])]
    return crt.trace(model, signature, output_names=[output_name]).graph


def label(node: Node) -> str:
    # A node's label: an input's name, or an operator's code after the part its name starts with.
    if node.kind == NodeKind.Input:
        return node.name
    part, slash, _ = node.name.partition("/")
    return (part + slash if slash else "") + node.op_code.name


def labels(graph: ModelGraph) -> str:
    return " ".join(label(node) for node in graph.nodes())


def io(graph: ModelGraph) -> str:
    # A graph's input names, then its output names.
    return " ".join(graph.input_names()) + " -> " + " ".join(graph.output_names())


def kind_name(kind: NameKind) -> str:
    # What a rename names: a node, an input or an output.
    return kind.name.lower()


def x() -> crt.Tensor:
    # The input every program runs, [2, 3].
    return crt.tensor([[1.0, -2.0, 3.0], [-4.0, 5.0, -6.0]], dtype=crt.float32)


def values(tensor: crt.Tensor) -> str:
    # A tensor's values, flattened and separated by spaces.
    return " ".join(f"{value:g}" for value in crt.to(tensor, "cpu").reshape(-1).tolist())


encoder = trace_graph(rectify, "x", "h")
head = trace_graph(twice, "h", "y")
pooled = trace_graph(pool, "x", "y")
for graph in (encoder, head, pooled):
    print(labels(graph), "|", io(graph))
    graph.finalize()   # a finalized graph runs, and merge and split refuse it
    print(values(graph.run([x()])[0]))
# x Relu | x -> h
# 1 0 3 0 5 0
# h Add | h -> y
# 2 -4 6 -8 10 -12
# x Relu Sum Neg | x -> y
# -4 -5
```

</LangTab>

</LangTabs>

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, the three models and the helpers).

## Merge graphs side by side

`graph::merge(parts)` takes a list of `GraphPart`s, each a graph and the name the merged graph knows it by. Parts that no connection joins sit side by side. The merged graph takes each part's inputs and outputs, parts in the order given, and names every operator of part `p` as `p/<name>`. `MergeOptions::io_names` decides the names of the inputs and outputs. Under `IoNames::PrefixOnCollision`, the default, a name two parts both use becomes `<part>/<name>` and every other name stays as it is. `IoNames::PrefixAlways` prefixes every name, and `IoNames::Keep` keeps every name and refuses a name two parts both use. `MergedGraph::renames` lists each name the merge changed as a `Rename`, which holds its part, whether it names a node, an input or an output, and the name before and after. Here the encoder and the head both read an input named `x`, so the default prefixes the two inputs and keeps the two outputs.

In Python, `graph.merge` takes the parts as a dict from name to graph, or as a list of `(name, graph)` pairs, and returns `(merged, renames)`. `io_names` takes the mode's name, such as `"prefix_always"`, and a `Rename`'s fields are `part`, `kind`, `from_` and `to`.

<LangTabs>

<LangTab value="cpp">

```cpp title="side_by_side.cpp"
int main() {
    // An encoder and a head that both read an input named x.
    const auto two_parts = [] {
        std::vector<GraphPart> parts;
        parts.push_back({"encoder", trace_graph(rectify, "x", "h")});
        parts.push_back({"head", trace_graph(twice, "x", "y")});
        return parts;
    };

    // IoNames::PrefixOnCollision, the default: a name both parts use takes its part's name as a prefix.
    std::vector<GraphPart> parts = two_parts();
    const ClikaRT::graph::MergedGraph merged = ClikaRT::graph::merge(parts);
    std::printf("%s | %s\n", labels(merged.graph).c_str(), io(merged.graph).c_str());
    // encoder/x head/x encoder/Relu head/Add | encoder/x head/x -> h y
    for (const Rename& rename : merged.renames) {   // every name the merge changed
        std::printf("%s %s %s -> %s\n", rename.part.c_str(), kind_name(rename.kind), rename.from.c_str(),
                    rename.to.c_str());
    }
    // encoder input x -> encoder/x
    // head input x -> head/x
    // encoder node Relu_1 -> encoder/Relu_1
    // head node Add_1 -> head/Add_1

    // IoNames::PrefixAlways: every input and output name takes the prefix.
    ClikaRT::graph::MergeOptions always;
    always.io_names = ClikaRT::graph::IoNames::PrefixAlways;
    parts = two_parts();
    std::printf("%s\n", io(ClikaRT::graph::merge(parts, {}, always).graph).c_str());
    // encoder/x head/x -> encoder/h head/y

    // IoNames::Keep keeps every name, so it refuses a name both parts use.
    ClikaRT::graph::MergeOptions keep;
    keep.io_names = ClikaRT::graph::IoNames::Keep;
    parts = two_parts();
    try {
        ClikaRT::graph::merge(parts, {}, keep);
    } catch (const Error& error) {
        std::printf("%s | %s\n", error.code_name().c_str(), error.what());
    }
    // INVALID_ARGUMENT | merge: parts 'encoder' and 'head' both have an input named 'x', which IoNames::Keep
    // keeps; merge with IoNames::PrefixOnCollision or PrefixAlways
    std::printf("%s | %s\n", io(parts[0].graph).c_str(), io(parts[1].graph).c_str());
    // x -> h | x -> y: a refused merge leaves every part as it was
    return 0;
}
```

</LangTab>

<LangTab value="python">

```python title="side_by_side.py"
def two_parts() -> list[tuple[str, ModelGraph]]:
    # An encoder and a head that both read an input named x.
    return [("encoder", trace_graph(rectify, "x", "h")), ("head", trace_graph(twice, "x", "y"))]


# io_names="prefix_on_collision", the default: a name both parts use takes its part's name as a prefix.
merged, renames = crt.graph.merge(two_parts())
print(labels(merged), "|", io(merged))
# encoder/x head/x encoder/Relu head/Add | encoder/x head/x -> h y
for rename in renames:   # every name the merge changed
    print(rename.part, kind_name(rename.kind), rename.from_, "->", rename.to)
# encoder input x -> encoder/x
# head input x -> head/x
# encoder node Relu_1 -> encoder/Relu_1
# head node Add_1 -> head/Add_1

# io_names="prefix_always": every input and output name takes the prefix.
always, _ = crt.graph.merge(two_parts(), io_names="prefix_always")
print(io(always))   # encoder/x head/x -> encoder/h head/y

# io_names="keep" keeps every name, so it refuses a name both parts use.
parts = two_parts()
try:
    crt.graph.merge(parts, io_names="keep")
except crt.InvalidArgumentError as error:
    print(error.code_name, "|", error)
# INVALID_ARGUMENT | merge: parts 'encoder' and 'head' both have an input named 'x', which IoNames::Keep
# keeps; merge with IoNames::PrefixOnCollision or PrefixAlways
print(io(parts[0][1]), "|", io(parts[1][1]))   # x -> h | x -> y: a refused merge leaves every part as it was
```

</LangTab>

</LangTabs>

Merge graphs side by side to run several models as one graph, with one `finalize()` and one `run()`.

## Feed one graph into another

A `Connection` feeds an output of one part into an input of another. `from` names the part and its output, `to` names the part and its input, and a connected output or input is not one of the merged graph's. Here the encoder's `h` feeds the head's `h`, so the merged graph reads `x`, returns `y`, and computes what running the two graphs one after the other computes. The two ends of a connection have one dtype and one rank, and every size both of them fix is the same. In Python, a connection is a pair of strings, `("encoder.h", "head.h")`.

<LangTabs>

<LangTab value="cpp">

```cpp title="pipeline.cpp"
int main() {
    std::vector<GraphPart> parts;
    parts.push_back({"encoder", trace_graph(rectify, "x", "h")});
    parts.push_back({"head", trace_graph(twice, "h", "y")});
    // The encoder's output h feeds the head's input h, and neither is the merged graph's any longer.
    const std::vector<Connection> connections = {{{"encoder", "h"}, {"head", "h"}}};
    ClikaRT::graph::MergedGraph merged = ClikaRT::graph::merge(parts, connections);
    std::printf("%s | %s\n", labels(merged.graph).c_str(), io(merged.graph).c_str());
    // x encoder/Relu head/Add | x -> y

    merged.graph.optimize();   // optional: the graph optimizer runs over the merged graph
    merged.graph.finalize();
    // The same input through the two graphs one after the other, for comparison.
    ModelGraph encoder = trace_graph(rectify, "x", "h");
    ModelGraph head = trace_graph(twice, "h", "y");
    encoder.finalize();
    head.finalize();
    std::printf("%s | %s\n", values(merged.graph.run({input()}).front()).c_str(),
                values(head.run(encoder.run({input()})).front()).c_str());
    // 2 0 6 0 10 0 | 2 0 6 0 10 0
    return 0;
}
```

</LangTab>

<LangTab value="python">

```python title="pipeline.py"
# The encoder's output h feeds the head's input h, and neither is the merged graph's any longer.
parts = {"encoder": trace_graph(rectify, "x", "h"), "head": trace_graph(twice, "h", "y")}
merged, _ = crt.graph.merge(parts, [("encoder.h", "head.h")])
print(labels(merged), "|", io(merged))   # x encoder/Relu head/Add | x -> y

merged.optimize()   # optional: the graph optimizer runs over the merged graph
merged.finalize()
# The same input through the two graphs one after the other, for comparison.
encoder = trace_graph(rectify, "x", "h")
head = trace_graph(twice, "h", "y")
encoder.finalize()
head.finalize()
print(values(merged.run([x()])[0]), "|", values(head.run(encoder.run([x()]))[0]))   # 2 0 6 0 10 0 | 2 0 6 0 10 0
```

</LangTab>

</LangTabs>

Feed one graph into another to serve a model with its preprocessing or its post-processing as one graph.

## Merge a second graph into this one

`ModelGraph::merge(second, io_map)` merges `second` into the graph it is called on. Each `io_map` entry feeds one of this graph's outputs into one of `second`'s inputs. This graph keeps every name, and a name of `second` that this graph already uses takes the first free `<name>_<n>`. The call returns those renames, each under the part `second`, and it consumes `second`. Everything else follows `graph::merge`, and a connection between two devices moves its value. Here both graphs are traces of `rectify`, so the second Relu takes a new name. In Python, `first.merge(second, [("h", "x")])` returns the renames.

<LangTabs>

<LangTab value="cpp">

```cpp title="merge_into.cpp"
int main() {
    ModelGraph first = trace_graph(rectify, "x", "h");
    ModelGraph second = trace_graph(rectify, "x", "h");
    // first's output h feeds second's input x. first keeps every name, and a name of second that
    // first already uses takes the first free <name>_<n>.
    const std::vector<Rename> renames = first.merge(second, {{"h", "x"}});
    std::printf("%s | %s\n", labels(first).c_str(), io(first).c_str());   // x Relu Relu | x -> h
    for (const Rename& rename : renames) {
        std::printf("%s %s %s -> %s\n", rename.part.c_str(), kind_name(rename.kind), rename.from.c_str(),
                    rename.to.c_str());
    }
    // second node Relu_1 -> Relu_1_1
    std::printf("%zu %zu\n", second.input_names().size(), second.output_names().size());   // 0 0: second is consumed

    first.finalize();
    std::printf("%s\n", values(first.run({input()}).front()).c_str());   // 1 0 3 0 5 0: Relu(Relu(x))
    return 0;
}
```

</LangTab>

<LangTab value="python">

```python title="merge_into.py"
first = trace_graph(rectify, "x", "h")
second = trace_graph(rectify, "x", "h")
# first's output h feeds second's input x. first keeps every name, and a name of second that
# first already uses takes the first free <name>_<n>.
renames = first.merge(second, [("h", "x")])
print(labels(first), "|", io(first))   # x Relu Relu | x -> h
for rename in renames:
    print(rename.part, kind_name(rename.kind), rename.from_, "->", rename.to)
# second node Relu_1 -> Relu_1_1
print(len(second.input_names()), len(second.output_names()))   # 0 0: second is consumed

first.finalize()
print(values(first.run([x()])[0]))   # 1 0 3 0 5 0: Relu(Relu(x))
```

</LangTab>

</LangTabs>

Merge a second graph into a graph to grow it in place while every name it has stays as it is.

## Connect graphs on two devices

A part's values sit on the device `to()` sent the part to, and otherwise where they are computed. When a connection's two ends sit on two devices, `CrossDevice::Move`, the default, moves the value onto the input's device, and `CrossDevice::Refuse` refuses the merge, naming the connection and both devices. When the parts do not share one device, the merge places each part on its device before it builds, and the merged graph keeps every part where it was placed. Here the head serves on the accelerator that `Device::gpu()` finds, so the merged graph computes the Relu on the CPU, moves `h` to the accelerator, and returns `y` there. On a host with no accelerator, the program prints `BLOCKED: this section needs an accelerator` and exits with status 3. In Python, the option is `cross_device="refuse"`.

<LangTabs>

<LangTab value="cpp">

```cpp title="across_devices.cpp"
int main() {
    const ClikaRT::Device accelerator = ClikaRT::Device::gpu();   // the CPU on a host with no accelerator
    if (accelerator.is_cpu()) {
        std::printf("BLOCKED: this section needs an accelerator\n");
        return 3;
    }
    // The encoder computes on the CPU, and the head serves on the accelerator.
    const auto two_parts = [&accelerator] {
        std::vector<GraphPart> parts;
        parts.push_back({"encoder", trace_graph(rectify, "x", "h")});
        parts.push_back({"head", trace_graph(twice, "h", "y")});
        parts[1].graph.to(accelerator);
        return parts;
    };
    const std::vector<Connection> connections = {{{"encoder", "h"}, {"head", "h"}}};

    // CrossDevice::Move, the default: the merged graph moves h onto the head's device.
    std::vector<GraphPart> parts = two_parts();
    ClikaRT::graph::MergedGraph merged = ClikaRT::graph::merge(parts, connections);
    merged.graph.finalize();
    const Tensor y = merged.graph.run({input()}).front();
    std::printf("%s %s\n", y.device() == accelerator ? "true" : "false",
                values(y.to(ClikaRT::Device::cpu())).c_str());   // true 2 0 6 0 10 0: y is on the accelerator

    // CrossDevice::Refuse: a connection between two devices refuses the merge.
    ClikaRT::graph::MergeOptions refuse;
    refuse.cross_device = ClikaRT::graph::CrossDevice::Refuse;
    parts = two_parts();
    try {
        ClikaRT::graph::merge(parts, connections, refuse);
    } catch (const Error& error) {
        std::printf("%s\n", error.code_name().c_str());   // INVALID_ARGUMENT, naming the connection and its devices
    }
    std::printf("%s | %s\n", io(parts[0].graph).c_str(), io(parts[1].graph).c_str());
    // x -> h | h -> y: a refused merge leaves every part as it was
    return 0;
}
```

</LangTab>

<LangTab value="python">

```python title="across_devices.py"
accelerator = crt.Device.gpu()   # the CPU on a host with no accelerator
if accelerator.type == "cpu":
    print("BLOCKED: this section needs an accelerator")
    raise SystemExit(3)


def two_parts() -> dict[str, ModelGraph]:
    # The encoder computes on the CPU, and the head serves on the accelerator.
    parts = {"encoder": trace_graph(rectify, "x", "h"), "head": trace_graph(twice, "h", "y")}
    parts["head"].to(accelerator)
    return parts


connections = [("encoder.h", "head.h")]

# cross_device="move", the default: the merged graph moves h onto the head's device.
merged, _ = crt.graph.merge(two_parts(), connections)
merged.finalize()
y = merged.run([x()])[0]
print(y.device == accelerator, values(y))   # True 2 0 6 0 10 0: y is on the accelerator

# cross_device="refuse": a connection between two devices refuses the merge.
parts = two_parts()
try:
    crt.graph.merge(parts, connections, cross_device="refuse")
except crt.InvalidArgumentError as error:
    print(error.code_name)   # INVALID_ARGUMENT, naming the connection and its devices
print(io(parts["encoder"]), "|", io(parts["head"]))
# x -> h | h -> y: a refused merge leaves every part as it was
```

</LangTab>

</LangTabs>

Connect graphs on two devices to keep one part on the CPU while the model runs on an accelerator.

## Split a graph at its cheapest cut

`min_value_cut()` finds the cut through a graph's values that hands the fewest bytes from the nodes before it to the nodes after it, as [The cheapest cut](query-a-graph.mdx#the-cheapest-cut) on the Query a graph page shows. `graph::split(graph, cut)` cuts the graph there into two parts. The part `before` holds the operators that produce the cut values and every operator they depend on, and the part `after` holds every other operator. Connections carry the cut values `after` reads and any graph input both parts read, and a value with no name of its own crosses under the name `<node>_<port>`. Here the cheapest values are the [2, 1] sums in `pool`, 8 bytes against 24 for `x`, so `before` computes the sums and `after` negates them. The split consumes the graph. In Python, `graph.split(graph, cut)` returns `(parts, connections)`, the parts as a dict from name to graph and the connections in the form `graph.merge` takes.

<LangTabs>

<LangTab value="cpp">

```cpp title="split_at_a_cut.cpp"
int main() {
    ModelGraph pooled = trace_graph(pool, "x", "y");
    // The cheapest cut: the values the nodes before it hand to the nodes after it, and their bytes.
    const ClikaRT::graph::ValueCut cut = pooled.min_value_cut();
    std::string crossing;
    for (const ClikaRT::graph::Value& value : cut.values) {
        crossing += (crossing.empty() ? "" : " ") + label(*value.producer());
    }
    std::printf("%s %llu\n", crossing.c_str(), static_cast<unsigned long long>(cut.bytes));
    // Sum 8: the [2, 1] float32 sums cross, the cheapest value to hand on

    ClikaRT::graph::SplitGraph split = ClikaRT::graph::split(pooled, cut);
    for (const GraphPart& part : split.parts) {
        std::printf("%s: %s | %s\n", part.name.c_str(), labels(part.graph).c_str(), io(part.graph).c_str());
    }
    // before: x Relu Sum | x -> Sum_2_0
    // after: Sum_2_0 Neg | Sum_2_0 -> y
    for (const Connection& connection : split.connections) {
        std::printf("%s.%s -> %s.%s\n", connection.from.part.c_str(), connection.from.name.c_str(),
                    connection.to.part.c_str(), connection.to.name.c_str());
    }
    // before.Sum_2_0 -> after.Sum_2_0: a value with no name of its own is named <node>_<port>
    std::printf("%zu\n", pooled.input_names().size());   // 0: the split consumed the graph
    return 0;
}
```

</LangTab>

<LangTab value="python">

```python title="split_at_a_cut.py"
pooled = trace_graph(pool, "x", "y")
# The cheapest cut: the values the nodes before it hand to the nodes after it, and their bytes.
cut = pooled.min_value_cut()
print(" ".join(label(value.producer()) for value in cut.values), cut.bytes)
# Sum 8: the [2, 1] float32 sums cross, the cheapest value to hand on

parts, connections = crt.graph.split(pooled, cut)
for name, part in parts.items():
    print(f"{name}:", labels(part), "|", io(part))
# before: x Relu Sum | x -> Sum_2_0
# after: Sum_2_0 Neg | Sum_2_0 -> y
for source, target in connections:
    print(source, "->", target)
# before.Sum_2_0 -> after.Sum_2_0: a value with no name of its own is named <node>_<port>
print(len(pooled.input_names()))   # 0: the split consumed the graph
```

</LangTab>

</LangTabs>

Split a graph at its cheapest cut to run its two halves on two devices, or in two processes, with the least data between them.

## Split a graph by a function

`graph::split(graph, partition)` asks a function you write for each operator's part, a name that is not empty and holds no `/` and no `.`. It asks once per operator, in `nodes()` order, and the parts come in an order every connection runs forward in. A value read in a part other than its producer's becomes an output of the producer's part and an input of each part that reads it, under one name. A graph input goes with the first part that reads it, and every later reader receives it through a connection. Here the function names each operator's part by its code. In Python, the function takes a `Node` and returns a `str`.

<LangTabs>

<LangTab value="cpp">

```cpp title="split_by_a_function.cpp"
int main() {
    ModelGraph pooled = trace_graph(pool, "x", "y");
    // Each operator's part, by its code. split asks once per operator, in nodes() order.
    const auto by_code = [](const Node& node) {
        if (node.op_code() == ClikaRT::graph::OpCode::Relu) return std::string("rectify");
        if (node.op_code() == ClikaRT::graph::OpCode::Sum) return std::string("pool");
        return std::string("negate");
    };
    ClikaRT::graph::SplitGraph split = ClikaRT::graph::split(pooled, by_code);
    for (const GraphPart& part : split.parts) {   // in an order every connection runs forward in
        std::printf("%s: %s | %s\n", part.name.c_str(), labels(part.graph).c_str(), io(part.graph).c_str());
    }
    // rectify: x Relu | x -> Relu_1_0
    // pool: Relu_1_0 Sum | Relu_1_0 -> Sum_2_0
    // negate: Sum_2_0 Neg | Sum_2_0 -> y
    for (const Connection& connection : split.connections) {
        std::printf("%s.%s -> %s.%s\n", connection.from.part.c_str(), connection.from.name.c_str(),
                    connection.to.part.c_str(), connection.to.name.c_str());
    }
    // rectify.Relu_1_0 -> pool.Relu_1_0
    // pool.Sum_2_0 -> negate.Sum_2_0
    return 0;
}
```

</LangTab>

<LangTab value="python">

```python title="split_by_a_function.py"
def by_code(node: Node) -> str:
    # Each operator's part, by its code. split asks once per operator, in nodes() order.
    if node.op_code == crt.graph.OpCode.Relu:
        return "rectify"
    if node.op_code == crt.graph.OpCode.Sum:
        return "pool"
    return "negate"


pooled = trace_graph(pool, "x", "y")
parts, connections = crt.graph.split(pooled, by_code)
for name, part in parts.items():   # in an order every connection runs forward in
    print(f"{name}:", labels(part), "|", io(part))
# rectify: x Relu | x -> Relu_1_0
# pool: Relu_1_0 Sum | Relu_1_0 -> Sum_2_0
# negate: Sum_2_0 Neg | Sum_2_0 -> y
for source, target in connections:
    print(source, "->", target)
# rectify.Relu_1_0 -> pool.Relu_1_0
# pool.Sum_2_0 -> negate.Sum_2_0
```

</LangTab>

</LangTabs>

Split a graph by a function to give the operators you choose, by code, by name or by position, a graph of their own.

## Extract the operators between values

`graph::extract(graph, inputs, outputs)` returns the operators that compute `outputs` from `inputs` as a graph of their own. A graph input keeps its name, and any other input value becomes an input named after it. An output keeps its graph output name when it has one, and is otherwise named after its value. The extract consumes the graph and frees the weights of the operators it leaves out. Weights stored inside an ONNX file, rather than as external data, share one buffer, which the extracted graph keeps in memory as the graph did. In Python, `graph.extract(graph, inputs, outputs)` takes lists of `Value`s.

<LangTabs>

<LangTab value="cpp">

```cpp title="extract.cpp"
int main() {
    ModelGraph pooled = trace_graph(pool, "x", "y");
    // The operators between x and the sums: the Relu and the Sum, not the Neg after them.
    const Node sum = pooled.find_nodes(ClikaRT::graph::OpCode::Sum).front();
    const std::vector<ClikaRT::graph::Value> inputs = {*pooled.node("x")->output(0)};
    const std::vector<ClikaRT::graph::Value> outputs = {*sum.output(0)};
    ModelGraph pooling = ClikaRT::graph::extract(pooled, inputs, outputs);
    std::printf("%s | %s\n", labels(pooling).c_str(), io(pooling).c_str());
    // x Relu Sum | x -> Sum_2_0: a graph input keeps its name, and an output with none is named after its value
    std::printf("%zu\n", pooled.input_names().size());   // 0: the extract consumed the graph

    pooling.finalize();
    std::printf("%s\n", values(pooling.run({input()}).front()).c_str());   // 4 5: each row's sum of Relu(x)
    return 0;
}
```

</LangTab>

<LangTab value="python">

```python title="extract.py"
pooled = trace_graph(pool, "x", "y")
# The operators between x and the sums: the Relu and the Sum, not the Neg after them.
(sums,) = pooled.find_nodes(crt.graph.OpCode.Sum)
pooling = crt.graph.extract(pooled, [pooled.node("x").output(0)], [sums.output(0)])
print(labels(pooling), "|", io(pooling))
# x Relu Sum | x -> Sum_2_0: a graph input keeps its name, and an output with none is named after its value
print(len(pooled.input_names()))   # 0: the extract consumed the graph

pooling.finalize()
print(values(pooling.run([x()])[0]))   # 4 5: each row's sum of Relu(x)
```

</LangTab>

</LangTabs>

Extract a block of a model to run it alone, to test it, or to serve it on its own.

## Merge the parts back

`graph::merge(split.parts, split.connections)` joins the parts of a split back into one graph, with the source's inputs and outputs by name and its operators named `<part>/<name>`. The merged graph computes what the source computes. In Python, `graph.merge(*graph.split(graph, cut))` does the same in one call.

<LangTabs>

<LangTab value="cpp">

```cpp title="round_trip.cpp"
int main() {
    ModelGraph pooled = trace_graph(pool, "x", "y");
    ClikaRT::graph::SplitGraph split = ClikaRT::graph::split(pooled, pooled.min_value_cut());
    // The parts merge back into the source's inputs and outputs, by name, with its operators named <part>/<name>.
    ClikaRT::graph::MergedGraph merged = ClikaRT::graph::merge(split.parts, split.connections);
    std::printf("%s | %s\n", labels(merged.graph).c_str(), io(merged.graph).c_str());
    // x before/Relu before/Sum after/Neg | x -> y

    merged.graph.finalize();
    ModelGraph traced = trace_graph(pool, "x", "y");   // the same model traced again, for comparison
    traced.finalize();
    std::printf("%s | %s\n", values(merged.graph.run({input()}).front()).c_str(),
                values(traced.run({input()}).front()).c_str());   // -4 -5 | -4 -5
    return 0;
}
```

</LangTab>

<LangTab value="python">

```python title="round_trip.py"
pooled = trace_graph(pool, "x", "y")
# split returns the parts and connections merge takes, so the parts merge back into the source's inputs and
# outputs, by name, with its operators named <part>/<name>.
merged, _ = crt.graph.merge(*crt.graph.split(pooled, pooled.min_value_cut()))
print(labels(merged), "|", io(merged))   # x before/Relu before/Sum after/Neg | x -> y

merged.finalize()
traced = trace_graph(pool, "x", "y")   # the same model traced again, for comparison
traced.finalize()
print(values(merged.run([x()])[0]), "|", values(traced.run([x()])[0]))   # -4 -5 | -4 -5
```

</LangTab>

</LangTabs>

Merge the parts back after you place, optimize or edit them one at a time.

## What the calls consume, and what a refusal leaves

A merge consumes every part's graph, and a split or an extract consumes its source. On success those graphs are empty, and a `Node`, `Value` or `Edge` view taken into one of them refuses with the code name `INVALID_ARGUMENT`. Each call checks everything before it changes anything, so a refused call leaves every graph as it was. The one exception is a merge that fixed a dynamic size through a connection before it refused, which leaves that size fixed, the size the merge would give it. The code name is `INVALID_ARGUMENT` for an argument the call cannot take, and `FAILED_PRECONDITION` for a graph it cannot take now, such as a finalized graph. A partition function that throws makes `split` refuse with the function's failure. A thrown `ClikaRT::Error` keeps its status, and any other exception comes back as `INTERNAL`. In Python, `graph.split` raises the function's own exception. [Handle errors by code](handle-errors-by-code.mdx) covers the channels every failure carries.

<LangTabs>

<LangTab value="cpp">

```cpp title="consumption.cpp"
int main() {
    std::vector<GraphPart> parts;
    parts.push_back({"encoder", trace_graph(rectify, "x", "h")});
    parts.push_back({"head", trace_graph(twice, "h", "y")});
    const std::vector<Connection> connections = {{{"encoder", "h"}, {"head", "h"}}};
    const Node relu = parts[0].graph.find_nodes(ClikaRT::graph::OpCode::Relu).front();   // a view taken before

    // A merge consumes every part's graph, and a view into one refuses from then on.
    const ClikaRT::graph::MergedGraph merged = ClikaRT::graph::merge(parts, connections);
    std::printf("%s | %zu %zu\n", io(merged.graph).c_str(), parts[0].graph.input_names().size(),
                parts[1].graph.input_names().size());   // x -> y | 0 0
    try {
        relu.check();
    } catch (const Error& error) {
        std::printf("%s | %s\n", error.code_name().c_str(), error.what());
    }
    // INVALID_ARGUMENT | the node 'Relu_1' is no longer in its graph: that graph no longer exists (a merge,
    // a split or an extract consumed it, or it was destroyed or assigned over)

    // A refused merge leaves every part as it was.
    std::vector<GraphPart> again;
    again.push_back({"encoder", trace_graph(rectify, "x", "h")});
    again.push_back({"head", trace_graph(twice, "h", "y")});
    const std::vector<Connection> wrong = {{{"encoder", "h"}, {"head", "z"}}};
    try {
        ClikaRT::graph::merge(again, wrong);
    } catch (const Error& error) {
        std::printf("%s | %s\n", error.code_name().c_str(), error.what());
    }
    // INVALID_ARGUMENT | merge: part 'head' has no input 'z'
    again[0].graph.finalize();
    try {
        ClikaRT::graph::merge(again, connections);
    } catch (const Error& error) {
        std::printf("%s | %s\n", error.code_name().c_str(), error.what());
    }
    // FAILED_PRECONDITION | merge: part 'encoder' is finalized; merge before finalize()
    std::printf("%s | %s\n", io(again[0].graph).c_str(), io(again[1].graph).c_str());   // x -> h | h -> y

    // A partition function that throws: split refuses with its failure, and the graph stays as it was.
    ModelGraph pooled = trace_graph(pool, "x", "y");
    try {
        ClikaRT::graph::split(pooled, [](const Node& node) -> std::string {
            if (node.op_code() == ClikaRT::graph::OpCode::Neg) {
                throw Error(ClikaRT::Status::InvalidArgument, "no part for Neg");
            }
            return "first";
        });
    } catch (const Error& error) {
        std::printf("%s | %s\n", error.code_name().c_str(), error.what());
    }
    // INVALID_ARGUMENT | split: the partition function failed on the node 'Neg_3': no part for Neg
    std::printf("%s\n", io(pooled).c_str());   // x -> y
    return 0;
}
```

</LangTab>

<LangTab value="python">

```python title="consumption.py"
def refuse_neg(node: Node) -> str:
    # A partition function that raises on the Neg.
    if node.op_code == crt.graph.OpCode.Neg:
        raise ValueError("no part for Neg")
    return "first"


parts = {"encoder": trace_graph(rectify, "x", "h"), "head": trace_graph(twice, "h", "y")}
connections = [("encoder.h", "head.h")]
(relu,) = parts["encoder"].find_nodes(crt.graph.OpCode.Relu)   # a view taken before

# A merge consumes every part's graph, and a view into one refuses from then on.
merged, _ = crt.graph.merge(parts, connections)
print(io(merged), "|", len(parts["encoder"].input_names()), len(parts["head"].input_names()))   # x -> y | 0 0
try:
    relu.check()
except crt.InvalidArgumentError as error:
    print(error.code_name, "|", error)
# INVALID_ARGUMENT | the node 'Relu_1' is no longer in its graph: that graph no longer exists (a merge,
# a split or an extract consumed it, or it was destroyed or assigned over)

# A refused merge leaves every part as it was.
again = {"encoder": trace_graph(rectify, "x", "h"), "head": trace_graph(twice, "h", "y")}
try:
    crt.graph.merge(again, [("encoder.h", "head.z")])
except crt.InvalidArgumentError as error:
    print(error.code_name, "|", error)   # INVALID_ARGUMENT | merge: part 'head' has no input 'z'
again["encoder"].finalize()
try:
    crt.graph.merge(again, connections)
except crt.InvalidArgumentError as error:
    print(error.code_name, "|", error)
# FAILED_PRECONDITION | merge: part 'encoder' is finalized; merge before finalize()
print(io(again["encoder"]), "|", io(again["head"]))   # x -> h | h -> y

# An exception the partition function raises leaves split as that same exception, and the graph as it was.
pooled = trace_graph(pool, "x", "y")
try:
    crt.graph.split(pooled, refuse_neg)
except ValueError as error:
    print(type(error).__name__, "|", error)   # ValueError | no part for Neg
print(io(pooled))   # x -> y
```

</LangTab>

</LangTabs>

Reuse the graphs a refused call leaves: they are as they were, and a corrected call takes them.
