Steps

Per-generator-type stepping logic for the samplers.

Generator is pure data; this module holds the behaviour, as plain functions keyed by the concrete generator type. Three tables cover the two solver backends:

NOTE: every STEP_FNS entry receives an already-evaluated Generator, so it gets exactly one network evaluation per step and cannot query the model at an intermediate stage. That fixes the manual solver to first-order explicit schemes — Euler, Euler-Maruyama, and the CTMC analogues. Multi-stage schemes (Heun, midpoint, higher-order stochastic Runge-Kutta) need the step function to own the network evaluation. The adaptive solver, Solver, is not subject to this: it owns the drift/diffusion closures and diffrax calls them as many times per step as its tableau requires.

A generator type that is a key of STEP_FNS but not of DRIFT_FNS (e.g. TransitionRates) can be simulated only by the manual solver: validate_generator_types() rejects it before the adaptive solver traces anything. Adding a new generator type is one data class plus one entry per table, with no solver edits.

Generator-type tables

stix.sampling.steps.STEP_FNS: dict[type[Generator], Callable[[Time, Var, Generator, float, Array, Callable[[Time], Scalar] | None, Direction], Var]] = {<class 'stix.core.generator.TransitionRates'>: <function ctmc_euler_step>, <class 'stix.core.generator.Velocity'>: <function velocity_euler_step>, <class 'stix.core.generator.VelocityAndScore'>: <function velocity_and_score_euler_maruyama_step>}

Single-stage update z_t -> z_{t+dt} keyed by generator type.

stix.sampling.steps.DRIFT_FNS: dict[type[Generator], Callable[[Time, Generator, Callable[[Time], Scalar] | None, Direction], Var]] = {<class 'stix.core.generator.Velocity'>: <function velocity_drift>, <class 'stix.core.generator.VelocityAndScore'>: <function velocity_and_score_drift>}

Diffrax drift term keyed by generator type (continuous modalities only).

stix.sampling.steps.DIFFUSION_FNS: dict[type[Generator], Callable[[Time, Var, Callable[[Time], Scalar] | None], Var]] = {<class 'stix.core.generator.Velocity'>: <function velocity_diffusion>, <class 'stix.core.generator.VelocityAndScore'>: <function velocity_and_score_diffusion>}

Diffrax diffusion term keyed by generator type (continuous modalities only).

Step updates

stix.sampling.steps.velocity_euler_step(t, z_t_k, generator_k, dt, key_k, stochasticity_scale_k, direction)

Euler ODE step for a Velocity.

One first-order update \(z_{t+\mathrm{d}t} = z_t + b_t(z_t)\,\mathrm{d}t\), evaluating the velocity once, at the left endpoint.

Parameters:
  • t (Time) – Current time (unused).

  • z_t_k (Var) – Current state for one modality.

  • generator_k (Generator) – Per-modality generator (must be a Velocity).

  • dt (float) – Signed time step.

  • key_k (Array) – PRNG key (unused).

  • stochasticity_scale_k (Callable[[Time], Scalar] | None) – Stochasticity scale (unused; must be None for ODE).

  • direction (Direction) – Integration direction (unused; the signed dt already encodes it).

Returns:

Updated state after one Euler step.

Return type:

Var

stix.sampling.steps.velocity_and_score_euler_maruyama_step(t, z_t_k, generator_k, dt, key_k, stochasticity_scale_k, direction)

Euler-Maruyama SDE step for a VelocityAndScore.

\[z_{t+\mathrm{d}t} = z_t + \left(b_t \pm \lambda_t s_t\right)\mathrm{d}t + \sqrt{2\lambda_t}\,\sqrt{|\mathrm{d}t|}\,\xi, \qquad \xi\sim\mathcal{N}(0, I),\]

with + forward and - reverse. Drift and score are evaluated once, at the left endpoint. When stochasticity_scale_k is None the noise term vanishes and this degenerates to the Euler ODE step.

Parameters:
  • t (Time) – Current time.

  • z_t_k (Var) – Current state for one modality.

  • generator_k (Generator) – Per-modality generator (must be a VelocityAndScore).

  • dt (float) – Signed time step.

  • key_k (Array) – PRNG key for the Brownian increment.

  • stochasticity_scale_k (Callable[[Time], Scalar] | None) – Time-dependent stochasticity scale, or None for ODE.

  • direction (Direction) – Integration direction; flips the score term on reverse.

Returns:

Updated state after one step.

Return type:

Var

