Generative models

class stix.core.gen_model.GenerativeModel(*args, **kwargs)

Bases: Module, ABC, Generic

Generative model base class.

GenerativeModel objects contain everything needed to perform sampling. It is responsible for:

  • Embedding the data

  • Computing the network output

  • Computing the sampling-time generator (velocity, velocity+score, or rates)

  • Computing the loss

To support classifier-style intrinsic guidance, the model must override the method get_guidance_loss() (consumed by stix.sampling.guidance.get_intrinsic_guidance_generator()).

Parameters:
  • args (Any)

  • kwargs (Any)

Return type:

Any

get_network_output(z_t, t, context_data, context_mask, attention_mask)

Get the output from the neural network.

Calls self.network positionally; the accepted call is the Network contract, which every network must satisfy.

Parameters:
  • z_t (PyTree[Var]) – Per-modality interpolated state in embedded space, at time t.

  • t (Time) – Interpolation time.

  • context_data (PyTree[Var] | None) – Context conditioning (a pytree) consumed by the network’s context encoder.

  • context_mask (PyTree[Mask | None] | None) – Context masks gating context_data.

  • attention_mask (PyTree[Mask | None] | None) – Per-modality attention masks.

Returns:

The network’s raw output (structure defined by the concrete network).

Return type:

Any

get_embeddings(raw_batch)

Embed raw data \((x_{\mathrm{src}}, x_{\mathrm{tgt}})\) to produce the embedded pairs \((z_{\mathrm{src}}, z_{\mathrm{tgt}})\).

Parameters:

raw_batch (PyTree[stix.typing.variables.RawSourceTargetPair]) – Per-modality raw (source, target) pairs \((x_{\mathrm{src}}, x_{\mathrm{tgt}})\).

Returns:

Per-modality embedded pairs \((z_{\mathrm{src}}, z_{\mathrm{tgt}})\). None raw source variables are mapped to None embedded source variables.

Return type:

PyTree[stix.typing.variables.EmbeddedSourceTargetPair]

abstractmethod get_generator(net_out, z_t, t)

Convert network output to per-modality sampling generators at \((z_t, t)\).

The returned pytree must have one Generator per modality, and each modality’s generator must be an instance of its declared modality.generator_type (inferred from the interpolant).

Parameters:
  • net_out (Any) – The network’s raw output.

  • z_t (PyTree[Var]) – Per-modality noisy state in embedded space.

  • t (Time) – Current time, \(t \in [0, 1]\).

Returns:

Per-modality sampling-time generators.

Return type:

PyTree[stix.core.generator.Generator]

get_guidance_loss(net_out, z_t, t, intrinsic_data, intrinsic_mask)

An optional scalar loss used for conditional intrinsic guidance.

It is consumed by stix.sampling.guidance.get_intrinsic_guidance_generator() at sampling time to guide the sampling process towards the conditional data.

More precisely, the gradient of this loss is used as an estimate of the conditional score \(\nabla_{z_t}\log p(c | z_t) \simeq \nabla_{z_t} \ell (c , z_t)\).

This method is not implemented by default, as the exact form of the guidance loss typically depends on the interpretation of the network output and the embedded variables, which are specified by the generative model. See the conditioning and guidance tutorial for a worked implementation.

Parameters:
  • net_out (Any) – Per-modality network output for this sample.

  • z_t (PyTree[Var]) – Per-modality state in embedded space.

  • t (Time) – Current time, \(t \in [0, 1]\).

  • intrinsic_data (PyTree[Var]) – Per-modality conditioning targets, one per registry modality.

  • intrinsic_mask (PyTree[Mask | None]) – Per-modality 0/1 masks; None (whole or per-entry) is treated as all-ones (no masking).

Returns:

A scalar guidance loss for this sample.

Return type:

Scalar

abstractmethod get_loss(net_out, t, z_t, epsilon, embedded_pairs, raw_pairs, loss_mask)

Compute the scalar training loss for one interpolated sample.

The loss defines what the network learns to predict. It is therefore inextricably linked to get_generator().

For example, if the network output is trained against the conditional velocity, then get_generator (typically) wraps the network output directly as the velocity, deriving the score from it when a VelocityAndScore is required; if it is trained against some other target, that conversion changes accordingly. For this reason there is no sensible default — each concrete model must implement get_loss alongside its prediction-to-generator conversion. For more details regarding the implementation of a tailored GenerativeModel, see the generative model tutorial.

A common pattern is to compute one scalar loss per modality and reduce with reduce_modality_losses():

