Skip to main content

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.

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.

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

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

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

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