Utilities¶
- stix.core.utils.check_gamma_is_not_zero(gamma_fn)¶
Check that
gamma_fnis non-zero on the interior of[0, 1].Evaluated on a 1000-point grid excluding the endpoints, so a
Trueresult is necessary but not sufficient.
- stix.core.utils.is_one_hot_vector(x)¶
Checks if all vectors along the last axis of
xare 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.
- stix.core.utils.infer_num_samples_from_raw_source(raw_source)¶
Leading batch size shared by all non-
Noneraw_sourceleaves.- Parameters:
raw_source (PyTree) – Pytree of raw source arrays or
None.- Returns:
The common leading axis, or
Noneif every leaf isNone.- Raises:
ValueError – If non-
Noneleaves disagree on the leading axis.- Return type:
int | None