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:
STEP_FNS— one updatez_t -> z_{t+dt}over a givendt, as used by the manual solver,ManualSolver.DRIFT_FNSandDIFFUSION_FNS— the drift/diffusion split the adaptive (diffrax) solver needs.
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
Nonefor ODE).direction (Direction) – Integration direction (unused; the signed
dtalready encodes it).
- Returns:
Updated state after one Euler step.
- Return type:
- 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. Whenstochasticity_scale_kisNonethe 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
Nonefor ODE.direction (Direction) – Integration direction; flips the score term on reverse.
- Returns:
Updated state after one step.
- Return type:
- 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_kis 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 Noneis \(\lambda\equiv 0\), the plain generative chain. OnREVERSEthe 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_kis 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
directionvia 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
Nonefor \(\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:
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
Velocitygenerator.Returns the velocity itself.
- Parameters:
- Returns:
Drift term equal to the generator velocity.
- Return type:
- stix.sampling.steps.velocity_diffusion(t, z_t_k, stochasticity_scale_k)¶
Diffrax diffusion for a
Velocity.Always zero (ODE).
- stix.sampling.steps.velocity_and_score_drift(t, generator_k, stochasticity_scale_k, direction)¶
Diffrax drift for a
VelocityAndScore.bwhenstochasticity_scale_kisNone(ODE), otherwise the SDE drift \(b_t \pm \lambda_t s_t\) (+forward,-reverse).- Parameters:
- Returns:
Drift term for the diffrax term.
- Return type:
- 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\)) isNone, otherwise \(\sqrt{2\lambda_t}\) broadcast to the state shape.