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: Protocol

Protocol for guidance recipes.

A guidance recipe builds per-modality Generator objects from a given GenerativeModel, embedded state \(z_t\), time \(t\) and conditioning data.

The solvers (Solver and its sibling ManualSolver) 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 the guidance_fn field of SolverConfig.

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.

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

  • 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:

PyTree[stix.core.generator.Generator]

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.

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

  • 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:

PyTree[stix.core.generator.Generator]

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’s get_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 a VelocityAndScore generator 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.

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

  • 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:
Return type:

PyTree[stix.core.generator.Generator]

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 a VelocityAndScore generator 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.

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

  • 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 implement velocity_from_score().

Return type:

PyTree[stix.core.generator.Generator]

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.

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

  • 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:

PyTree[stix.core.generator.Generator]