Skip to main content

Structure inputs and outputs as pytrees

A pytree is a nested arrangement of Python containers whose ends are the values you care about: a dict of lists of tensors, a tuple of dicts, a dataclass holding both. clika_runtime.pytree takes such a tree apart into a flat list of leaves plus a TreeSpec that remembers the shape, and puts the two halves back together. Every boundary of the Python surface that accepts several tensors at once (compile, trace, eval, save and load) accepts a pytree, so a model's inputs and outputs keep the structure the code was written in.

The functions carry PyTorch's spellings (tree_flatten, tree_map, register_pytree_node), so code written against torch.utils._pytree reads the same here. Two rules differ from a plain container walk and are worth reading first.

Leaves, and why None is one of them​

Lists, tuples, dicts, OrderedDict, defaultdict, deque, namedtuples, and every class registered with the node registry are nodes; everything else (numbers, strings, bytes, tensors, objects the registry does not know) is a leaf. Leaves come out depth first, left to right, and a flatten followed by an unflatten reproduces the same container types.

flatten.py
from clika_runtime import pytree

tree = {"a": [1, (2, None)], "b": 3, "c": {"d": (4,), "e": []}}
leaves, spec = pytree.tree_flatten(tree)
print(leaves) # [1, 2, None, 3, 4]
print(spec.num_leaves) # 5
print(spec) # TreeSpec({'a': [*, (*, *)], 'b': *, 'c': {'d': (*,), 'e': []}})

rebuilt = pytree.tree_unflatten(spec, leaves)
assert rebuilt == tree and type(rebuilt) is dict

None is a leaf. A model that takes an optional input (a KV cache that is empty on the first step) keeps the same flat index for every other input whether or not the optional one is present, so a compiled graph's argument slots stay stable. none_is_leaf=False turns None into a childless node instead, and is_leaf=lambda x: isinstance(x, list) stops the descent at every list, which then counts as one leaf.

leaf_rules.py
from clika_runtime import pytree

with_cache = {"input_ids": 1, "cache": (2, 3)}
without_cache = {"input_ids": 1, "cache": None}
leaves_a, _ = pytree.tree_flatten(with_cache)
leaves_b, spec_b = pytree.tree_flatten(without_cache)
print(leaves_a[0], leaves_b[0]) # 1 1
print(leaves_b) # [1, None]
print(pytree.tree_unflatten(spec_b, [10, None])) # {'input_ids': 10, 'cache': None}

leaves, spec = pytree.tree_flatten({"x": 1, "cache": None}, none_is_leaf=False)
print(leaves) # [1]
print(spec) # TreeSpec({'x': *, 'cache': None}, none_is_leaf=False)

Map and reduce​

tree_map applies a function to every leaf and rebuilds the same structure; extra trees pair their leaves by position and must share the first tree's structure (a mismatch raises ValueError). tree_map_ runs the function for its side effect and returns the tree it was given; tree_map_only filters by type or by predicate. The reductions (tree_all, tree_any, tree_reduce, tree_sum, tree_max, tree_min) read the leaves without rebuilding anything, tree_leaves returns the flat list and tree_iter yields it lazily.

map.py
from clika_runtime import pytree

print(pytree.tree_map(lambda x: x is None, {"x": 1, "y": None})) # {'x': False, 'y': True}
print(pytree.tree_map(lambda x, y: x + y, {"a": 1, "b": 2}, {"a": 10, "b": 20})) # {'a': 11, 'b': 22}

mixed = {"a": 1, "b": "text", "c": [2, None, 2.5]}
print(pytree.tree_map_only(int, lambda x: x + 1, mixed)) # {'a': 2, 'b': 'text', 'c': [3, None, 2.5]}

Name every leaf by its key path​

tree_flatten_with_path returns (path, leaf) pairs; a path is a tuple of key entries (MappingKey, SequenceKey, GetAttrKey, DataclassKey), keystr renders it as the indexing expression that reaches the leaf, and key_get follows it. Key paths are how the file format names the leaves of a saved tree (below) and how a diagnostic points at one tensor inside a large state.

paths.py
from typing import NamedTuple

from clika_runtime import pytree


class Point(NamedTuple):
x: object
y: object


tree = {"a": [1, 2], "p": Point(3, 4)}
pairs = pytree.tree_flatten_with_path(tree)[0]
print([leaf for _, leaf in pairs]) # [1, 2, 3, 4]
print([pytree.keystr(path) for path, _ in pairs]) # ["['a'][0]", "['a'][1]", "['p'].x", "['p'].y"]

path = (pytree.MappingKey("a"), pytree.SequenceKey(1))
print(pytree.key_get(tree, path)) # 2

