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.