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:
objectConfiguration 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 aNoneleaf) is \(\lambda\equiv 0\). For aVelocityAndScoremodality it is the SDE scale, soNonemeans ODE sampling. For aTransitionRatesmodality it is the CTMC corrector scale of \(\bar u_t = (1 + \lambda_t)\,\hat u_t + \lambda_t\,\check u_t\), soNonemeans the plain generative chain (see Introduction for the notations and more details). Setting it for a modality whosegenerator_typecarries no score (Velocity), or setting a callable that is negative on[0, 1], raises aValueError.direction (Direction) – Integration direction (“forward” or “reverse”). Defaults to “forward”. Forward integrates
t=0->t=1, reverset=1->t=0. ForTransitionRatesmodalities, 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 at0 + t_min_tolerance). Defaults to 1e-4.t_max_tolerance (float) – Pad below
t=1(integration starts/ends at1 - t_max_tolerance). Defaults to 1e-4.guidance_fn (GuidanceFn | None) – A
GuidanceFnrecipe returning the per-modalityGenerator.None(default) means no guidance: the solver falls back toget_context_conditional_generator(), a single network forward that is unconditional (withcontext_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
argmaxof the terminal denoising posterior. This parameter is a no-op for modalities without aMaskDiscreteInterpolant, so it is safe to leave on. It is also a no-op onREVERSE: a mask interpolant has data at \(t=1\).
- class stix.sampling.solver_manual.ManualSolver(solver_config)¶
Bases:
objectUnconditional 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.vmapover samples. Each step calls the configuredguidance_fnrecipe for the per-modalityGenerator, then applies the update ruleSTEP_FNSholds for that modality’sgenerator_type:ODE path (
stochasticity_scale is None): Euler on the velocity.SDE path: Euler-Maruyama, drift = velocity + score term. Only
VelocityAndScoremodalities may carry astochasticity_scale.CTMC path: Euler on the transition probabilities of
TransitionRatesmodalities, wherestochasticity_scaleweighs the backward rates into the corrector.
This is the only solver that handles continuous SDEs and discrete CTMCs simultaneously —
Solverrejects aTransitionRatesmodality outright — so it is the solver for discrete and mixed-modality sampling.Guidance is opt-in: with the default
guidance_fn=Nonethe recipe is a single plain network forward, so an unconfigured solver performs unconditional sampling (context-conditional if you passcontext_data). Setguidance_fnto a recipe fromstix.sampling.guidancefor 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
Solverfor 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_fnrecipe, forwarded to the network. Seestix.sampling.guidancefor 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_fnrecipe.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_typethis solver cannot step.ValueError – If
stochasticity_scaleis set for a modality whose generator carries no score or is negative on[0, 1], or ifguidance_scaleis not structurally compatible with the modality registry.
- Return type:
- 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 fromintegrate_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_fnrecipe, forwarded to the network. Seestix.sampling.guidancefor 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_fnrecipe.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_typethis solver cannot use.ValueError – If
stochasticity_scaleis set for a modality whose generator carries no score or is negative on[0, 1], or ifguidance_scaleis not structurally compatible with the modality registry.
- Return type:
- 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_fnrecipe, forwarded to the network. Seestix.sampling.guidancefor 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_fnrecipe.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 aPyTree[Var]with a leading(num_steps,)axis, and the time values at each step with shape(num_steps,).trajectoryholds the state entering each step, so it starts atz_initand stops one step short ofz_final.resolve_terminal_maskacts onz_finalalone, never on the states stored intrajectory.- Return type:
A triple
(z_final, trajectory, trajectory_t)- Raises:
TypeError – If a modality declares a
generator_typethis solver cannot use.ValueError – If
stochasticity_scaleis set for a modality whose generator carries no score or is negative on[0, 1], or ifguidance_scaleis not structurally compatible with the modality registry.