Skip to main content

Merge and split

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().

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;
}

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 GraphParts, 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.

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;
}

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").

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;
}

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.

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;
}

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".

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;
}

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 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.

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;
}

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.

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;
}

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 Values.

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;
}

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.

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;
}

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 covers the channels every failure carries.

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;
}

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