Loss pipeline

class stix.training.loss_pipeline.LossPipeline(coupling=None, time_sampler=None, antithetical=False)

Bases: object

Base class for loss pipelines.

The loss pipeline orchestrates the computation of the Monte Carlo estimator of the training loss for a batch of data.

It performs the following steps: - Embeds the raw data. - Samples the embedded source for two-sided modalities that have an embedded source prior. - Applies the coupling strategy to the raw and embedded pairs. - Generates the noise for each modality. - Interpolates the embedded pairs. - Computes the network output for each modality. - Computes the loss for each modality. - Returns the mean loss and the metrics.

Parameters:
__call__(gen_model, batch, key)

Compute the training loss for a batch of data.

Parameters:
  • gen_model (GenerativeModel) – The generative model (nnx.Module). Must be passed explicitly so that nnx.value_and_grad can trace its parameters.

  • batch (Batch) – A Batch carrying the raw training batch, context data/masks, attention masks and loss masks.

  • key (Array) – PRNG key for stochasticity.

Returns:

A pair (loss, metrics). metrics is a dict that always contains "loss" and may be extended by subclasses with extra diagnostics (e.g. per-modality losses, variance estimates).

Return type:

tuple[Scalar, PyTree[Scalar]]