Register your own containers​

A class the registry does not know is a leaf. register_pytree_node(cls, flatten_fn, unflatten_fn) opens it: flatten_fn returns (children, context) or (children, context, entries) (the entries name the children for key paths), and unflatten_fn(children, context) rebuilds an instance. The argument order is PyTorch's. After registration, every tree function descends into the class and every rebuild produces an instance of it.

register.py
from clika_runtime import pytree


class Point:
def __init__(self, x, y):
self.x, self.y = x, y


pytree.register_pytree_node(
Point,
lambda p: ((p.x, p.y), None, ("x", "y")),
lambda children, _ctx: Point(*children),
)

print(pytree.tree_leaves({"p": Point(1, 2), "n": 3})) # [1, 2, 3]
mapped = pytree.tree_map(lambda v: v * 10, Point(1, 2))
print(type(mapped).__name__, mapped.x, mapped.y) # Point 10 20
spec = pytree.tree_structure(Point(1, 2))
print(spec.entries(), spec.paths()) # ['x', 'y'] [('x',), ('y',)]

A class can carry its own flattening as a method pair and register with the decorator. Note the method order: __tree_unflatten__(cls, context, children) takes the context first, while the function form above takes (children, context).

register_class.py
from clika_runtime import pytree


@pytree.register_pytree_node_class
class Pair:
def __init__(self, a, b):
self.a, self.b = a, b

def __tree_flatten__(self):
return (self.a, self.b), "pair", ("a", "b")

@classmethod
def __tree_unflatten__(cls, context, children):
return cls(*children)


leaves, spec = pytree.tree_flatten(Pair(1, [2, 3]))
print(leaves, spec.context) # [1, 2, 3] pair
rebuilt = pytree.tree_unflatten(spec, [10, 20, 30])
print(rebuilt.a, rebuilt.b) # 10 [20, 30]

Dataclasses register by field: the data fields become children, the meta fields ride the context and come back unchanged.

register_dataclass.py
import dataclasses

from clika_runtime import pytree


@dataclasses.dataclass
class Sample:
weight: object
bias: object
name: str


pytree.register_dataclass(Sample, data_fields=["weight", "bias"], meta_fields=["name"])

leaves, spec = pytree.tree_flatten(Sample(1, 2, "s"))
print(leaves) # [1, 2]
print(pytree.tree_unflatten(spec, [10, 20])) # Sample(weight=10, bias=20, name='s')
print([pytree.keystr(p) for p, _ in pytree.tree_flatten_with_path(Sample(1, 2, "s"))[0]]) # ['.weight', '.bias']

Registering a built-in container, an instance instead of a class, or the same class twice raises; unregister_pytree_node(cls) makes instances leaves again. A registration meant for one library goes into a namespace so it never changes what other code sees: register_pytree_node(..., namespace="mylib") opens the class only for calls that pass namespace="mylib" (tree_leaves(tree, namespace="mylib")), is_registered(cls, namespace="mylib") reports it, and a namespace with no registration of its own falls back to the global one.

Dictionary order​

Dicts flatten in insertion order, and the order is part of the structure: {"b": 1, "a": 2} and {"a": 2, "b": 1} have different specs. Code that wants two dicts with the same keys to share one structure regardless of insertion order wraps the calls in dict_insertion_ordered(False), which flattens by sorted key inside the block; the rebuilt dict still comes back in the original order.

dict_order.py
from clika_runtime import pytree

tree = {"b": 1, "a": 2, "c": 3}
leaves, spec = pytree.tree_flatten(tree)
print(leaves) # [1, 2, 3]
print(list(pytree.tree_unflatten(spec, leaves))) # ['b', 'a', 'c']

with pytree.dict_insertion_ordered(False):
print(pytree.tree_leaves({"b": 1, "a": 2})) # [2, 1]
print(pytree.tree_structure({"b": 1, "a": 2}) == pytree.tree_structure({"a": 0, "b": 0})) # True

Treespecs​

A TreeSpec is a value: it compares and hashes by structure, prints as the container shape with a * per leaf, and rebuilds a tree from any sequence of the right length (spec.unflatten(leaves); a wrong count raises ValueError naming the expected number). is_prefix asks whether one structure is the top of another, compose grafts an inner structure onto every leaf of an outer one, and flatten_up_to flattens a tree only as deep as the spec goes.

treespec.py
from clika_runtime import pytree

spec = pytree.tree_structure({"a": [1, (2, None)], "b": 3})
print(spec) # TreeSpec({'a': [*, (*, *)], 'b': *})
print(spec.num_leaves, spec.num_nodes, spec.num_children) # 4 7 2
print(spec.paths()) # [('a', 0), ('a', 1, 0), ('a', 1, 1), ('b',)]

