Trace eager code to graphs
Eager code runs op by op. Tracing runs your function ONCE over data-free stand-ins (shapes and dtypes matter, values are never read) and captures the operator graph: a ModelGraph, the same runnable type the ONNX loader compiles to, with the same run.
Two rules make a function traceable:
- List in, list out. A traced callable receives its input tensors as one list and returns its outputs as a list. A bare tensor return does not auto-wrap; return
[y]. - No value reads. The stand-ins carry no data, so reading a value during tracing raises. Shape-driven math is fine.
- Python
- C++
In Python the two rules relax to pytrees: a traced function may take a list, a dict or any nested container of tensors and return one, and the graph keeps the names. No value is read during the capture in either form.
import numpy as np
import clika_runtime as crt
w = crt.tensor(np.full((4, 3), 0.1, dtype=np.float32))
b = crt.tensor(np.zeros(3, dtype=np.float32))
def fn(inputs: list[crt.Tensor]) -> list[crt.Tensor]:
return [crt.softmax(crt.relu(crt.add(crt.matmul(inputs[0], w), b)), dim=-1)]
x = crt.tensor(np.ones((2, 4), dtype=np.float32))
# Capture: fn runs once over stand-ins; the graph names its inputs and outputs.
graph = crt.trace(fn, example_inputs=[x], input_names=["x"], output_names=["p"])
print(graph.input_names(), graph.output_names()) # ['x'] ['p']
# The graph runs like the function did, by position or by name.
(positional,) = graph.run([x])
(named,) = graph.run({"x": x})
print(positional.numpy()[0]) # [0.33333334 0.33333334 0.33333334]
A module traces the same way, and the traced graph is inspectable node by node: each node carries an op_code from crt.graph.OpCode and a kind.
import numpy as np
import clika_runtime as crt
import clika_runtime.nn as nn
class Head(nn.Module):
def __init__(self, w: np.ndarray, b: np.ndarray) -> None:
super().__init__()
self.weight = nn.Parameter(crt.tensor(w))
self.bias = nn.Parameter(crt.tensor(b))
def forward(self, x: crt.Tensor) -> crt.Tensor:
return crt.softmax(crt.relu(crt.add(crt.matmul(x, self.weight), self.bias)), dim=-1)
rng = np.random.default_rng(0)
head = Head(rng.standard_normal((4, 3)).astype(np.float32), np.zeros(3, dtype=np.float32))
x = crt.tensor(np.ones((2, 4), dtype=np.float32))
traced = crt.trace(head, example_inputs=(x,))
graph = traced.graph # the runnable ModelGraph
codes = {node.op_code for node in graph.nodes()}
print(crt.graph.OpCode.MatMul in codes, crt.graph.OpCode.Softmax in codes) # True True
print(sum(node.kind == crt.graph.NodeKind.Input for node in graph.nodes())) # 1
The matmul, the bias add and the relu lower to one MatMul node with the bias and the activation folded in, then the Softmax; the graph has one input node.
The same capture surface in C++: ClikaRT::graph::trace takes the callable and the example inputs and returns the ModelGraph; ClikaRT::compile wraps a callable in capture-and-replay exactly as below. The shipped examples include a full tracing walkthrough.
The compile wrapper
compile wraps a callable in capture-and-replay: the first call runs eagerly AND captures; later calls replay the graph. A shape change recaptures, invisible to values and visible on the counter.
- Python
step = crt.compile(fn, [crt.TensorSpec("x", crt.float32, [2, 4])])
print(step.state) # State.Pending (State.Pending: nothing captured yet)
(first,) = step([x]) # runs eagerly and records the graph
print(step.state) # State.Compiled (State.Compiled: later calls replay)
(again,) = step([x]) # replays the captured graph; Python is not called
print(step.recapture_count) # 0 (0)
graph = step.take_graph() # the captured ModelGraph, for standalone use
(replayed,) = graph.run([x])
The pytree form takes a module or a function over dicts and returns dicts; dynamic=True keeps one graph across batch sizes instead of recapturing on a shape change.
class Wrapped(nn.Module):
def __init__(self) -> None:
super().__init__()
self.head = head
def forward(self, batch: dict[str, crt.Tensor]) -> dict[str, crt.Tensor]:
p = self.head(batch["x"])
return {"probs": p, "argmax": crt.argmax(p, dims=[1])}
compiled = crt.compile(Wrapped(), dynamic=True)
out = compiled({"x": x})
print(sorted(out)) # ['argmax', 'probs'] (['argmax', 'probs'])
again = compiled({"x": crt.tensor(np.ones((4, 4), dtype=np.float32))})
print(again["probs"].shape) # clika_runtime.Size([4, 3])
fullgraph=True turns a fallback into an error: a Python branch on a tensor value cannot be recorded, and crt.compile(fn, fullgraph=True) raises crt.ClikaRTError at the call instead of serving it eagerly. step.reset() re-arms the wrapper, and step.state reads Pending, Compiled or Fallback.