Criteria¶
- class stix.core.loss.Criterion¶
Bases:
ABCAbstract per-modality criterion: a bare prediction-vs-ground-truth comparison.
- class stix.core.loss.MSECriterion¶
Bases:
CriterionMean squared error criterion for continuous modalities.
- __call__(prediction, ground_truth, t, mask)¶
Compute masked mean squared error between prediction and ground-truth.
- Parameters:
- Returns:
The masked MSE loss averaged over the unmasked entries. Note
LossPipelinevmaps over the batch, so this sees one sample, not a batch.- Return type:
- class stix.core.loss.CrossEntropyCriterion¶
Bases:
CriterionCross 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_entropyto 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:
- Returns:
The softmax cross entropy loss averaged over the unmasked entries. Note
LossPipelinevmaps over the batch, so this sees one sample, not a batch.- Return type:
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 viamodality_registry.map), then collapse them here. Withoutweightsthis is a plain mean over the modality leaves; withweightsit 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 Pythonfloator a JAX scalar). This is how a customget_lossweights different modalities.Noneweights every modality equally.
- Returns:
The reduced scalar loss.
- Return type: