partition_static#

t3toolbox.backend.common.partition_static(tree)#
def partition_static(
        tree,   # a tuple/list tree of arrays and static structure (a frame's .data, a frame sweep, ...)
) -> typ.Tuple[
    typ.Tuple,        # the dynamic leaves, in traversal order -- the jax pytree children
    StaticSkeleton,   # the tree shape + static values -- the jax pytree aux_data
]:

Split a raw backend data tuple into what jax should TRACE and what it must keep STATIC.

Backend data tuples mix both – a uniform frame is (4 supercores, shape, masks) – and a bare tuple is a jax pytree whose every element is a leaf, so flattening one naively traces the masks and the shape. That raises (require_concrete_masks) or silently produces a program that cannot do host-integer shape arithmetic. The frontend avoids it by giving UT3Frame a registered pytree with the masks as aux; the backend keeps plain tuples, so the split happens here instead.

Round-trips exactly through rebuild_static().

Return type:

Tuple[Tuple, StaticSkeleton]