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().
- C++
- Python
#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;
}
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
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.
- C++
- Python
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;
}
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
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").
- C++
- Python
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;
}
# 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
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.
- C++
- Python
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;
}
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))
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".
- C++
- Python
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;
}
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
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.
- C++
- Python
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;
}
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
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.
- C++
- Python
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;
}
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
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.
- C++
- Python
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;
}
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)
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.
- C++
- Python
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;
}
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
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.
- C++
- Python
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;
}
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
Reuse the graphs a refused call leaves: they are as they were, and a corrected call takes them.