Manual solver

The fixed-step ODE/SDE/CTMC solver.

class stix.sampling.solver_manual.ManualSolverConfig(stochasticity_scale=None, direction=Direction.FORWARD, num_steps=300, t_min_tolerance=0.0001, t_max_tolerance=0.0001, guidance_fn=None, guidance_scale=<factory>, resolve_terminal_mask=True)

Bases: object

Configuration for the fixed-step solver.

Parameters:
  • stochasticity_scale (PyTree[Callable[[Time], Scalar] | None] | None) – Per-modality time-dependent stochasticity scale. Each callable maps \(t\) to \(\lambda_t\in\mathbb{R}_{\geq 0}\), and None (default, or a None leaf) is \(\lambda\equiv 0\). For a VelocityAndScore modality it is the SDE scale, so None means ODE sampling. For a TransitionRates modality it is the CTMC corrector scale of \(\bar u_t = (1 + \lambda_t)\,\hat u_t + \lambda_t\,\check u_t\), so None means the plain generative chain (see Introduction for the notations and more details). Setting it for a modality whose generator_type carries no score (Velocity), or setting a callable that is negative on [0, 1], raises a ValueError.

  • direction (Direction) – Integration direction (“forward” or “reverse”). Defaults to “forward”. Forward integrates t=0 -> t=1, reverse t=1 -> t=0. For TransitionRates modalities, reverse swaps the two corrector weights, namely \((1 + \lambda_t)\,\hat u_t + \lambda_t\,\check u_t`\) becomes \(\lambda_t\,\hat u_t + (1 + \lambda_t)\,\check u_t\), so the same \(\lambda_t\) time-reverses the CTMC.

  • num_steps (int) – Number of fixed steps. Defaults to 300.

  • t_min_tolerance (float) – Pad above t=0 (integration starts/ends at 0 + t_min_tolerance). Defaults to 1e-4.

  • t_max_tolerance (float) – Pad below t=1 (integration starts/ends at 1 - t_max_tolerance). Defaults to 1e-4.

  • guidance_fn (GuidanceFn | None) – A GuidanceFn recipe returning the per-modality Generator. None (default) means no guidance: the solver falls back to get_context_conditional_generator(), a single network forward that is unconditional (with context_data).

  • guidance_scale (PyTree[Callable[[Time], Scalar]]) – Per-modality time-dependent guidance scale. Defaults to an empty dict (sufficient for the default guidance_fn, which ignores it).

  • resolve_terminal_mask (bool) – If True (default), a mask-source discrete modality (masked diffusion) whose CTMC left some positions on the mask symbol has them resolved after the final step, by argmax of the terminal denoising posterior. This parameter is a no-op for modalities without a MaskDiscreteInterpolant, so it is safe to leave on. It is also a no-op on REVERSE: a mask interpolant has data at \(t=1\).

class stix.sampling.solver_manual.ManualSolver(solver_config)

Bases: object

Unconditional or guided fixed-step SDE/ODE/CTMC solver for the generative model.

A fixed-step integrator that is easy to inspect and debug, and compatible with jax.vmap over samples. Each step calls the configured guidance_fn recipe for the per-modality Generator, then applies the update rule STEP_FNS holds for that modality’s generator_type:

  • ODE path (stochasticity_scale is None): Euler on the velocity.

  • SDE path: Euler-Maruyama, drift = velocity + score term. Only VelocityAndScore modalities may carry a stochasticity_scale.

  • CTMC path: Euler on the transition probabilities of TransitionRates modalities, where stochasticity_scale weighs the backward rates into the corrector.

This is the only solver that handles continuous SDEs and discrete CTMCs simultaneously — Solver rejects a TransitionRates modality outright — so it is the solver for discrete and mixed-modality sampling.

Guidance is opt-in: with the default guidance_fn=None the recipe is a single plain network forward, so an unconfigured solver performs unconditional sampling (context-conditional if you pass context_data). Set guidance_fn to a recipe from stix.sampling.guidance for e.g. classifier-free or intrinsic guidance.

The gen_model is split internally: the graph definition is captured in closures (static) and the state flows as a dynamic pytree.

Note

These three rules are the only schemes this class can express. The network is evaluated once per step, before dispatch, so no step function can re-evaluate it mid-step: Heun, midpoint and higher-order stochastic Runge-Kutta are out of reach by construction. See Solver for higher-order integration of purely continuous modalities.

Parameters:

solver_config (ManualSolverConfig)

__call__(gen_model, z_init, key, *, context_data=None, context_mask=None, intrinsic_data=None, intrinsic_mask=None, attention_mask=None)

Integrate the SDE/ODE/CTMC and return decoded samples.

Safe to use under jax.vmap — the gen_model state is traced, the graph definition is captured in the closure.

Parameters:
  • gen_model (GenerativeModel) – The generative model (owns the network and the modality registry).

  • z_init (PyTree[Var]) – Initial samples in embedding space, one array per modality.

  • key (Array) – PRNG key for the Brownian increments and CTMC draws; split per step and per modality (deterministic ODE steps ignore theirs).

  • context_data (PyTree[Var] | None) – Context conditioning consumed by the configured guidance_fn recipe, forwarded to the network. See stix.sampling.guidance for the context contract.

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

  • intrinsic_data (PyTree[Var] | None) – Per-modality intrinsic-guidance targets consumed by the configured guidance_fn recipe.

  • intrinsic_mask (PyTree[Mask | None] | None) – Per-modality mask paired with intrinsic_data.

  • attention_mask (PyTree[Mask | None] | None) – Per-modality attention masks (sequence padding etc.); independent of conditioning.

Returns:

Decoded samples in raw data space.

Raises:
  • TypeError – If a modality declares a generator_type this solver cannot step.

  • ValueError – If stochasticity_scale is set for a modality whose generator carries no score or is negative on [0, 1], or if guidance_scale is not structurally compatible with the modality registry.

Return type:

PyTree[RawVar]

integrate(gen_model, z_init, key, *, context_data=None, context_mask=None, intrinsic_data=None, intrinsic_mask=None, attention_mask=None)

Integrate the SDE/ODE/CTMC and return embedding-space z_final.

Use this (instead of __call__()) when you need the final state in embedding space — for example to compose two solver passes for a round-trip, or to hand the output to another component that operates in embedding space — without paying the cost of trajectory storage from integrate_with_trajectory().

Parameters:
  • gen_model (GenerativeModel) – The generative model (owns the network and the modality registry).

  • z_init (PyTree[Var]) – Initial samples in embedding space, one array per modality.

  • key (Array) – PRNG key for the Brownian increments and CTMC draws; split per step and per modality (deterministic ODE steps ignore theirs).

  • context_data (PyTree[Var] | None) – Context conditioning consumed by the configured guidance_fn recipe, forwarded to the network. See stix.sampling.guidance for the context contract.

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

  • intrinsic_data (PyTree[Var] | None) – Per-modality intrinsic-guidance targets consumed by the configured guidance_fn recipe.

  • intrinsic_mask (PyTree[Mask | None] | None) – Per-modality mask paired with intrinsic_data.

  • attention_mask (PyTree[Mask | None] | None) – Per-modality attention masks (sequence padding etc.); independent of conditioning.

Returns:

Final state in embedding space, PyTree[Var] (no per-modality decode applied).

Raises:
  • TypeError – If a modality declares a generator_type this solver cannot use.

  • ValueError – If stochasticity_scale is set for a modality whose generator carries no score or is negative on [0, 1], or if guidance_scale is not structurally compatible with the modality registry.

Return type:

PyTree[Var]

integrate_with_trajectory(gen_model, z_init, key, *, context_data=None, context_mask=None, intrinsic_data=None, intrinsic_mask=None, attention_mask=None)

Integrate and return the full trajectory in embedding space.

Parameters:
  • gen_model (GenerativeModel) – The generative model (owns the network and the modality registry).

  • z_init (PyTree[Var]) – Initial samples in embedding space, one array per modality.

  • key (Array) – PRNG key for the Brownian increments and CTMC draws; split per step and per modality (deterministic ODE steps ignore theirs).

  • context_data (PyTree[Var] | None) – Context conditioning consumed by the configured guidance_fn recipe, forwarded to the network. See stix.sampling.guidance for the context contract.

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

  • intrinsic_data (PyTree[Var] | None) – Per-modality intrinsic-guidance targets consumed by the configured guidance_fn recipe.

  • intrinsic_mask (PyTree[Mask | None] | None) – Per-modality mask paired with intrinsic_data.

  • attention_mask (PyTree[Mask | None] | None) – Per-modality attention masks (sequence padding etc.); independent of conditioning.

Returns:

the final state as a PyTree[Var], the per-step states as a PyTree[Var] with a leading (num_steps,) axis, and the time values at each step with shape (num_steps,). trajectory holds the state entering each step, so it starts at z_init and stops one step short of z_final. resolve_terminal_mask acts on z_final alone, never on the states stored in trajectory.

Return type:

A triple (z_final, trajectory, trajectory_t)

Raises:
  • TypeError – If a modality declares a generator_type this solver cannot use.

  • ValueError – If stochasticity_scale is set for a modality whose generator carries no score or is negative on [0, 1], or if guidance_scale is not structurally compatible with the modality registry.