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

Configuration 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 a ValueError. If None (default), ODE sampling is used, otherwise SDE sampling is used. Setting it for a modality whose generator_type carries no score (Velocity) raises a ValueError. A TransitionRates modality is rejected by this solver outright, with a TypeError, scale or no scale: use ManualSolver for CTMCs.

  • direction (Direction) – Integration direction (“forward” or “reverse”). Defaults to “forward”. Forward integrates t=0 -> t=1, reverse t=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 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).

Note

The single guidance_scale field is sufficient for every recipe in stix.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: object

Unconditional 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_fn recipe to obtain the per-modality Generator, then steps it via the drift/diffusion tables in stix.sampling.steps.

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 (conditional, if you pass context_data). Set guidance_fn to a recipe from stix.sampling.guidance for 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_FNS and DIFFUSION_FNS, keyed by generator_type, and the integration scheme is diffrax’s:

  • ODE path (stochasticity_scale is None): drift = velocity only, integrated with Tsit5.

  • SDE path: drift = velocity + score term, with Brownian noise from a VirtualBrownianTree, integrated with ShARK. Only VelocityAndScore modalities may carry a stochasticity_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 a TypeError naming the solver. Use ManualSolver instead.

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 under jax.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_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 (only Velocity / VelocityAndScore).

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

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_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 (only Velocity / VelocityAndScore).

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