Sampling generators

class stix.core.generator.Generator

Bases: PyTreeNode

Base dataclass describing the infinitesimal generator of the Markov process simulated at sampling time for a single modality.

It carries no behaviour. What the solver does with it (an Euler ODE step, an Euler-Maruyama SDE step, or a continuous-time-Markov-chain jump) is decided entirely by the solver, keyed on the concrete Generator subclass. See stix.sampling.steps for the step/drift/diffusion dispatch tables.

Each interpolant declares its associated generator variant via its generator_type class attribute; modalities expose that as generator_type. The generators returned by get_generator() must match those declared for each modality.

This is a frozen flax.struct pytree node, so a per-modality PyTree[Generator] flows through jax.tree.map / stix.core.modality.ModalityRegistry.map() and diffrax args unchanged, with its array leaves traced.

Subclasses only add array fields; the solver interprets those fields.

replace(**updates)

Returns a new object replacing the specified fields with new values.

class stix.core.generator.Velocity(velocity)

Bases: Generator

A deterministic drift \(b_t\) (an ODE generator).

Parameters:

velocity (Var)

velocity

The per-modality velocity field \(b_t\) at the current state and time.

Type:

stix.typing.core.Var

replace(**updates)

Returns a new object replacing the specified fields with new values.

class stix.core.generator.VelocityAndScore(velocity, score)

Bases: Generator

A drift \(b_t\) together with a score \(s_t\).

Supports both ODE sampling and SDE sampling (the score enters the drift correction and sets the diffusion coefficient).

Parameters:
velocity

The per-modality velocity field \(b_t\).

Type:

stix.typing.core.Var

score

The per-modality score \(s_t = \nabla_{z_t} \log p_t(z_t)\).

Type:

stix.typing.core.Var

replace(**updates)

Returns a new object replacing the specified fields with new values.

class stix.core.generator.TransitionRates(forward_rates, backward_rates)

Bases: Generator

Forward and backward transition rates of a CTMC generator.

forward_rates holds the generative probability velocity \(\hat u_t\) out of the current discrete state \(z_t\). backward_rates holds the time-reversed velocity \(\check u_t\), which walks the same path of marginals with decreasing \(t\) and is what rates_from_mixture_distributions() returns for backward=True. Both are valid transition rates, and both feed the corrector mix \(\bar u_t = (1 + \lambda_t)\,\hat u_t + \lambda_t\,\check u_t\). Each trailing axis has length equal to the number of states. The diagonal entry (the rate to stay, \(y = z\)) is conventionally the negative sum of the off-diagonal rates.

Supports continuous-time Markov chains (CTMCs) sampling.

Parameters:
  • forward_rates (Var)

  • backward_rates (Var)

forward_rates

Generative rates \(\hat u_t\) out of the current state, shape (..., num_states).

Type:

stix.typing.core.Var

backward_rates

Time-reversed rates \(\check u_t\), same shape.

Type:

stix.typing.core.Var

replace(**updates)

Returns a new object replacing the specified fields with new values.