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 givingUT3Framea 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]