stix.sampling.steps.ctmc_euler_step(t, z_t_k, generator_k, dt, key_k, stochasticity_scale_k, direction)

Euler CTMC step for a TransitionRates.

z_t_k is the current state as a one-hot vector over the CTMC states, shape (..., num_states). The solver mixes the generator’s forward and backward rates into the corrector of stochasticity scale \(\lambda_t\):

\[\bar u_t = (1 + \lambda_t)\,\hat u_t + \lambda_t\,\check u_t,\]

the one-parameter family that walks the marginals \(p_t\) forward, since \(\hat u_t + \check u_t\) is divergence-free against \(p_t\). stochasticity_scale_k is None is \(\lambda\equiv 0\), the plain generative chain. On REVERSE the two weights are swapped, 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 process. Euler-forward on the transition probabilities, over the unsigned interval \(|\mathrm{d}t|\):

\[p(\cdot \mid z) = \delta(\cdot, z) + \mathrm{d}t \, \bar u_t(\cdot \mid z),\]

where the one-hot z_t_k is the \(\delta(\cdot, z)\) term. The probabilities are clipped to be non-negative and renormalised, then a categorical index is drawn independently at each position and re-encoded as a one-hot vector so the sampling state stays one-hot.

Parameters:
  • t (Time) – Current time, used to evaluate stochasticity_scale_k.

  • z_t_k (Var) – One-hot CTMC state for one modality.

  • generator_k (Generator) – Per-modality generator (must be a TransitionRates).

  • dt (float) – Time step. Only its magnitude is used, the direction coming from direction via the swap of the two weights.

  • key_k (Array) – PRNG key for the categorical draw.

  • stochasticity_scale_k (Callable[[Time], Scalar] | None) – Time-dependent corrector scale \(\lambda_t \geq 0\), or None for \(\lambda\equiv 0\). Clamped at zero.

  • direction (Direction) – Integration direction; reverse swaps the weights of the forward and backward rates.

Returns:

Updated one-hot state after one Euler step.

Return type:

Var

Note

One categorical draw per position per step means at most one jump per position per step, so this is the Euler discretisation of the CTMC.

Drift and diffusion

stix.sampling.steps.velocity_drift(t, generator_k, stochasticity_scale_k, direction)

Diffrax drift for a Velocity generator.

Returns the velocity itself.

Parameters:
  • t (Time) – Current time (unused).

  • generator_k (Generator) – Per-modality generator (must be a Velocity).

  • stochasticity_scale_k (Callable[[Time], Scalar] | None) – Stochasticity scale (unused).

  • direction (Direction) – Integration direction (unused; the signed time interval already encodes it).

Returns:

Drift term equal to the generator velocity.

Return type:

Var

stix.sampling.steps.velocity_diffusion(t, z_t_k, stochasticity_scale_k)

Diffrax diffusion for a Velocity.

Always zero (ODE).

Parameters:
  • t (Time) – Current time (unused).

  • z_t_k (Var) – Current state (used only for shape/dtype).

  • stochasticity_scale_k (Callable[[Time], Scalar] | None) – Stochasticity scale (unused).

Returns:

A zero array matching z_t_k.

Return type:

Var

stix.sampling.steps.velocity_and_score_drift(t, generator_k, stochasticity_scale_k, direction)

Diffrax drift for a VelocityAndScore.

b when stochasticity_scale_k is None (ODE), otherwise the SDE drift \(b_t \pm \lambda_t s_t\) (+ forward, - reverse).

Parameters:
  • t (Time) – Current time.

  • generator_k (Generator) – Per-modality generator (must be a VelocityAndScore).

  • stochasticity_scale_k (Callable[[Time], Scalar] | None) – Time-dependent stochasticity scale, or None for ODE.

  • direction (Direction) – Integration direction; flips the score term on reverse.

Returns:

Drift term for the diffrax term.

Return type:

Var

stix.sampling.steps.velocity_and_score_diffusion(t, z_t_k, stochasticity_scale_k)

Diffrax diffusion for a VelocityAndScore.

Zero when stochasticity_scale_k (\(\lambda_t\)) is None, otherwise \(\sqrt{2\lambda_t}\) broadcast to the state shape.

Parameters:
  • t (Time) – Current time.

  • z_t_k (Var) – Current state (used for shape/dtype).

  • stochasticity_scale_k (Callable[[Time], Scalar] | None) – Time-dependent stochasticity scale, or None for ODE.

Returns:

Diagonal diffusion coefficient matching z_t_k.

Return type:

Var