Skip to main content

clika_runtime.pytree

Pytrees: nested Python containers of tensors (or anything else), flattened to a list and rebuilt from it.

A pytree is a tree of Python containers: tuples, lists, dicts, OrderedDict, defaultdict, deque, namedtuples, and any class you register. Everything else is a leaf: a :class:~clika_runtime.Tensor, a number, a string, an nn.Module, an arbitrary object. Model inputs and outputs are pytrees, so one call handles a batch dict, a tuple of tensors, a key-value cache, or a nested mix of them.

The three calls most code needs::

>>> from clika_runtime import pytree
>>> batch = {"input_ids": [1, 2], "mask": (3, None)}
>>> leaves, spec = pytree.tree_flatten(batch)
>>> leaves
[1, 2, 3, None]
>>> spec
TreeSpec({'input_ids': [*, *], 'mask': (*, *)})
>>> pytree.tree_unflatten(spec, [10, 20, 30, 40])
{'input_ids': [10, 20], 'mask': (30, 40)}
>>> pytree.tree_map(lambda x: x if x is None else x + 1, batch)
{'input_ids': [2, 3], 'mask': (4, None)}

None is a leaf by default: it keeps its slot in the flat list, so the index of every other leaf is the same whether an optional input is present or absent. Pass none_is_leaf=False to make None a node with no children instead. Dictionaries flatten in key insertion order; :func:dict_insertion_ordered switches a block of code to sorted keys.

A :class:TreeSpec records the structure and answers questions about it: leaf and node counts, children, whether one structure is a prefix of another, and how two structures broadcast against each other. It serializes to JSON (:func:treespec_dumps / :func:treespec_loads) and names each leaf's position as a key path::

>>> pairs, spec = pytree.tree_flatten_with_path(batch)
>>> [pytree.keystr(kp) for kp, _ in pairs]
["['input_ids'][0]", "['input_ids'][1]", "['mask'][0]", "['mask'][1]"]
>>> pytree.treespec_loads(pytree.treespec_dumps(spec)) == spec
True

Dataclasses register in one line (:func:register_dataclass, or the :func:dataclass decorator with :func:field), splitting their fields into children and configuration; :func:register_generic covers plain classes.

Custom containers join with :func:register_pytree_node (a class plus a flatten and an unflatten function) or :func:register_pytree_node_class (a decorator for a class carrying __tree_flatten__ / __tree_unflatten__). A namespace scopes a registration to the calls that name it, so two libraries registering the same type never collide.

NameKind
AutoEntryclass
DataclassEntryclass
DataclassKeyclass
FlattenedEntryclass
GetAttrEntryclass
GetAttrKeyclass
GetItemEntryclass
KeyEntryclass
LeafSpecclass
MappingEntryclass
MappingKeyclass
NamedTupleEntryclass
NodeRegistryEntryclass
PyTreeAccessorclass
PyTreeEntryclass
PyTreeSpecclass
SequenceEntryclass
SequenceKeyclass
StructSequenceEntryclass
TreeKindclass
TreeSpecclass
TreeSpecSubspecsclass
functionsmodule functions