Utilities

stix.core.utils.check_gamma_is_not_zero(gamma_fn)

Check that gamma_fn is non-zero on the interior of [0, 1].

Evaluated on a 1000-point grid excluding the endpoints, so a True result is necessary but not sufficient.

Parameters:

gamma_fn (Callable[[Time], Scalar])

Return type:

bool

stix.core.utils.is_one_hot_vector(x)

Checks if all vectors along the last axis of x are valid one-hot vectors.

One-hot vectors are detected by checking that: 1. They hold a non-integer dtype, 2. They hold binary values (through jnp.logical_or(x == 0, x == 1)), 3. They sum to 1.

Returns a single scalar JAX boolean.

NOTE: One-hot encodings are produced by jax.nn.one_hot(), which yields a float dtype.

Parameters:

x (Array) – The input array to check.

Returns:

A scalar JAX boolean indicating whether each vector along the last axis is a valid one-hot encoding.

Return type:

Array

stix.core.utils.infer_num_samples_from_raw_source(raw_source)

Leading batch size shared by all non-None raw_source leaves.

Parameters:

raw_source (PyTree) – Pytree of raw source arrays or None.

Returns:

The common leading axis, or None if every leaf is None.

Raises:

ValueError – If non-None leaves disagree on the leading axis.

Return type:

int | None