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.
| Name | Kind |
|---|---|
AutoEntry | class |
DataclassEntry | class |
DataclassKey | class |
FlattenedEntry | class |
GetAttrEntry | class |
GetAttrKey | class |
GetItemEntry | class |
KeyEntry | class |
LeafSpec | class |
MappingEntry | class |
MappingKey | class |
NamedTupleEntry | class |
NodeRegistryEntry | class |
PyTreeAccessor | class |
PyTreeEntry | class |
PyTreeSpec | class |
SequenceEntry | class |
SequenceKey | class |
StructSequenceEntry | class |
TreeKind | class |
TreeSpec | class |
TreeSpecSubspecs | class |
| functions | module functions |