Criteria

class stix.core.loss.Criterion

Bases: ABC

Abstract per-modality criterion: a bare prediction-vs-ground-truth comparison.

abstractmethod __call__(prediction, ground_truth, t, mask)

Compute the per-modality loss for a single prediction/ground-truth pair.

Parameters:
Return type:

Scalar

class stix.core.loss.MSECriterion

Bases: Criterion

Mean squared error criterion for continuous modalities.

__call__(prediction, ground_truth, t, mask)

Compute masked mean squared error between prediction and ground-truth.

Parameters:
  • prediction (Var) – Predicted variable, of shape (*dims, dim).

  • ground_truth (Var) – Ground-truth variable, of shape (*dims, dim).

  • t (Time) – Time step.

  • mask (Mask) – Mask, of shape (*dims, dim).

Returns:

The masked MSE loss averaged over the unmasked entries. Note LossPipeline vmap s over the batch, so this sees one sample, not a batch.

Return type:

Scalar

class stix.core.loss.CrossEntropyCriterion

Bases: Criterion

Cross entropy for discrete modalities, computed on logits prediction vs one-hot ground-truth labels.

__call__(prediction, ground_truth, t, mask)

Compute masked softmax cross entropy between predicted logits and one-hot ground-truth labels.

Use optax.softmax_cross_entropy to compute the standard per-example cross entropy

\[-\sum_k \mathrm{label}_k \cdot \mathrm{log-softmax}(\mathrm{logits})_k\]

and returns its average over the unmasked entries.

Parameters:
  • prediction (Var) – Predicted logits, of shape (*dims, num_categories).

  • ground_truth (Var) – One-hot ground-truth labels, of shape (*dims, num_categories).

  • t (Time) – Time step.

  • mask (Mask) – Mask, of shape (*dims, num_categories).

Returns:

The softmax cross entropy loss averaged over the unmasked entries. Note LossPipeline vmap s over the batch, so this sees one sample, not a batch.

Return type:

Scalar

Loss utilities

stix.core.loss.reduce_modality_losses(losses, weights=None)

Reduce a per-modality pytree of scalar losses to a single scalar.

A small helper for writing GenerativeModel.get_loss: compute one scalar loss per modality (typically via modality_registry.map), then collapse them here. Without weights this is a plain mean over the modality leaves; with weights it is the weight-normalised mean

\[\mathcal{L} = \frac{\sum_i w_i\,\mathcal{L}_i}{\sum_j w_j}.\]
Parameters:
  • losses (PyTree[Scalar]) – A pytree whose leaves are per-modality scalar losses (e.g. the output of modality_registry.map(...)).

  • weights (PyTree[Scalar | float] | None) – Optional pytree of per-modality weights, structurally compatible with losses; each leaf is a scalar (a Python float or a JAX scalar). This is how a custom get_loss weights different modalities. None weights every modality equally.

Returns:

The reduced scalar loss.

Return type:

Scalar