Loss pipeline¶
- class stix.training.loss_pipeline.LossPipeline(coupling=None, time_sampler=None, antithetical=False)¶
Bases:
objectBase 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:
coupling (Coupling | None)
time_sampler (TimeSampler | None)
antithetical (bool)
- __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 thatnnx.value_and_gradcan trace its parameters.batch (Batch) – A
Batchcarrying the raw training batch, context data/masks, attention masks and loss masks.key (Array) – PRNG key for stochasticity.
- Returns:
A pair
(loss, metrics).metricsis a dict that always contains"loss"and may be extended by subclasses with extra diagnostics (e.g. per-modality losses, variance estimates).- Return type: