Skip to main content

clika_runtime.pytree functions

all_leaves​

all_leaves(iterable: 'Iterable[Any]', *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'bool'

Whether every element of iterable is a leaf.

pytree.all_leaves([1, 2, None]) True pytree.all_leaves([1, [2]]) False

broadcast_common​

broadcast_common(tree: 'Any', other_tree: 'Any', *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'tuple[list[Any], list[Any]]'

The leaves of :func:tree_broadcast_common, two flat lists of equal length.

pytree.broadcast_common([1, (2, 3), 4], [5, 6, (7, 8)]) ([1, 2, 3, 4, 4], [5, 6, 6, 7, 8])

broadcast_prefix​

broadcast_prefix(prefix_tree: 'Any', full_tree: 'Any', *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'list[Any]'

The leaves of :func:tree_broadcast_prefix as a flat list.

pytree.broadcast_prefix([1, 2, 3], [4, 5, (6, 7)]) [1, 2, 3, 3]

dataclass​

dataclass(cls: '_ClassT | None' = None, /, *, namespace: 'str' = '', serialized_type_name: 'str | None' = None, **dataclass_kwargs: 'Any') -> '_ClassT | Callable[[_ClassT], _ClassT]'

:func:dataclasses.dataclass that also registers the class as a pytree node; every field is a child unless declared field(pytree_node=False).

from clika_runtime import pytree @pytree.dataclass ... class Point: ... x: float ... y: float ... unit: str = pytree.field(default="m", pytree_node=False) pytree.tree_leaves(Point(1.0, 2.0)) [1.0, 2.0] pytree.tree_map(lambda v: v * 2, Point(1.0, 2.0)) Point(x=2.0, y=4.0, unit='m')

Accepts the keyword arguments of :func:dataclasses.dataclass (frozen=True, slots=True, ...) plus namespace and serialized_type_name.

dict_insertion_ordered​

dict_insertion_ordered(mode: 'bool', /, *, namespace: 'str' = '') -> 'Iterator[None]'

Choose, for the duration of the block, whether :class:dict and :class:~collections.defaultdict keys flatten in insertion order (True, the default) or in sorted order (False).

Sorted order makes two dictionaries with the same keys flatten to the same treespec whatever order their keys were inserted in; unflatten restores each dictionary's own insertion order either way. :class:~collections.OrderedDict always keeps insertion order.

tree = {"b": 1, "a": 2} tree_leaves(tree) [1, 2] with dict_insertion_ordered(False): ... tree_leaves(tree) [2, 1]

Args: mode: True for insertion order, False for sorted keys. namespace: Apply the choice to tree operations called with this namespace only; "" (the default) sets the global choice, which every namespace without its own choice follows.

The flag is process-wide state, not thread-local: a block in one thread is visible to flattening in another.

ensure_hf_batch_containers​

ensure_hf_batch_containers() -> 'int'

ensure_hf_batch_containers() -> int

Register the batch containers when transformers is already imported (never importing it); return how many registrations were new. One dict lookup once the registrations exist, so a call site on the compile path pays nothing after the first call.

entry_to_slice​

entry_to_slice(treespec: 'TreeSpec', entry: 'Any') -> 'slice'

The flat leaf index range under the top-level entry (a key, an index, a field name or a key entry) of treespec.

from clika_runtime import pytree pytree.entry_to_slice(pytree.tree_structure({"a": (1, 2), "b": 3}), "b") slice(2, 3, None)

Raises: ValueError: If entry names no child of the root.

field​

field(*, pytree_node: 'bool' = True, metadata: 'dict[Any, Any] | None' = None, **kwargs: 'Any') -> 'Any'

:func:dataclasses.field with one extra flag: pytree_node.

pytree_node=True (the default) makes the field a child of the node; False keeps it in the treespec context as configuration. The flag is stored in field.metadata["pytree_node"] (a flag already present in metadata is replaced by the argument).

Raises: ValueError: If pytree_node=True is combined with init=False (a child must be an __init__ argument so the node can be rebuilt; write field(init=False, pytree_node=False)).

is_namedtuple​

is_namedtuple(obj: 'object') -> 'bool'

Return whether obj is a namedtuple class or an instance of one.

is_namedtuple_class​

is_namedtuple_class(cls: 'object') -> 'bool'

Return whether cls is a namedtuple class (a tuple subclass with the _fields / _make / _asdict shape :func:collections.namedtuple and :class:typing.NamedTuple produce).

is_namedtuple_instance​

is_namedtuple_instance(obj: 'object') -> 'bool'

Return whether obj is an instance of a namedtuple class.

is_registered​

is_registered(cls: 'type', *, namespace: 'str' = '') -> 'bool'

Return whether instances of cls are pytree nodes under namespace (through a registration, or because cls is a built-in node type, a namedtuple class, or a PyStructSequence type).

is_registered(dict) True is_registered(set) False

is_structseq​

is_structseq(obj: 'object') -> 'bool'

Return whether obj is a PyStructSequence type or an instance of one.

is_structseq_class​

is_structseq_class(cls: 'object') -> 'bool'

Return whether cls is a PyStructSequence type (os.stat_result, sys.version_info, time.struct_time, ...).

is_structseq_instance​

is_structseq_instance(obj: 'object') -> 'bool'

Return whether obj is an instance of a PyStructSequence type.

key_get​

key_get(obj: 'Any', kp: 'KeyPath') -> 'Any'

Follow kp into obj and return the value it reaches.

key_get({"layer": [object(), 3]}, (MappingKey("layer"), SequenceKey(1))) 3

keystr​

keystr(kp: 'KeyPath') -> 'str'

Render a key path as the code that reaches the leaf.

keystr((MappingKey("layer"), SequenceKey(0), GetAttrKey("weight"))) "['layer'][0].weight"

make_dataclass​

make_dataclass(cls_name: 'str', fields: 'Iterable[Any]', *, namespace: 'str' = '', serialized_type_name: 'str | None' = None, **kwargs: 'Any') -> 'type'

:func:dataclasses.make_dataclass that also registers the new class as a pytree node. Field specs may use :func:field to mark meta fields.

from clika_runtime import pytree Pair = pytree.make_dataclass("Pair", [("a", int), ("b", int)]) pytree.tree_leaves(Pair(1, 2)) [1, 2]

namedtuple_fields​

namedtuple_fields(obj: 'object') -> 'tuple[str, ...]'

Return the field names of a namedtuple class or instance.

Raises: TypeError: If obj is neither a namedtuple class nor an instance of one.

path_to_slice​

path_to_slice(treespec: 'TreeSpec', path: 'Sequence[Any]') -> 'slice'

The flat leaf index range under path of treespec. path is a key path or a tuple of raw entries (keys, indices, field names); the empty path is the whole spec.

from clika_runtime import pytree pytree.path_to_slice(pytree.tree_structure({"a": (1, 2), "b": 3}), ("a",)) slice(0, 2, None)

Raises: ValueError: If path names no node of treespec.

reexport​

reexport(*, namespace: 'str' = '', module: 'str | None' = None) -> 'types.ModuleType'

Mint a module that offers this package's API with namespace bound as the default of every function that takes one.

A library that registers its own node types under a namespace exposes this module to its users, so they flatten and map without spelling the namespace at each call::

# mylib/__init__.py
from clika_runtime import pytree as _pytree
pytree = _pytree.reexport(namespace="mylib", module="mylib.pytree")

# user code
import mylib
mylib.pytree.tree_leaves(tree) # namespace="mylib" applied

Args: namespace: The namespace to bind ("" binds the global one). module: The dotted name registered in :data:sys.modules; default "<caller module>.pytree".

Returns: The new module; module.namespace holds the bound namespace.

Raises: ValueError: If module is not a dotted identifier or already exists.

register_dataclass​

register_dataclass(cls: 'type', data_fields: 'Iterable[str] | None' = None, meta_fields: 'Iterable[str] | None' = None, *, drop_fields: 'Iterable[str]' = (), namespace: 'str' = '', serialized_type_name: 'str | None' = None) -> 'type'

Register an existing dataclass as a pytree node, naming its data and meta fields.

data_fields become the node's children (in that order); meta_fields stay in the treespec context. With neither given, each field's field(pytree_node=...) flag decides (True by default). With only data_fields given, every other __init__ field is meta. Fields in drop_fields are neither; they must have defaults for the rebuild.

import dataclasses from clika_runtime import pytree @dataclasses.dataclass ... class Batch: ... ids: list ... mask: list ... name: str = "train" _ = pytree.register_dataclass(Batch, ["ids", "mask"], ["name"]) leaves, spec = pytree.tree_flatten(Batch([1, 2], [1, 1])) leaves [1, 2, 1, 1] pytree.tree_unflatten(spec, leaves) == Batch([1, 2], [1, 1]) True

Raises: TypeError: If cls is not a dataclass, has InitVar fields, or a data field is not an __init__ argument. ValueError: If a field name is unknown, listed on both sides, or the class is already registered as a dataclass node.

register_generic​

register_generic(cls: 'type', fields: 'Iterable[str] | None' = None, *, namespace: 'str' = '', serialized_type_name: 'str | None' = None) -> 'type'

Register an arbitrary class as a pytree node from its attribute state.

Children are the values of fields when given, otherwise the object's set __slots__ and __dict__ attributes in sorted name order (the names travel in the treespec context, so a tree always rebuilds the attributes it was flattened with).

Rebuilding bypasses __init__: a new instance comes from cls.__new__(cls), each attribute is set back, and __post_init__ is called when the class defines one. A class whose __init__ computes derived state or validates its arguments should register through :func:~clika_runtime.pytree.register_pytree_node or :func:~clika_runtime.pytree.register_dataclass instead.

from clika_runtime import pytree class Output: ... def init(self, logits, hidden): ... self.logits, self.hidden = logits, hidden _ = pytree.register_generic(Output) leaves, spec = pytree.tree_flatten(Output(1, [2, 3])) leaves [2, 3, 1] vars(pytree.tree_unflatten(spec, leaves)) {'hidden': [2, 3], 'logits': 1}

Raises: TypeError: If cls is not a class. ValueError: If cls is already registered in namespace.

register_hf_batch_containers​

register_hf_batch_containers(transformers_module: 'Any' = None) -> 'int'

register_hf_batch_containers(transformers_module=None) -> int

Register transformers.BatchEncoding and transformers.BatchFeature as pytree nodes; return how many registrations were new. Imports transformers when no module is given and returns 0 when it is not installed. Idempotent.

register_leaf_meta​

register_leaf_meta(cls: 'type', meta_fn: 'Callable[[Any], Sequence[Any]]', *, namespace: 'str' = '') -> 'None'

Declare how a registered node type annotates the leaves under it.

meta_fn(node) returns one metadata value per DIRECT child, in the children's flatten order. A child's metadata reaches every leaf beneath it; a nested node's own annotation wins over an inherited one (two dict annotations merge, the nearer keys winning). None means "no annotation from this node".

from clika_runtime import pytree class Cache: ... def init(self, keys, values): ... self.keys, self.values = keys, values _ = pytree.register_pytree_node(Cache, lambda c: ((c.keys, c.values), None, ("keys", "values")), ... lambda children, _: Cache(*children)) pytree.register_leaf_meta(Cache, lambda c: ({"is_key": True}, {"is_key": False})) leaves, metas, spec = pytree.tree_flatten_with_meta({"c": Cache([1, 2], [3, 4]), "x": 5}) leaves, metas ([1, 2, 3, 4, 5], [{'is_key': True}, {'is_key': True}, {'is_key': False}, {'is_key': False}, None])

register_namedtuple​

register_namedtuple(cls: 'type', *, serialized_type_name: 'str | None' = None, namespace: 'str' = '') -> 'type'

Attach a serialized type name to a namedtuple class.

Namedtuple classes are pytree nodes without registration; this call only records the name a serialized treespec uses for cls.

Raises: TypeError: If cls is not a namedtuple class. ValueError: If cls is already registered in namespace.

register_node​

register_node(cls: '_ClassT | str | None' = None, /, *, namespace: 'str | None' = None, serialized_type_name: 'str | None' = None) -> '_ClassT | Callable[[_ClassT], _ClassT]'

Register an existing dataclass as a pytree node using its field(pytree_node=...) flags; a plain call or a decorator (the one positional string is the namespace).

from clika_runtime import pytree @pytree.register_node ... @dataclasses.dataclass ... class Pair: ... a: int ... b: int pytree.tree_leaves(Pair(1, 2)) [1, 2]

register_pytree_node​

register_pytree_node(cls: 'type', flatten_fn: 'FlattenFn', unflatten_fn: 'UnflattenFn', *, serialized_type_name: 'str | None' = None, to_dumpable_context: 'ToDumpableContextFn | None' = None, from_dumpable_context: 'FromDumpableContextFn | None' = None, flatten_with_keys_fn: 'FlattenWithKeysFn | None' = None, namespace: 'str' = '') -> 'type'

Make cls a pytree node: instances open into children and close back.

Args: cls: The class to register. flatten_fn: Called with an instance; returns (children, context) or (children, context, entries). children is any iterable of the values to descend into. context is whatever the type needs to rebuild itself and must not contain tensors; it is stored in the treespec and compared with == when two treespecs are compared. entries (optional) names each child for key paths, for example attribute names; None or an omitted third item means range(len(children)). unflatten_fn: Called with (children, context), where children is a tuple in the same order flatten_fn produced them; returns a new instance. (A class carrying __tree_unflatten__(cls, context, children) registers through :func:register_pytree_node_class, which adapts that method's argument order.) serialized_type_name: The name a serialized treespec records for this type; defaults to the class's qualified name at serialization time. to_dumpable_context: Converts a context into a JSON-serializable value for treespec serialization. Both dumpable hooks or neither. from_dumpable_context: The inverse of to_dumpable_context. flatten_with_keys_fn: flatten_with_keys_fn(node) -> ([(key, child), ...], context), the key-path form used by the *_with_path functions; derived from flatten_fn and entries when omitted. namespace: Scope the registration to tree operations called with the same namespace. "" (the default) registers globally.

Returns: cls, so the call can be used as a decorator body.

Raises: TypeError: If cls is not a class or namespace is not a string. ValueError: If cls is a built-in node type, or is already registered in the same namespace, or only one dumpable hook is given.

Example: >>> class Point: ... def init(self, x, y): ... self.x, self.y = x, y >>> register_pytree_node( ... Point, ... lambda p: ((p.x, p.y), None, ("x", "y")), ... lambda children, _: Point(*children), ... ) # doctest: +ELLIPSIS <class '...Point'> >>> from clika_runtime import pytree >>> pytree.tree_leaves({"p": Point(1, 2), "n": 3}) [1, 2, 3] >>> pytree.tree_map(lambda v: v * 10, Point(1, 2)).x 10

register_pytree_node_class​

register_pytree_node_class(cls: '_ClassT | str | None' = None, /, *, namespace: 'str | None' = None, serialized_type_name: 'str | None' = None, to_dumpable_context: 'ToDumpableContextFn | None' = None, from_dumpable_context: 'FromDumpableContextFn | None' = None) -> '_ClassT | Callable[[_ClassT], _ClassT]'

Register a class that carries its own flatten and unflatten methods.

The class defines __tree_flatten__(self) -> (children, context[, entries]) and the classmethod __tree_unflatten__(cls, context, children). Note the method form takes context first; the function form registered with :func:register_pytree_node takes unflatten_fn(children, context). The registry adapts the method form, so a class written for either convention needs no change. A class that spells the pair tree_flatten / tree_unflatten (same argument orders) is accepted too; the dunder pair is added to it.

Usable as a bare decorator, as a decorator with a namespace (keyword or as the one positional string), or as a plain call on the class::

@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), None, ("a", "b")
@classmethod
def __tree_unflatten__(cls, context, children):
return cls(*children)

@register_pytree_node_class("mylib") # namespace as the string
class Scoped(Pair): ...

register_pytree_node_class(Pair, namespace="other")

Returns: The class when called on one; otherwise a decorator that registers the class it receives and returns it.

Raises: TypeError: If the class defines neither method pair, or a namespace is not a string. ValueError: If both a positional namespace string and namespace= are given, or the class is already registered in the namespace.

structseq_fields​

structseq_fields(obj: 'object') -> 'tuple[str, ...]'

Return the names of the positional fields of a PyStructSequence type or instance.

Raises: TypeError: If obj is neither a PyStructSequence type nor an instance of one.

tree_accessors​

tree_accessors(tree: 'Any', *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'list[PyTreeAccessor]'

The accessor of every leaf of tree, in flatten order.

tree_all​

tree_all(pred_or_tree: 'Any', tree: 'Any' = <missing>, *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'bool'

Whether every leaf is true (or satisfies a predicate); True for an empty tree.

Two call forms: tree_all(tree) tests the truth value of each leaf; tree_all(pred, tree) tests pred(leaf).

pytree.tree_all({"a": 1, "b": [2, 3]}) True pytree.tree_all(lambda x: x > 1, {"a": 1, "b": [2, 3]}) False pytree.tree_all({"a": None}) # None is a leaf and is false False pytree.tree_all({"a": None}, none_is_leaf=False) True

tree_all_only​

tree_all_only(type_or_types_or_pred: 'type | tuple[type, ...] | Callable[[Any], bool]', pred: 'Callable[[Any], bool]', tree: 'Any', *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'bool'

Whether pred holds for every leaf of the given type(s); leaves of other types are skipped.

pytree.tree_all_only(int, lambda x: x > 0, {"a": 1, "b": "text"}) True

tree_any​

tree_any(pred_or_tree: 'Any', tree: 'Any' = <missing>, *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'bool'

Whether some leaf is true (or satisfies a predicate); False for an empty tree.

Two call forms: tree_any(tree) tests the truth value of each leaf; tree_any(pred, tree) tests pred(leaf).

pytree.tree_any({"a": 0, "b": [0, 3]}) True pytree.tree_any(lambda x: x > 5, {"a": 0, "b": [0, 3]}) False

tree_any_only​

tree_any_only(type_or_types_or_pred: 'type | tuple[type, ...] | Callable[[Any], bool]', pred: 'Callable[[Any], bool]', tree: 'Any', *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'bool'

Whether pred holds for some leaf of the given type(s); leaves of other types are skipped.

pytree.tree_any_only(str, lambda s: s.startswith("t"), {"a": 1, "b": "text"}) True

tree_broadcast_common​

tree_broadcast_common(tree: 'Any', other_tree: 'Any', *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'tuple[Any, Any]'

Expand two trees to their common structure: wherever one has a leaf and the other a subtree, the leaf is copied over that subtree.

pytree.tree_broadcast_common([1, (2, 3), 4], [5, 6, (7, 8)]) ([1, (2, 3), (4, 4)], [5, (6, 6), (7, 8)])

Raises: ValueError: Where the two structures disagree (different node types, arities, keys, or contexts at the same position).

tree_broadcast_prefix​

tree_broadcast_prefix(prefix_tree: 'Any', full_tree: 'Any', *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'Any'

Expand prefix_tree to full_tree's structure: each leaf of the prefix is copied over the matching subtree of the full tree.

pytree.tree_broadcast_prefix(1, [2, 3, 4]) [1, 1, 1] pytree.tree_broadcast_prefix([1, 2, 3], [4, 5, (6, 7)]) [1, 2, (3, 3)]

Raises: ValueError: If prefix_tree is not a prefix of full_tree; the message names the first node that differs.

tree_flatten​

tree_flatten(tree: 'Any', *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'tuple[list[Any], TreeSpec]'

Split a pytree into its leaves and its structure.

Leaves come out in left-to-right, depth-first order. Dictionaries flatten in key insertion order (see :func:dict_insertion_ordered for sorted order). Everything the registry does not know is a leaf: tensors, numbers, strings, modules, arbitrary objects.

from clika_runtime import pytree leaves, spec = pytree.tree_flatten({"a": [1, (2, None)], "b": 3}) leaves [1, 2, None, 3] spec TreeSpec({'a': [, (, *)], 'b': *}) pytree.tree_unflatten(spec, [x * 10 for x in leaves if x is not None] + [0]) {'a': [10, (20, 30)], 'b': 0}

None is a leaf by default, so a missing optional input keeps every other leaf at the same flat index. With none_is_leaf=False it is a node without children and disappears from the leaf list:

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

is_leaf stops the descent early; the whole subtree becomes one leaf:

pytree.tree_leaves({"a": [1, 2], "b": [3]}, is_leaf=lambda x: isinstance(x, list)) [[1, 2], [3]]

Args: tree: The pytree to flatten. is_leaf: An optional predicate; True makes the node a leaf. none_is_leaf: Whether None is a leaf (default True). namespace: The registry namespace for custom node types.

Returns: (leaves, treespec): the list of leaves and the :class:TreeSpec that :func:tree_unflatten rebuilds the tree from.

tree_flatten_with_accessor​

tree_flatten_with_accessor(tree: 'Any', *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'tuple[list[PyTreeAccessor], list[Any], TreeSpec]'

Flatten tree into (accessors, leaves, treespec).

from clika_runtime import pytree accessors, leaves, spec = pytree.tree_flatten_with_accessor({"a": (1, 2)}) [a.codify() for a in accessors], leaves (["['a'][0]", "['a'][1]"], [1, 2])

tree_flatten_with_meta​

tree_flatten_with_meta(tree: 'Any', meta_fn: 'Callable[[Any], Any] | None' = None, *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'tuple[list[Any], list[Any], TreeSpec]'

Flatten tree into (leaves, metas, treespec): one metadata value per leaf.

metas[i] combines what the nodes above leaf i declare through :func:register_leaf_meta with meta_fn(leaf) when a meta_fn is given (two dicts merge, the leaf function's keys winning; otherwise the leaf function's value wins). With no declarations and no meta_fn every entry is None.

The metadata rides beside the leaves instead of inside the treespec, so a consumer that classifies flat slots (which slot is a key, which a value) reads metas[i] and never re-derives it from indices.

from clika_runtime import pytree leaves, metas, spec = pytree.tree_flatten_with_meta({"a": 1, "b": "x"}, lambda leaf: type(leaf).name) metas ['int', 'str']

tree_flatten_with_path​

tree_flatten_with_path(tree: 'Any', *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'tuple[list[tuple[KeyPath, Any]], TreeSpec]'

Flatten tree into (key_path, leaf) pairs and its treespec.

from clika_runtime import pytree pairs, spec = pytree.tree_flatten_with_path({"a": [1, 2], "b": 3}) [(pytree.keystr(kp), leaf) for kp, leaf in pairs] [("['a'][0]", 1), ("['a'][1]", 2), ("['b']", 3)]

tree_is_leaf​

tree_is_leaf(tree: 'Any', *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'bool'

Whether tree is a leaf (not a node the registry opens).

pytree.tree_is_leaf(1), pytree.tree_is_leaf([1]), pytree.tree_is_leaf(None) (True, False, True) pytree.tree_is_leaf(None, none_is_leaf=False) False

tree_iter​

tree_iter(tree: 'Any', *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'Iterator[Any]'

Iterate over the leaves of tree in flatten order, lazily.

list(pytree.tree_iter({"a": [1, 2], "b": 3})) [1, 2, 3]

tree_leaves​

tree_leaves(tree: 'Any', *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'list[Any]'

The leaves of tree in flatten order.

pytree.tree_leaves({"a": [1, 2], "b": (3, None)}) [1, 2, 3, None]

tree_leaves_with_path​

tree_leaves_with_path(tree: 'Any', *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'list[tuple[KeyPath, Any]]'

The (key_path, leaf) pairs of tree in flatten order.

tree_map​

tree_map(fn: 'Callable[..., Any]', tree: 'Any', *rests: 'Any', is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'Any'

Apply fn to every leaf and rebuild the tree with the results.

from clika_runtime import pytree pytree.tree_map(lambda x: x * 2, {"a": [1, 2], "b": 3}) {'a': [2, 4], 'b': 6}

With more trees, fn receives one leaf from each, position by position. The first tree fixes the structure; every other tree must have it as a prefix (it may be deeper where the first has a leaf):

pytree.tree_map(lambda x, y: x + y, {"a": 1, "b": 2}, {"a": 10, "b": 20}) {'a': 11, 'b': 22} pytree.tree_map(lambda x, sub: [x, *sub], [1, 2], [[7, 9], [8]]) [[1, 7, 9], [2, 8]]

None is a leaf by default and reaches fn; with none_is_leaf=False it stays in place untouched:

pytree.tree_map(lambda x: x is None, {"x": 1, "y": None}) {'x': False, 'y': True} pytree.tree_map(lambda x: x is None, {"x": 1, "y": None}, none_is_leaf=False) {'x': False, 'y': None}

Args: fn: Takes 1 + len(rests) arguments, one leaf from each tree. tree: The tree that fixes the structure and supplies the first argument. *rests: Trees with tree's structure as a prefix. is_leaf: An optional predicate; True makes the node a leaf. none_is_leaf: Whether None is a leaf (default True). namespace: The registry namespace for custom node types.

Returns: A new tree with tree's structure holding fn's results.

Raises: ValueError: If a tree in rests does not have tree as a prefix; the message names the first node that differs.

tree_map_​

tree_map_(fn: 'Callable[..., Any]', tree: 'Any', *rests: 'Any', is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'Any'

Call fn on every leaf for its side effect and return tree itself.

seen = [] _ = pytree.tree_map_(seen.append, {"a": [1, 2], "b": 3}) seen [1, 2, 3]

tree_map_only​

tree_map_only(type_or_types_or_pred: 'type | tuple[type, ...] | Callable[[Any], bool]', fn: 'Callable[[Any], Any]', tree: 'Any', *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'Any'

Apply fn only to the leaves that are instances of the given type(s) or satisfy the predicate; every other leaf is kept as is.

pytree.tree_map_only(int, lambda x: x + 1, {"a": 1, "b": "text", "c": [2, None]}) {'a': 2, 'b': 'text', 'c': [3, None]} pytree.tree_map_only(lambda x: x is None, lambda _: 0, [1, None]) [1, 0]

tree_map_only_​

tree_map_only_(type_or_types_or_pred: 'type | tuple[type, ...] | Callable[[Any], bool]', fn: 'Callable[[Any], Any]', tree: 'Any', *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'Any'

Like :func:tree_map_only, calling fn for its side effect and returning tree itself.

tree_map_with_path​

tree_map_with_path(fn: 'Callable[..., Any]', tree: 'Any', *rests: 'Any', is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'Any'

Like :func:~clika_runtime.pytree.tree_map, with each leaf's key path as the first argument: fn(key_path, leaf, *rest_leaves).

from clika_runtime import pytree pytree.tree_map_with_path(lambda kp, x: f"{pytree.keystr(kp)}={x}", {"a": [1, 2]}) {'a': ["['a'][0]=1", "['a'][1]=2"]}

tree_map_with_path_​

tree_map_with_path_(fn: 'Callable[..., Any]', tree: 'Any', *rests: 'Any', is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'Any'

Like :func:tree_map_with_path, calling fn for its side effect and returning tree itself.

tree_max​

tree_max(tree: 'Any', *, default: 'Any' = <missing>, key: 'Callable[[Any], Any] | None' = None, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'Any'

The largest leaf (by key when given); default for an empty tree.

pytree.tree_max({"a": 1, "b": [5, 3]}) 5 pytree.tree_max([], default=0) 0

Raises: ValueError: If the tree has no leaves and no default is given.

tree_min​

tree_min(tree: 'Any', *, default: 'Any' = <missing>, key: 'Callable[[Any], Any] | None' = None, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'Any'

The smallest leaf (by key when given); default for an empty tree.

pytree.tree_min({"a": 1, "b": [5, 3]}) 1

Raises: ValueError: If the tree has no leaves and no default is given.

tree_partition​

tree_partition(pred: 'Callable[[Any], bool]', tree: 'Any', *, fillvalue: 'Any' = None, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'tuple[Any, Any]'

Split the leaves by pred into two trees of tree's structure: the first keeps the leaves where pred is true, the second the rest; a leaf that goes to one side is fillvalue on the other.

left, right = pytree.tree_partition(lambda x: x > 10, {"x": 7, "y": (42, 64)}) left {'x': None, 'y': (42, 64)} right {'x': 7, 'y': (None, None)}

tree_paths​

tree_paths(tree: 'Any', *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'list[tuple[Any, ...]]'

The raw entry path of every leaf (indices, keys, field positions), in flatten order.

from clika_runtime import pytree pytree.tree_paths({"a": [1, 2], "b": 3}) [('a', 0), ('a', 1), ('b',)]

tree_ravel​

tree_ravel(tree: 'Any', *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'tuple[Any, Callable[[Any], Any]]'

Pack a tree of tensors into one 1-D tensor; return it with the function that unpacks such a tensor back into the tree's structure.

Every leaf must be a :class:~clika_runtime.Tensor. The result dtype is promoted across the leaves: by category first (bool < unsigned integer < signed integer < float), then to the widest width within the top category; a signed integer that must hold an unsigned operand of the same or a wider width widens one step (int32 with uint32 gives int64, int64 with uint64 stays int64); float16 with bfloat16 gives float32. Each leaf is cast to that dtype, flattened and concatenated in flatten order. The unpack function casts each piece back to the leaf's original dtype and shape. Quantized and sub-byte dtypes are refused.

import clika_runtime as crt from clika_runtime import pytree tree = {"w": crt.ones(2, 3), "b": crt.zeros(2)} flat, unravel = pytree.tree_ravel(tree) tuple(flat.shape) (8,) tuple(unravel(flat)["w"].shape) (2, 3)

Returns: (flat, unravel): the 1-D tensor and unravel(flat_like) -> tree. An empty tree gives an empty float32 tensor.

Raises: ImportError: If the runtime package is not importable. TypeError: If a leaf is not a tensor or has an unsupported dtype.

tree_reduce​

tree_reduce(fn: 'Callable[[Any, Any], Any]', tree: 'Any', initial: 'Any' = <missing>, *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'Any'

Reduce the leaves left to right with fn(accumulator, leaf), like :func:functools.reduce.

pytree.tree_reduce(lambda acc, x: acc + x, {"a": 1, "b": [2, 3]}) 6 pytree.tree_reduce(lambda acc, x: acc + x, {}, 0) 0

Raises: TypeError: If the tree has no leaves and no initial is given.

tree_replace_nones​

tree_replace_nones(sentinel: 'Any', tree: 'Any', *, namespace: 'str' = '') -> 'Any'

Replace every None leaf with sentinel.

pytree.tree_replace_nones(0, {"a": 1, "b": None, "c": (2, None)}) {'a': 1, 'b': 0, 'c': (2, 0)}

tree_structure​

tree_structure(tree: 'Any', *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'TreeSpec'

The :class:TreeSpec of tree alone.

pytree.tree_structure({"a": [1, 2], "b": (3, None)}) TreeSpec({'a': [*, ], 'b': (, *)})

tree_sum​

tree_sum(tree: 'Any', start: 'Any' = 0, *, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'Any'

Sum start and every leaf, left to right.

pytree.tree_sum({"a": 1, "b": [2, 3]}) 6 pytree.tree_sum({"a": "x", "b": ["y"]}, start="") 'xy'

tree_transpose​

tree_transpose(outer_treespec: 'TreeSpec', inner_treespec: 'TreeSpec', tree: 'Any', *, is_leaf: 'LeafPredicate | None' = None) -> 'Any'

Turn a tree of structure (outer, inner) inside out into (inner, outer).

outer = pytree.tree_structure({"a": 1, "b": 2}) inner = pytree.tree_structure((1, 2)) pytree.tree_transpose(outer, inner, {"a": (1, 2), "b": (3, 4)}) ({'a': 1, 'b': 3}, {'a': 2, 'b': 4})

Only the leaf count of tree is checked against the two specs; the caller guarantees the (outer, inner) nesting.

Raises: ValueError: If the two specs disagree on none_is_leaf or namespace, if either has no leaves, or if tree's leaf count is not outer.num_leaves * inner.num_leaves.

tree_transpose_map​

tree_transpose_map(fn: 'Callable[..., Any]', tree: 'Any', *rests: 'Any', inner_treespec: 'TreeSpec | None' = None, is_leaf: 'LeafPredicate | None' = None, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'Any'

Map fn over the leaves, then transpose: the results' shared structure comes out on top, with a tree of tree's structure at each of its leaves.

pytree.tree_transpose_map(lambda x: (x, x * 10), {"a": 1, "b": [2, 3]}) ({'a': 1, 'b': [2, 3]}, {'a': 10, 'b': [20, 30]})

inner_treespec fixes the results' structure; by default it is read from the first result.

Raises: ValueError: If tree has no leaves, or a result does not have the inner structure as a prefix.

tree_unflatten​

tree_unflatten(treespec: 'TreeSpec', leaves: 'Iterable[Any]') -> 'Any'

Rebuild a pytree from a :class:TreeSpec and its leaves.

leaves, spec = pytree.tree_flatten({"a": (1, 2)}) pytree.tree_unflatten(spec, ["x", "y"]) {'a': ('x', 'y')}

Raises: TypeError: If treespec is not a :class:TreeSpec. ValueError: If leaves does not hold exactly treespec.num_leaves values; the message states both counts.

treespec_accessors​

treespec_accessors(treespec: 'TreeSpec') -> 'list[PyTreeAccessor]'

One accessor per leaf of treespec, in leaf order.

treespec_custom_types​

treespec_custom_types(treespec: 'TreeSpec') -> 'set[type]'

The registered (custom) node types treespec contains: what a process must have registered before it can rebuild trees of this shape.

treespec_defaultdict​

treespec_defaultdict(default_factory: 'Callable[[], Any] | None' = None, mapping: 'Mapping[Any, TreeSpec] | Iterable[tuple[Any, TreeSpec]]' = (), *, none_is_leaf: 'bool' = True, namespace: 'str' = '', **kwargs: 'TreeSpec') -> 'TreeSpec'

The spec of a :class:~collections.defaultdict with the given factory whose values have the given specs.

pytree.treespec_defaultdict(list, a=pytree.treespec_leaf()) TreeSpec(defaultdict(<class 'list'>, {'a': *}))

treespec_deque​

treespec_deque(iterable: 'Iterable[TreeSpec]' = (), maxlen: 'int | None' = None, *, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'TreeSpec'

The spec of a :class:~collections.deque whose items have the given specs.

pytree.treespec_deque([pytree.treespec_leaf()], maxlen=4) TreeSpec(deque([*], maxlen=4))

treespec_dict​

treespec_dict(mapping: 'Mapping[Any, TreeSpec] | Iterable[tuple[Any, TreeSpec]]' = (), *, none_is_leaf: 'bool' = True, namespace: 'str' = '', **kwargs: 'TreeSpec') -> 'TreeSpec'

The spec of a dict whose values have the given specs.

pytree.treespec_dict({"a": pytree.treespec_leaf()}, b=pytree.tree_structure((1, 2))) TreeSpec({'a': , 'b': (, *)})

treespec_dumps​

treespec_dumps(treespec: 'TreeSpec', protocol: 'int | None' = None) -> 'str'

Serialize treespec to a JSON string.

The text is one object with "version" (the format version), the spec's "none_is_leaf" and "namespace", and the "root" node; every node is {"type", "context", "children_spec"} ("type": null marks a leaf). Object keys are sorted, so equal specs produce identical text.

from clika_runtime import pytree spec = pytree.tree_structure({"a": (1, 2)}) text = pytree.treespec_dumps(spec) pytree.treespec_loads(text) == spec True

Args: treespec: The spec to serialize. protocol: The format version to write; None picks the newest (:data:SERIALIZATION_PROTOCOL).

Raises: ValueError: If protocol is unknown, a custom node type has no serialized_type_name, or a defaultdict factory has no importable name. TypeError: If a context (or dict key) is not JSON data and the node type registered no dumpable-context hooks; the message names the type.

treespec_empty_template​

treespec_empty_template(treespec: 'TreeSpec', *, keep_custom: 'bool' = False) -> 'Any'

A tree with treespec's shape and None at every leaf.

By default every custom node becomes a list with the same number of children, so the template is built from built-in containers only, needs no registered types, and flattens to the same number of slots as the original. It stands in for an absent optional input whose leaves must keep their flat positions. With keep_custom=True custom nodes are rebuilt through their registered unflatten function with None children instead.

from clika_runtime import pytree spec = pytree.tree_structure({"ids": [1, 2], "mask": (3,)}) pytree.treespec_empty_template(spec) {'ids': [None, None], 'mask': (None,)}

treespec_from_collection​

treespec_from_collection(collection: 'Any', *, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'TreeSpec'

The spec of a container whose children are specs: collection is flattened one level and each child must be a :class:TreeSpec.

pytree.treespec_from_collection([pytree.LeafSpec(), pytree.tree_structure((1, 2))]) TreeSpec([, (, *)])

treespec_keypaths​

treespec_keypaths(treespec: 'TreeSpec') -> 'list[KeyPath]'

One key path per leaf of treespec, in leaf order.

from clika_runtime import pytree spec = pytree.tree_structure({"a": [1, 2], "b": 3}) [keystr(kp) for kp in treespec_keypaths(spec)] ["['a'][0]", "['a'][1]", "['b']"]

treespec_leaf​

treespec_leaf(*, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'TreeSpec'

The spec of a single leaf.

pytree.treespec_leaf() TreeSpec(*) pytree.treespec_leaf() == pytree.tree_structure(1) True

treespec_list​

treespec_list(iterable: 'Iterable[TreeSpec]' = (), *, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'TreeSpec'

The spec of a list whose items have the given specs.

pytree.treespec_list([pytree.treespec_leaf(), pytree.treespec_leaf()]) TreeSpec([*, *])

treespec_loads​

treespec_loads(data: 'str') -> 'TreeSpec'

Rebuild a :class:~clika_runtime.pytree.TreeSpec from :func:treespec_dumps text.

Custom node types resolve through the registry by their serialized_type_name (in the spec's namespace); they must be registered in the loading process.

Raises: ValueError: If the text is not a treespec dump, carries an unknown version, names a type that is not registered (the message names the type and the namespace) or not importable, or is malformed.

treespec_namedtuple​

treespec_namedtuple(namedtuple: 'tuple[Any, ...]', *, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'TreeSpec'

The spec of a namedtuple whose fields hold the given specs.

from collections import namedtuple Point = namedtuple("Point", ["x", "y"]) pytree.treespec_namedtuple(Point(pytree.treespec_leaf(), pytree.treespec_leaf())) TreeSpec(Point(x=, y=))

Raises: ValueError: If namedtuple is not a namedtuple instance.

treespec_none​

treespec_none(*, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'TreeSpec'

The spec of None: a leaf when none_is_leaf (the default), else a node with no children.

pytree.treespec_none() TreeSpec(*) pytree.treespec_none(none_is_leaf=False) TreeSpec(None, none_is_leaf=False)

treespec_ordereddict​

treespec_ordereddict(mapping: 'Mapping[Any, TreeSpec] | Iterable[tuple[Any, TreeSpec]]' = (), *, none_is_leaf: 'bool' = True, namespace: 'str' = '', **kwargs: 'TreeSpec') -> 'TreeSpec'

The spec of an :class:~collections.OrderedDict whose values have the given specs.

pytree.treespec_ordereddict(a=pytree.treespec_leaf()) TreeSpec(OrderedDict({'a': *}))

treespec_path_map​

treespec_path_map(treespec: 'TreeSpec') -> 'dict[KeyPath, int]'

Map the key path of every leaf of treespec to its flat index.

from clika_runtime import pytree {pytree.keystr(kp): i for kp, i in pytree.treespec_path_map(pytree.tree_structure({"a": (1, 2), "b": 3})).items()} {"['a'][0]": 0, "['a'][1]": 1, "['b']": 2}

treespec_pprint​

treespec_pprint(treespec: 'TreeSpec') -> 'str'

Render treespec over several indented lines.

Containers whose one-line form fits in 80 columns stay on one line; wider ones put each child on its own line.

from clika_runtime import pytree print(pytree.treespec_pprint(pytree.tree_structure({"a": (1, 2), "b": 3}))) TreeSpec({'a': (*, *), 'b': *})

treespec_slices​

treespec_slices(treespec: 'TreeSpec') -> 'dict[KeyPath, slice]'

The flat leaf index range under every node of treespec, keyed by the node's key path (the empty path is the whole spec).

from clika_runtime import pytree slices = pytree.treespec_slices(pytree.tree_structure({"a": (1, 2), "b": 3})) slices[(pytree.MappingKey("a"),)], slices[()] (slice(0, 2, None), slice(0, 3, None))

treespec_subspecs​

treespec_subspecs(treespec: 'TreeSpec') -> 'TreeSpecSubspecs'

Map every path of treespec to the structure below it.

Returns (path_to_spec, entry_to_spec, path_to_entry): the spec below each raw entry path (the empty path is the whole spec), the spec below each top-level entry, and each path's last entry. A custom node below the root stays atomic (nothing under it is mapped); the root itself is always opened, so a dataclass root maps its fields, named by field.

from clika_runtime import pytree maps = pytree.treespec_subspecs(pytree.tree_structure({"a": (1, 2), "b": 3})) sorted(maps.path_to_spec) [(), ('a',), ('a', 0), ('a', 1), ('b',)] maps.entry_to_spec["a"] TreeSpec((*, *))

treespec_tuple​

treespec_tuple(iterable: 'Iterable[TreeSpec]' = (), *, none_is_leaf: 'bool' = True, namespace: 'str' = '') -> 'TreeSpec'

The spec of a tuple whose items have the given specs.

pytree.treespec_tuple([pytree.treespec_leaf(), pytree.tree_structure([1, 2])]) TreeSpec((, [, *]))

treespec_widest​

treespec_widest(treespecs: 'Iterable[TreeSpec]') -> 'TreeSpec'

The spec with the most leaves; the first one on a tie.

This reconciles the structures a binding sees across sample inputs (one sample passes past_key_values=None, another a full cache): the widest one becomes the template every call is filled up to. Each other spec is expected to be a prefix of the result (a leaf where the widest has a subtree); a spec that is not one is reported.

Raises: ValueError: If treespecs is empty, or a spec is not a prefix of the widest one (the message names both).

unregister_leaf_meta​

unregister_leaf_meta(cls: 'type', *, namespace: 'str' = '') -> 'None'

Remove the leaf-metadata declaration of cls under namespace.

unregister_pytree_node​

unregister_pytree_node(cls: 'type', *, namespace: 'str' = '') -> 'NodeRegistryEntry'

Remove the registration of cls under namespace; instances become leaves again (for namespaced registrations, only under that namespace).

Returns: The removed entry.

Raises: TypeError: If cls is not a class or namespace is not a string. ValueError: If cls is a built-in node type or is not registered in namespace.