Adaptive solver¶
The diffrax-based adaptive ODE/SDE solvers.
- class stix.sampling.solver.SolverConfig(stochasticity_scale=None, direction=Direction.FORWARD, rtol=0.001, atol=0.001, max_steps=4096, t_min_tolerance=0.0001, t_max_tolerance=0.0001, guidance_fn=None, guidance_scale=<factory>)¶
Bases:
objectConfiguration for the diffrax SDE/ODE 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}\); a callable found negative on
[0, 1]raises aValueError. IfNone(default), ODE sampling is used, otherwise SDE sampling is used. Setting it for a modality whosegenerator_typecarries no score (Velocity) raises aValueError. ATransitionRatesmodality is rejected by this solver outright, with aTypeError, scale or no scale: useManualSolverfor CTMCs.direction (Direction) – Integration direction (“forward” or “reverse”). Defaults to “forward”. Forward integrates
t=0->t=1, reverset=1->t=0.rtol (float) – Relative tolerance for adaptive step-size control. Defaults to 1e-3.
atol (float) – Absolute tolerance for adaptive step-size control. Defaults to 1e-3.
max_steps (int) – Maximum number of solver steps. Defaults to 4096.
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).
Note
The single
guidance_scalefield is sufficient for every recipe instix.sampling.guidance. A recipe that genuinely needs multiple independent scales should bake them into a custom guidance function.
- class stix.sampling.solver.Solver(solver_config)¶
Bases:
objectUnconditional or guided SDE/ODE solver for sampling from the generative model.
A diffrax-based solver which uses adaptive step-sizes. Each integration step calls the configured
guidance_fnrecipe to obtain the per-modalityGenerator, then steps it via the drift/diffusion tables instix.sampling.steps.Guidance is opt-in: with the default
guidance_fn=Nonethe recipe is a single plain network forward, so an unconfigured solver performs unconditional sampling (conditional, if you passcontext_data). Setguidance_fnto a recipe fromstix.sampling.guidancefor classifier-free or intrinsic guidance; the solver class does not change.Two-path design. This solver supplies only the terms — the drift and diffusion formulas come from
DRIFT_FNSandDIFFUSION_FNS, keyed bygenerator_type, and the integration scheme is diffrax’s:ODE path (
stochasticity_scale is None): drift = velocity only, integrated withTsit5.SDE path: drift = velocity + score term, with Brownian noise from a
VirtualBrownianTree, integrated withShARK. OnlyVelocityAndScoremodalities may carry astochasticity_scale.
The gen_model is split internally: the graph definition is captured in closures (static) and the state flows as a dynamic pytree.
Note
A
TransitionRates(CTMC) modality cannot be integrated by diffrax;integrate()raises aTypeErrornaming the solver. UseManualSolverinstead.- Parameters:
solver_config (SolverConfig)
- __call__(gen_model, z_init, key, *, context_data=None, context_mask=None, intrinsic_data=None, intrinsic_mask=None, attention_mask=None)¶
Solve the ODE/SDE and return decoded samples.
The gen_model is split internally: the graph definition is captured in closures (static) and the state flows through diffrax
args(dynamic, traced). Safe to use underjax.vmap.- 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 motion (required; ignored for ODE).
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 (onlyVelocity/VelocityAndScore).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)¶
Solve the ODE/SDE 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.- 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 motion (required; ignored for ODE).
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 (onlyVelocity/VelocityAndScore).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: