Generative models¶
- class stix.core.gen_model.GenerativeModel(*args, **kwargs)¶
-
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 bystix.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.networkpositionally; the accepted call is theNetworkcontract, 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:
- 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}})\).
Noneraw source variables are mapped toNoneembedded source variables.- Return type:
- 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
Generatorper modality, and each modality’s generator must be an instance of its declaredmodality.generator_type(inferred from the interpolant).
- 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.
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:
- 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 aVelocityAndScoreis 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 implementget_lossalongside 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:
net_out (Any) – The network’s raw output for this sample.
t (Time) – Interpolation time, shared across modalities.
z_t (PyTree[Var]) – Per-modality interpolated state in embedded space.
epsilon (PyTree[NoiseVar]) – Per-modality noise sample drawn for the interpolation path.
embedded_pairs (PyTree[stix.typing.variables.EmbeddedSourceTargetPair]) – Per-modality embedded
(source, target)pairs.raw_pairs (PyTree[stix.typing.variables.RawSourceTargetPair]) – Per-modality raw-space
(source, target)pairs.loss_mask (PyTree[Mask]) – Per-modality loss mask (all-ones when unmasked).
- Returns:
A scalar loss for this sample.
- Return type:
- class stix.core.gen_model.factory.VelocityOneSidedGenerativeModel(*args, **kwargs)¶
Bases:
GenerativeModel,GenericOne-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 aVelocityAndScore.
get_guidance_loss()is not implemented here. To use this model withget_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:
net_out (Any) – The network’s raw output for this sample.
t (Time) – Interpolation time, shared across modalities.
z_t (PyTree[Var]) – Per-modality interpolated state in embedded space.
epsilon (PyTree[NoiseVar]) – Per-modality noise sample drawn for the interpolation path.
embedded_pairs (PyTree[stix.typing.variables.EmbeddedSourceTargetPair]) – Per-modality embedded
(source, target)pairs.raw_pairs (PyTree[stix.typing.variables.RawSourceTargetPair]) – Per-modality raw-space
(source, target)pairs.loss_mask (PyTree[Mask]) – Per-modality loss mask (all-ones when unmasked).
- Returns:
A scalar loss for this sample.
- Return type:
- 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_scaleisNone.
- class stix.core.gen_model.factory.VelocityTwoSidedGenerativeModel(*args, **kwargs)¶
Bases:
GenerativeModel,GenericTwo-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 aVelocity. 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:
net_out (Any) – The network’s raw output for this sample.
t (Time) – Interpolation time, shared across modalities.
z_t (PyTree[Var]) – Per-modality interpolated state in embedded space.
epsilon (PyTree[NoiseVar]) – Per-modality noise sample drawn for the interpolation path.
embedded_pairs (PyTree[stix.typing.variables.EmbeddedSourceTargetPair]) – Per-modality embedded
(source, target)pairs.raw_pairs (PyTree[stix.typing.variables.RawSourceTargetPair]) – Per-modality raw-space
(source, target)pairs.loss_mask (PyTree[Mask]) – Per-modality loss mask (all-ones when unmasked).
- Returns:
A scalar loss for this sample.
- Return type:
- class stix.core.gen_model.factory.NoiseOneSidedGenerativeModel(*args, **kwargs)¶
Bases:
GenerativeModel,GenericOne-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 noiseepsilondrawn 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 withget_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:
net_out (Any) – The network’s raw output for this sample.
t (Time) – Interpolation time, shared across modalities.
z_t (PyTree[Var]) – Per-modality interpolated state in embedded space.
epsilon (PyTree[NoiseVar]) – Per-modality noise sample drawn for the interpolation path.
embedded_pairs (PyTree[stix.typing.variables.EmbeddedSourceTargetPair]) – Per-modality embedded
(source, target)pairs.raw_pairs (PyTree[stix.typing.variables.RawSourceTargetPair]) – Per-modality raw-space
(source, target)pairs.loss_mask (PyTree[Mask]) – Per-modality loss mask (all-ones when unmasked).
- Returns:
A scalar loss for this sample.
- Return type:
- get_generator(net_out, z_t, t)¶
Derive a
VelocityAndScorefrom the predicted noise.Velocity and score are both derived from the predicted noise via the interpolant. Samplers ignore the score when
stochasticity_scaleisNone, resulting in a ODE solver.
- class stix.core.gen_model.factory.PosteriorMixtureGenerativeModel(*args, **kwargs)¶
Bases:
GenerativeModel,GenericAn all-discrete model over
DiscreteInterpolantmodalities withrates_from_target_posterior.The network head predicts target-posterior logits over the
Kdata 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-modalityTransitionRatesgenerators computed using the learned posterior logits and the interpolant’srates_from_target_posteriormethod.
- 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:
net_out (Any) – The network’s output, interpreted as posterior logits.
t (Time) – Interpolation time.
z_t (PyTree[Var]) – Per-modality interpolated state in embedded space.
epsilon (PyTree[NoiseVar]) – Per-modality noise sample drawn for the interpolation path.
embedded_pairs (PyTree[stix.typing.variables.EmbeddedSourceTargetPair]) – Per-modality embedded
(source, target)pairs.raw_pairs (PyTree[stix.typing.variables.RawSourceTargetPair]) – Per-modality raw-space
(source, target)pairs.loss_mask (PyTree[Mask]) – Per-modality loss mask (all-ones when unmasked).
- Returns:
A scalar loss for this sample.
- Return type:
- get_generator(net_out, z_t, t)¶
Turn target-posterior logits into per-modality
TransitionRates.Uses the interpolant’s
rates_from_target_posteriormethod.
- class stix.core.gen_model.factory.VelocityAndPosteriorGenerativeModel(*args, **kwargs)¶
Bases:
GenerativeModel,GenericJoint model over continuous one-sided-linear and discrete DFM modalities.
Continuous modalities behave like
VelocityOneSidedGenerativeModel(velocity MSE,VelocityAndScoregenerator); discrete modalities behave likePosteriorMixtureGenerativeModel(denoising cross-entropy,TransitionRatesgenerator). 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-sidedDiscreteInterpolantimplementingrates_from_target_posterior, and setnum_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:
net_out (Any) – The network’s raw output for this sample.
t (Time) – Interpolation time, shared across modalities.
z_t (PyTree[Var]) – Per-modality interpolated state in embedded space.
epsilon (PyTree[NoiseVar]) – Per-modality noise sample drawn for the interpolation path.
embedded_pairs (PyTree[stix.typing.variables.EmbeddedSourceTargetPair]) – Per-modality embedded
(source, target)pairs.raw_pairs (PyTree[stix.typing.variables.RawSourceTargetPair]) – Per-modality raw-space
(source, target)pairs.loss_mask (PyTree[Mask]) – Per-modality loss mask (all-ones when unmasked).
- Returns:
A scalar loss for this sample.
- Return type:
- get_generator(net_out, z_t, t)¶
Per-modality
VelocityAndScore(continuous) orTransitionRates(discrete).