common#

Backend infrastructure: numpy/jax dispatch, array predicates, scans/maps, shared mixins.

get_backend(is_uniform, use_jax) -> (xnp, xmap, xscan) selects the array module and loop machinery; is_ndarray/is_jax_ndarray/tree_contains_jax are the type-inference predicates behind the no-use_jax-parameter convention (dispatch is inferred from the input arrays at the lowest level); ValueHashedMasks is the value-based hash/eq mixin that keeps a rebuilt-but-identical mask holder on the same jit cache key.

Attributes#

Classes#

ExplicitEquality

Mixin making == a directed error and the class unhashable: equality must be EXPLICIT.

ValueHashedMasks

Mixin giving uniform-layer mask holders VALUE-based __hash__/__eq__ over mask content.

ValueHashedFields

Mixin giving a frozen dataclass VALUE-based __hash__/__eq__ over all its fields,

StaticSkeleton

The static half of a partitioned data tuple -- tree shape plus the static values, hashed by value

Functions#

is_boolean_ndarray(x)

jax_or_warn(what)

The jax-absent policy: a request for jax (use_jax=True, to_jax, use_jit=True) on a machine

to_jax(x)

jnp.array(x) when jax is available; otherwise np.array(x) with a warning (see jax_or_warn).

ragged_scan(f, init, xs)

Similar to jax.lax.scan, except for ragged-sized arrays

numpy_scan(f, init, xs)

Similar to jax.lax.scan, except returns numpy arrays instead of jax arrays.

ragged_map(f, xs)

numpy_map(f, xs)

get_backend(is_uniform, use_jax)

xwhile(cond, body, init_state[, use_jit])

Data-dependent while with the numpy / eager-jax / jit dispatch -- the xscan precedent for a

xcat(x, y)

Concatenate arrays or sequences.

xappend(S, x)

Append slice to array or element to sequence

xprepend(x, S)

Prepend slice to array or element to sequence

randn(*args, use_jax)

tree_contains_jax(T)

tree_to_jax(T)

Move every array leaf of a pytree (nested tuples/lists of arrays) onto jax, preserving the tree

items_are_uniform(xx)

Checks if an object can be treated as uniform for the purposes of jax.scan and jax.map.

save_core_families(file, families)

Save a sequence of core-families (each a sequence of arrays) to a .npz file.

load_core_families(file)

Inverse of save_core_families(): load a .npz file into a tuple of core-families.

readonly_mask_copies(masks)

Defensive READ-ONLY numpy copies of a mask tuple, for storage in value-hashed aux objects:

partition_static(tree)

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

rebuild_static(dynamic, skeleton)

Reassemble the tuple tree partition_static() split, substituting dynamic in order.

require_concrete_masks(*masks)

Guard the uniform-mask contract: masks are concrete host (numpy) arrays, never jax tracers.

prefix_mask(ranks, pad)

Boolean prefix indicator: slot j is real iff j < rank -- the canonical (prefix) form.

Module Contents#

jax_available = False#
NDArray#
is_ndarray#
is_jax_ndarray#
is_numpy_ndarray#
to_numpy#
jax_scan#
jax_map#