def get_loss(self, net_out, t, z_t, epsilon,
             embedded_pairs, raw_pairs, loss_mask):
    criterion = MSECriterion()
    losses = self.modality_registry.map(
        lambda modality, net_out_k, pair_k, eps_k, mask_k: criterion(
            net_out_k,
            modality.interpolant.get_conditional_velocity(
                pair_k, t, eps_k
            ),
            t,
            mask_k,
        ),
        net_out, embedded_pairs, epsilon, loss_mask,
    )
    return reduce_modality_losses(losses)

Losses that cannot be decomposed per modality (e.g. cross-modality / coupled losses) are equally valid — simply implement them directly here rather than mapping over the registry.

Parameters:
Returns:

A scalar loss for this sample.

Return type:

Scalar

class stix.core.gen_model.factory.VelocityOneSidedGenerativeModel(*args, **kwargs)

Bases: GenerativeModel, Generic

One-sided model whose network output is the (embedded) velocity field.

A ready-made, predefined GenerativeModel: the single choice to train the network on the velocity fixes both methods the rest of the library relies on.

  • get_loss() — a per-modality MSE between the network output and the interpolant’s conditional velocity, reduced across modalities.

  • get_generator() — the network output already is the velocity; the score is derived from it via the interpolant and both are returned as a VelocityAndScore.

get_guidance_loss() is not implemented here. To use this model with get_intrinsic_guidance_generator(), subclass it and override that hook. See the conditioning and guidance tutorial for a worked implementation.

The loss is a velocity MSE (MSECriterion) — that is what makes this a velocity model, so it is not configurable here.

Parameters:
  • args (Any)

  • kwargs (Any)

Return type:

Any

get_loss(net_out, t, z_t, epsilon, embedded_pairs, raw_pairs, loss_mask)

Velocity MSE: compare the network output to the conditional velocity.

One scalar loss is computed per modality (the interpolant provides the ground-truth conditional velocity), then reduced across modalities with reduce_modality_losses().

Parameters:
Returns:

A scalar loss for this sample.

Return type:

Scalar

get_generator(net_out, z_t, t)

Wrap the predicted velocity as a VelocityAndScore.

The network output already is the velocity; the score is derived from it via the interpolant. ODE sampling ignores the score when stochasticity_scale is None.

Parameters:
  • net_out (Any) – The network’s raw output.

  • z_t (PyTree[Var]) – Per-modality noisy state in embedded space.

  • t (Time) – Current time, \(t \in [0, 1]\).

Returns:

Per-modality sampling-time generators.

Return type:

PyTree[stix.core.generator.Generator]

class stix.core.gen_model.factory.VelocityTwoSidedGenerativeModel(*args, **kwargs)

Bases: GenerativeModel, Generic

Two-sided model whose network output is the (embedded) velocity field.

The two-sided counterpart of VelocityOneSidedGenerativeModel: source and target are both data distributions, transported into one another, rather than noise into data.

  • get_loss() — a per-modality MSE between the network output and the interpolant’s conditional velocity, reduced across modalities.

  • get_generator() — wraps the network output as a Velocity. A deterministic interpolant carries no noise term, so no score exists and sampling is ODE-only (stochasticity_scale=None); that generator type is inferred from the interpolant.

The loss is a velocity MSE (MSECriterion) — that is what makes this a velocity model, so it is not configurable here.

Parameters:
  • args (Any)

  • kwargs (Any)

Return type:

Any

get_loss(net_out, t, z_t, epsilon, embedded_pairs, raw_pairs, loss_mask)

Velocity MSE: compare the network output to the conditional velocity.

One scalar loss is computed per modality (the interpolant provides the ground-truth conditional velocity), then reduced across modalities with reduce_modality_losses().

Parameters:
Returns:

A scalar loss for this sample.

Return type:

Scalar

get_generator(net_out, z_t, t)

Wrap the predicted velocity as a Velocity generator.

A deterministic interpolant carries no noise term and therefore defines no score, so the inferred generator type is Velocity; SDE sampling is unavailable and the process is an ODE.

Parameters:
  • net_out (Any) – The network’s raw output.

  • z_t (PyTree[Var]) – Per-modality noisy state in embedded space.

  • t (Time) – Current time, \(t \in [0, 1]\).

Returns:

Per-modality sampling-time generators.

Return type:

PyTree[stix.core.generator.Generator]

class stix.core.gen_model.factory.NoiseOneSidedGenerativeModel(*args, **kwargs)

Bases: GenerativeModel, Generic

One-sided model whose network output is the (embedded) noise \(\epsilon\).