short = pytree.tree_structure([1, 2])
deep = pytree.tree_structure([1, (2, 3)])
print(short.is_prefix(deep), deep.is_prefix(short)) # True False

outer = pytree.tree_structure([1, 2])
inner = pytree.tree_structure((1, 2))
print(outer.compose(inner)) # TreeSpec([(*, *), (*, *)])

treespec_dumps writes a spec as JSON and treespec_loads reads it back, so a structure can travel beside a file or a request. Built-in containers and namedtuples serialize as they are; a registered class needs a serialized_type_name, and a context that is not JSON needs the to_dumpable_context / from_dumpable_context pair at registration. An unnamed custom node or an unknown type name in the document raises ValueError.

serialize.py
import json

from clika_runtime import pytree

spec = pytree.tree_structure({"a": [1, (2, None)], "b": 3})
text = pytree.treespec_dumps(spec)
print(json.loads(text)["version"] == pytree.SERIALIZATION_PROTOCOL) # True
loaded = pytree.treespec_loads(text)
print(loaded == spec) # True
print(loaded.unflatten([1, 2, None, 3])) # {'a': [1, (2, None)], 'b': 3}

Pytrees at the boundaries​

crt.compile takes a callable whose arguments and return value are pytrees of tensors: it flattens the call's inputs, treats the tensor leaves as graph inputs and every other leaf as a static value that is part of the guard, and rebuilds the output structure on the way out; dynamic=True keeps one graph across batch sizes.

compile_tree.py
import numpy as np

import clika_runtime as crt
import clika_runtime.nn as nn

BATCH, D_IN, D_OUT = 5, 131, 64
rng = np.random.default_rng(0)
w = rng.standard_normal((D_IN, D_OUT)).astype(np.float32) * 0.1
b = rng.standard_normal(D_OUT).astype(np.float32)
x = rng.standard_normal((BATCH, D_IN)).astype(np.float32)


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)


class Wrapped(nn.Module):
def __init__(self) -> None:
super().__init__()
self.head = Head(w, b)

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": crt.tensor(x)})
print(sorted(out)) # ['argmax', 'probs']
again = compiled({"x": crt.tensor(np.tile(x, (2, 1)))})
print(tuple(again["probs"].shape)) # (10, 64)

crt.trace records a function once over stand-ins and returns a graph that runs by position or by name; the traced function takes its tensor inputs as one list and returns a list, and Trace eager code to graphs walks through it.

Operators return before their work runs. crt.eval(*trees) flattens whatever trees it is given, skips the leaves that are not tensors, and returns once every tensor leaf holds its value; a later read then costs no wait.

eval_tree.py
import numpy as np

import clika_runtime as crt

x = np.random.default_rng(0).standard_normal((5, 131)).astype(np.float32)
t = crt.tensor(x)
tree = {"a": crt.exp(t), "b": [crt.sum(t), None, "text"], "c": (crt.relu(t),)}
crt.eval(tree, crt.abs(t))
print(np.allclose(tree["a"].numpy(), np.exp(x.astype(np.float64)), rtol=1e-5, atol=1e-6)) # True

crt.save writes a tensor, a state dict, or any pytree of tensors as a safetensors file, and crt.load reads it back. A state dict (crt.save(state, path), crt.load(path, device="cpu")) saves under its own keys in insertion order, so another safetensors reader sees it as it is; a deeper tree saves one entry per leaf named by keystr of its key path, with the structure carried in the file's metadata, and crt.load rebuilds the same tree.

save_tree.py
from pathlib import Path

import numpy as np

import clika_runtime as crt
from clika_runtime import pytree

rng = np.random.default_rng(0)
tree = {
"layers": [
{"w": crt.tensor(rng.standard_normal((32, 131)).astype(np.float32)), "b": crt.tensor(np.zeros(32, np.float32))}
for _ in range(2)
],
"step": crt.tensor(np.array([7], dtype=np.int64)),
}
crt.save(tree, Path("tree.safetensors"))
back = crt.load(Path("tree.safetensors"))
print(pytree.tree_structure(back) == pytree.tree_structure(tree)) # True
print(sorted(pytree.keystr(kp) for kp, _ in pytree.tree_flatten_with_path(tree)[0]))
# ["['layers'][0]['b']", "['layers'][0]['w']", "['layers'][1]['b']", "['layers'][1]['w']", "['step']"]

The same structure rule holds for a module's state: crt.save(module.state_dict(), path) followed by other.load_state_dict(crt.load(path)) reproduces the forward, and the file is plain safetensors any other tool reads. Use ClikaRT from Python covers the tensor and module surface the leaves belong to.