Sampling generators¶
- class stix.core.generator.Generator¶
Bases:
PyTreeNodeBase 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
Generatorsubclass. Seestix.sampling.stepsfor the step/drift/diffusion dispatch tables.Each interpolant declares its associated generator variant via its
generator_typeclass attribute; modalities expose that asgenerator_type. The generators returned byget_generator()must match those declared for each modality.This is a frozen
flax.structpytree node, so a per-modalityPyTree[Generator]flows throughjax.tree.map/stix.core.modality.ModalityRegistry.map()and diffraxargsunchanged, 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:
GeneratorA 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:
- replace(**updates)¶
Returns a new object replacing the specified fields with new values.
- class stix.core.generator.VelocityAndScore(velocity, score)¶
Bases:
GeneratorA 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).
- velocity¶
The per-modality velocity field \(b_t\).
- Type:
- score¶
The per-modality score \(s_t = \nabla_{z_t} \log p_t(z_t)\).
- Type:
- replace(**updates)¶
Returns a new object replacing the specified fields with new values.
- class stix.core.generator.TransitionRates(forward_rates, backward_rates)¶
Bases:
GeneratorForward and backward transition rates of a CTMC generator.
forward_ratesholds the generative probability velocity \(\hat u_t\) out of the current discrete state \(z_t\).backward_ratesholds the time-reversed velocity \(\check u_t\), which walks the same path of marginals with decreasing \(t\) and is whatrates_from_mixture_distributions()returns forbackward=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.
- forward_rates¶
Generative rates \(\hat u_t\) out of the current state, shape
(..., num_states).- Type:
- backward_rates¶
Time-reversed rates \(\check u_t\), same shape.
- Type:
- replace(**updates)¶
Returns a new object replacing the specified fields with new values.