A ready-made, predefined GenerativeModel: the single choice to train the network on the noise fixes both methods the rest of the library relies on.

  • get_loss() — a per-modality MSE between the network output and the noise epsilon drawn for the interpolation path, reduced across modalities.

  • get_generator() — derives the velocity and score from the predicted noise via each modality’s (linear) interpolant.

get_guidance_loss() is not implemented here. To use this model with get_intrinsic_guidance_generator(), subclass it and override that hook. See the conditioning and guidance tutorial for a worked implementation.

The conversions above are only valid for a one-sided linear interpolant, so the constructor validates that every modality’s interpolant is an OneSidedLinearStochasticInterpolant.

The loss is a noise MSE (MSECriterion) — that is what makes this a noise model, so it is not configurable here.

Parameters:
  • args (Any)

  • kwargs (Any)

Return type:

Any

get_loss(net_out, t, z_t, epsilon, embedded_pairs, raw_pairs, loss_mask)

Noise MSE: compare the network output to the sampled noise epsilon.

One scalar loss is computed per modality, then reduced across modalities with reduce_modality_losses().

Parameters:
Returns:

A scalar loss for this sample.

Return type:

Scalar

get_generator(net_out, z_t, t)

Derive a VelocityAndScore from the predicted noise.

Velocity and score are both derived from the predicted noise via the interpolant. Samplers ignore the score when stochasticity_scale is None, resulting in a ODE solver.

Parameters:
  • net_out (Any) – The network’s raw output.

  • z_t (PyTree[Var]) – Per-modality noisy state in embedded space.

  • t (Time) – Current time, \(t \in [0, 1]\).

Returns:

Per-modality sampling-time generators.

Return type:

PyTree[stix.core.generator.Generator]

class stix.core.gen_model.factory.PosteriorMixtureGenerativeModel(*args, **kwargs)

Bases: GenerativeModel, Generic

An all-discrete model over DiscreteInterpolant modalities with rates_from_target_posterior.

The network head predicts target-posterior logits over the K data categories (for a mask interpolant the mask symbol is excluded from the logits).

This model fixes:

  • get_loss() — cross-entropy of the posterior logits against the embedded target one-hot.

  • get_generator() — per-modality TransitionRates generators computed using the learned posterior logits and the interpolant’s rates_from_target_posterior method.

Parameters:
  • args (Any)

  • kwargs (Any)

Return type:

Any

get_loss(net_out, t, z_t, epsilon, embedded_pairs, raw_pairs, loss_mask)

Cross-entropy of the posterior logits against the embedded target one-hot.

Parameters:
Returns:

A scalar loss for this sample.

Return type:

Scalar

get_generator(net_out, z_t, t)

Turn target-posterior logits into per-modality TransitionRates.

Uses the interpolant’s rates_from_target_posterior method.

Parameters:
  • net_out (Any) – The network’s raw output.

  • z_t (PyTree[Var]) – Per-modality noisy state in embedded space.

  • t (Time) – Current time, \(t \in [0, 1]\).

Returns:

Per-modality sampling-time generators.

Return type:

PyTree[stix.core.generator.TransitionRates]

class stix.core.gen_model.factory.VelocityAndPosteriorGenerativeModel(*args, **kwargs)

Bases: GenerativeModel, Generic

Joint model over continuous one-sided-linear and discrete DFM modalities.

Continuous modalities behave like VelocityOneSidedGenerativeModel (velocity MSE, VelocityAndScore generator); discrete modalities behave like PosteriorMixtureGenerativeModel (denoising cross-entropy, TransitionRates generator). Dispatch is based on the modality’s interpolant type, resolved per modality.

Continuous modality’s must have interpolant of type OneSidedLinearStochasticInterpolant. Discrete modalities must have a one-sided DiscreteInterpolant implementing rates_from_target_posterior, and set num_categories.

Parameters:
  • args (Any)

  • kwargs (Any)

Return type:

Any

get_loss(net_out, t, z_t, epsilon, embedded_pairs, raw_pairs, loss_mask)

Per-modality velocity MSE (continuous) or weighted CE (discrete).

Parameters:
Returns:

A scalar loss for this sample.

Return type:

Scalar

get_generator(net_out, z_t, t)

Per-modality VelocityAndScore (continuous) or TransitionRates (discrete).

Parameters:
  • net_out (Any) – The network’s raw output.

  • z_t (PyTree[Var]) – Per-modality noisy state in embedded space.

  • t (Time) – Current time, \(t \in [0, 1]\).

Returns:

Per-modality sampling-time generators.

Return type:

PyTree[stix.core.generator.Generator]