Guidance¶
Guidance recipes for conditional sampling. Each recipe returns the per-modality
Generator consumed by
Solver and
ManualSolver via their
guidance_fn config.
- class stix.sampling.guidance.GuidanceFn(*args, **kwargs)¶
Bases:
ProtocolProtocol for guidance recipes.
A guidance recipe builds per-modality
Generatorobjects from a givenGenerativeModel, embedded state \(z_t\), time \(t\) and conditioning data.The solvers (
Solverand its siblingManualSolver) call the configured recipe once per integration step to get the generators that drive the step. Every recipe shares this one signature so recipes slot interchangeably into theguidance_fnfield ofSolverConfig.Conditioning data are passed through two independent channels:
context (
context_data/context_mask): a pytree fed to the generative model’s network forward. Typically used for classifier-free guidance.intrinsic (
intrinsic_data/intrinsic_mask): conditioning data that are not fed to the network forward. Typically, these are data encoding conditions on the desired target state (e.g. a desired part of an image for image generation).
This is a structural (callback) protocol, not a base class to subclass. Any plain function whose positional parameters align with
__call__satisfies it:def __call__( self, gen_model: GenerativeModel, z_t: PyTree[Var], t: Time, context_data: PyTree[Var] | None, context_mask: PyTree[Mask | None] | None, intrinsic_data: PyTree[Var] | None, intrinsic_mask: PyTree[Mask | None] | None, attention_mask: PyTree[Mask | None] | None, guidance_scale: PyTree[Callable[[Time], Scalar]], ) -> PyTree[Generator]:
To add a recipe you merely write a function such as:
def context_generator(gen_model, z_t, t, context_data, context_mask, *, attention_mask, guidance_scale): net_out = gen_model.get_network_output( z_t, t, context_data, context_mask, attention_mask, ) return gen_model.get_generator(net_out, z_t, t)
and select it via
SolverConfig(guidance_fn=context_generator).Help on defining custom implementations of this protocol can be found in the conditioning and guidance tutorial.
- __call__(gen_model, z_t, t, context_data=None, context_mask=None, intrinsic_data=None, intrinsic_mask=None, *, attention_mask=None, guidance_scale)¶
Return the per-modality generator at time
t.- Parameters:
gen_model (GenerativeModel) – Generative model computing the generators.
t (Time) – Current time.
context_data (PyTree[Var] | None) – Context conditioning data for the recipe.
context_mask (PyTree[Mask | None] | None) – Context masks paired with
context_data.intrinsic_data (PyTree[Var] | None) – Intrinsic-guidance conditioning data.
intrinsic_mask (PyTree[Mask | None] | None) – Intrinsic masks paired with
intrinsic_data.attention_mask (PyTree[Mask | None] | None) – Per-modality attention masks.
guidance_scale (PyTree[Callable[[Time], Scalar]]) – Per-modality time-dependent guidance scale.
- Returns:
Per-modality
Generator.- Return type:
- stix.sampling.guidance.get_context_conditional_generator(gen_model, z_t, t, context_data=None, context_mask=None, intrinsic_data=None, intrinsic_mask=None, *, attention_mask=None, guidance_scale)¶
Produce the generator from a single network forward, passing the context data to the network.
The intrinsic channel data is ignored.
Used as the default recipe in
SolverConfig.- Parameters:
gen_model (GenerativeModel) – Generative model computing the generators.
t (Time) – Current time.
context_data (PyTree[Var] | None) – Context conditioning data for the recipe.
context_mask (PyTree[Mask | None] | None) – Context masks paired with
context_data.intrinsic_data (PyTree[Var] | None) – Intrinsic-guidance conditioning data.
intrinsic_mask (PyTree[Mask | None] | None) – Intrinsic masks paired with
intrinsic_data.attention_mask (PyTree[Mask | None] | None) – Per-modality attention masks.
guidance_scale (PyTree[Callable[[Time], Scalar]]) – Per-modality time-dependent guidance scale. Ignored by this recipe.
- Returns:
Per-modality
Generator.- Return type:
- stix.sampling.guidance.get_intrinsic_guidance_generator(gen_model, z_t, t, context_data=None, context_mask=None, intrinsic_data=None, intrinsic_mask=None, *, attention_mask=None, guidance_scale)¶
Score-based classifier guidance recipe.
This recipe is only valid for generative models whose modalities’ generators are all of type
VelocityAndScore.Up to terms constant in \(z_t\), Bayes’ rule gives
\[s_{\text{cond}}(z_t, t) = \underbrace{\nabla_{z_t} \log p_t(z_t)}_{\text{unconditional score}} + \omega(t) \underbrace{\nabla_{z_t} \log p_t(c | z_t)}_{\text{guidance gradient}},\]where \(c\) is the intrinsic conditioning data.
The unconditional score is computed from the network output’s generator (
get_generator()); the guidance gradient is obtained by automatic differentiation of the model’sget_guidance_loss(), which is interpreted as an approximation to \(-\log p(c | z_t)\). (See that method for precise details and caveats!).The conditional score \(s_{\text{cond}}(z_t, t)\) is then used to re-derive a consistent velocity via the interpolant’s
velocity_from_score()method, and packaged into aVelocityAndScoregenerator that is returned by the recipe.The context data are not ignored, they are passed to the network forward.
For more information see the conditioning and guidance tutorial.
- Parameters:
gen_model (GenerativeModel) – Generative model computing the generators.
t (Time) – Current time.
context_data (PyTree[Var] | None) – Context conditioning data for the recipe.
context_mask (PyTree[Mask | None] | None) – Context masks paired with
context_data.intrinsic_data (PyTree[Var] | None) – Intrinsic-guidance conditioning data.
intrinsic_mask (PyTree[Mask | None] | None) – Intrinsic masks paired with
intrinsic_data.attention_mask (PyTree[Mask | None] | None) – Per-modality attention masks.
guidance_scale (PyTree[Callable[[Time], Scalar]]) – Per-modality time-dependent guidance scale.
- Returns:
Per-modality intrinsic-guided generator (a
VelocityAndScore).- Raises:
NotImplementedError – If
gen_modeldoes not overrideget_guidance_loss(): the recipe differentiates that model-side hook, so a model without it cannot drive intrinsic guidance.TypeError – If a modality’s generator is not a
VelocityAndScore, or if its interpolant does not implementvelocity_from_score().
- Return type:
- stix.sampling.guidance.get_classifier_free_guidance_generator(gen_model, z_t, t, context_data=None, context_mask=None, intrinsic_data=None, intrinsic_mask=None, *, attention_mask=None, guidance_scale)¶
Score-based classifier-free guidance recipe.
This recipe is only valid for generative models whose modalities’ generators are all of type
VelocityAndScore.Up to terms constant in \(z_t\), Bayes’ rule gives
\[s_{\text{cond}}(z_t, t) \approx \omega(t)\, s(z_t, t, \text{context}=c) + (1 - \omega(t))\, s(z_t, t, \text{context}=\text{null}).\]The two scores are computed from the network output’s generators, the first one using the context data, the second one using no context data (i.e.
context_data=None, context_mask=None).The combined score is then used to re-derive a consistent velocity via the interpolant’s
velocity_from_score()method, and packaged into aVelocityAndScoregenerator that is returned by the recipe.The intrinsic channel data is ignored.
For more information see the conditioning and guidance tutorial.
- Parameters:
gen_model (GenerativeModel) – Generative model computing the generators.
t (Time) – Current time.
context_data (PyTree[Var] | None) – Context conditioning data for the recipe.
context_mask (PyTree[Mask | None] | None) – Context masks paired with
context_data.intrinsic_data (PyTree[Var] | None) – Intrinsic-guidance conditioning data.
intrinsic_mask (PyTree[Mask | None] | None) – Intrinsic masks paired with
intrinsic_data.attention_mask (PyTree[Mask | None] | None) – Per-modality attention masks.
guidance_scale (PyTree[Callable[[Time], Scalar]]) – Per-modality time-dependent guidance scale.
- Returns:
Per-modality classifier-free-guided generator (a
VelocityAndScore).- Raises:
TypeError – If a modality’s generator is not a
VelocityAndScore, or if its interpolant does not implementvelocity_from_score().- Return type:
- stix.sampling.guidance.get_discrete_classifier_free_guidance_generator(gen_model, z_t, t, context_data=None, context_mask=None, intrinsic_data=None, intrinsic_mask=None, *, attention_mask=None, guidance_scale)¶
Classifier-free guidance recipe for discrete modalities, in logit space.
This recipe is only valid for generative models whose modalities’ generators are all of type
TransitionRates. It also assumes that the network outputs the logits of the embedded target posterior distribution.The unconditional and conditional posterior distributions are combined through a tempered geometric mean
\[\tilde{p}_{\mathrm{tgt}|t}(\cdot \mid z_t, c) \propto p_{\mathrm{tgt}|t}(\cdot \mid z_t, c)^{\omega(t)}\, p_{\mathrm{tgt}|t}(\cdot \mid z_t, \varnothing)^{1-\omega(t)},\]which is obtained from a linear combination of the pre-softmax logits
\[\tilde{\hat{\ell}}(\mathrm{tgt}\mid z_t) = \omega(t)\,\hat{\ell}(\mathrm{tgt}\mid z_t, c) + (1 - \omega(t))\,\hat{\ell}(\mathrm{tgt}\mid z_t, \varnothing),\]where \(\hat{\ell}\) is the network output approximating the logits.
(The two logits are computed from the network outputs, the first one using the context data, the second one using no context data (i.e.
context_data=None, context_mask=None).)The intrinsic channel data is ignored.
For more information see the conditioning and guidance tutorial.
- Parameters:
gen_model (GenerativeModel) – Generative model computing the generators.
t (Time) – Current time.
context_data (PyTree[Var] | None) – Context conditioning data for the recipe.
context_mask (PyTree[Mask | None] | None) – Context masks paired with
context_data.intrinsic_data (PyTree[Var] | None) – Intrinsic-guidance conditioning data.
intrinsic_mask (PyTree[Mask | None] | None) – Intrinsic masks paired with
intrinsic_data.attention_mask (PyTree[Mask | None] | None) – Per-modality attention masks.
guidance_scale (PyTree[Callable[[Time], Scalar]]) – Per-modality time-dependent guidance scale.
- Returns:
Per-modality generator built from the guidance-combined network output (a
TransitionRates).- Return type: