Interpolants¶
Base classes¶
- class stix.core.interpolant.Interpolant¶
Bases:
ABCA base class for a generic interpolant.
The general form of the interpolant is:
\(z_t = I_t(z_{\mathrm{src}}, z_{\mathrm{tgt}},\epsilon)\)
with \(z_{\mathrm{src}}\) and \(z_{\mathrm{tgt}}\) the embedded source and target variables, and \(\epsilon\) the noise variable. \(\epsilon\) is always assumed to be independent of \(z_{\mathrm{src}}\) and \(z_{\mathrm{tgt}}\).
See Introduction for more details.
- generator_type: type[Generator]¶
Sampling-time generator subclass returned for this interpolant family.
- abstractmethod sample_noise(key, shape)¶
Sample the interpolant’s noise variable \(\epsilon\).
- abstractmethod sample_initial_state(z_src, epsilon)¶
Sampling-time initial state \(z_0\).
At training time, \(z_0\) is computed using \(z_0 = I_0(z_{\mathrm{src}}, z_{\mathrm{tgt}}, \epsilon)\). At sampling time, \(z_{\mathrm{tgt}}\) is not available, and \(I_0\) might still depend on \(z_{\mathrm{tgt}}\).
Instead of feeding an arbitrary \(z_{\mathrm{tgt}}\) to \(I_0\), this method is used as the sampling-time proxy to compute \(z_0\) from \(z_{\mathrm{src}}\) and \(\epsilon\) only.
Note: For one-sided interpolants,
z_{\mathrm{src}}is alwaysNone, so the initial state must be computed from \(\epsilon\) only.- Parameters:
z_src (EmbeddedVar | None) – Embedded source \(z_{\mathrm{src}}\), or
Nonefor a one-sided interpolant.epsilon (NoiseVar) – The interpolant’s noise variable \(\epsilon\).
- Returns:
The initial embedded state \(z_0\).
- Return type:
- abstractmethod interpolate(embedded_pairs, t, epsilon)¶
Compute the interpolated state at time \(t\).
\[z_t = I_t(z_{\mathrm{src}}, z_{\mathrm{tgt}}, \epsilon).\]- Parameters:
embedded_pairs (EmbeddedSourceTargetPair) – Embedded source-target pair \((z_{\mathrm{src}}, z_{\mathrm{tgt}})\).
t (Time) – Interpolation time.
epsilon (NoiseVar) – The interpolant’s noise variable \(\epsilon\).
- Returns:
The interpolated state \(z_t\).
- Return type:
- class stix.core.interpolant.OneSidedInterpolant¶
Bases:
InterpolantA base class for one-sided interpolant.
The general form of the one-sided interpolant is:
\[z_t = I_t(z_{\mathrm{tgt}},\epsilon)\]with \(z_{\mathrm{tgt}}\) the embedded target variable, and \(\epsilon\) the noise variable. As for the general interpolant, \(\epsilon\) is always assumed to be independent of \(z_{\mathrm{tgt}}\).
By construction, one-sided interpolants do not use any source variable \(z_{\mathrm{src}}\).
At training time, \(z_t\) is built from \(z_{\mathrm{tgt}}\) and \(\epsilon\).
At sampling time, \(z_{\mathrm{tgt}}\) is not available, so
sample_initial_state()maps \(\epsilon\) to \(z_0\) and ignores \(z_{\mathrm{src}}\).
Continuous interpolants¶
- class stix.core.interpolant.ContinuousInterpolant(gamma_fn=None)¶
Bases:
InterpolantA base class for continuous interpolants, following the Stochastic Interpolants framework of [Albergo et al. 2024](https://www.jmlr.org/papers/v26/23-1605.html).
\[z_t = J_t(z_{\mathrm{src}}, z_{\mathrm{tgt}}) + \gamma_t\,\epsilon, \qquad \epsilon \sim \mathcal{N}(0, I).\]Concrete subclasses implement
interpolant_fn(), which returns the deterministic part of the interpolant \(J_t(z_{\mathrm{src}}, z_{\mathrm{tgt}})\), and are initialized with a noise schedule function \(\gamma_t\).- abstractmethod interpolant_fn(embedded_pairs, t)¶
Compute the deterministic part of the interpolant.
\[J_t(z_{\mathrm{src}}, z_{\mathrm{tgt}}).\]- Parameters:
embedded_pairs (EmbeddedSourceTargetPair) – Embedded source-target pair \((z_{\mathrm{src}}, z_{\mathrm{tgt}})\).
t (Time) – Interpolation time.
- Returns:
The deterministic part of the interpolant at time \(t\).
- Return type:
- sample_noise(key, shape)¶
Sample standard Gaussian noise \(\epsilon \sim \mathcal{N}(0, I)\).
- interpolate(embedded_pairs, t, epsilon)¶
Compute the full interpolation including noise.
\[z_t = J_t(z_{\mathrm{src}}, z_{\mathrm{tgt}}) + \gamma_t\,\epsilon.\]- Parameters:
embedded_pairs (EmbeddedSourceTargetPair) – Embedded source-target pair \((z_{\mathrm{src}}, z_{\mathrm{tgt}})\).
t (Time) – Interpolation time.
epsilon (NoiseVar) – The interpolant’s noise variable \(\epsilon\).
- Returns:
The interpolated state \(z_t\).
- Return type:
- get_conditional_velocity_field(embedded_pairs, t, epsilon)¶
Compute the conditional velocity of the deterministic interpolation.
\[v(t, z_t \mid z_{\mathrm{src}}, z_{\mathrm{tgt}}) = \partial_t J_t(z_{\mathrm{src}}, z_{\mathrm{tgt}}).\]- Parameters:
embedded_pairs (EmbeddedSourceTargetPair) – Embedded source-target pair \((z_{\mathrm{src}}, z_{\mathrm{tgt}})\).
t (Time) – Interpolation time.
epsilon (NoiseVar) – The interpolant’s noise variable. Unused; accepted so the signature matches
get_conditional_velocity().
- Returns:
The time derivative of the deterministic interpolation at \(t\).
- Return type:
- get_conditional_velocity(embedded_pairs, t, epsilon)¶
Compute the conditional velocity of the full (noisy) interpolation.
\[b(t, z_t \mid z_{\mathrm{src}}, z_{\mathrm{tgt}}) = \partial_t J_t(z_{\mathrm{src}}, z_{\mathrm{tgt}}) + \dot{\gamma}(t)\,\epsilon.\]- Parameters:
embedded_pairs (EmbeddedSourceTargetPair) – Embedded source-target pair \((z_{\mathrm{src}}, z_{\mathrm{tgt}})\).
t (Time) – Interpolation time.
epsilon (NoiseVar) – The interpolant’s noise variable.
- Returns:
The conditional velocity at time
t.- Return type:
- class stix.core.interpolant.ContinuousStochasticInterpolant(gamma_fn)¶
Bases:
ContinuousInterpolant,InterpolantContinuous-state stochastic interpolant.
\[z_t = J_t(z_{\mathrm{src}}, z_{\mathrm{tgt}}) + \gamma_t\,\epsilon, \qquad \epsilon \sim \mathcal{N}(0, I),\]with \(\gamma_t\) non-zero for \(t\) in \((0, 1)\).
- generator_type¶
alias of
VelocityAndScore
- get_conditional_score(embedded_pairs, t, epsilon)¶
Compute the conditional score.
\[s(t, z_t \mid z_{\mathrm{src}}, z_{\mathrm{tgt}}) = -\epsilon / \gamma_t.\]- Parameters:
embedded_pairs (EmbeddedSourceTargetPair) – Embedded source-target pair \((z_{\mathrm{src}}, z_{\mathrm{tgt}})\). Unused; accepted for signature parity with the velocity helpers.
t (Time) – Interpolation time.
epsilon (NoiseVar) – The interpolant’s noise variable \(\epsilon\).
- Returns:
The conditional score at time \(t\).
- Return type:
- class stix.core.interpolant.ContinuousDeterministicInterpolant¶
Bases:
ContinuousInterpolantA base class for the generic deterministic interpolant.
\[z_t = J_t(z_{\mathrm{src}}, z_{\mathrm{tgt}}).\]
- class stix.core.interpolant.ContinuousOneSidedStochasticInterpolant(gamma_fn)¶
Bases:
ContinuousStochasticInterpolant,OneSidedInterpolantContinuous one-sided stochastic interpolant.
\[z_t = J_t(z_{\mathrm{tgt}}) + \gamma_t\,\epsilon.\]with \(\gamma_t\) non-zero for \(t\) in \([0, 1)\). Note that \(\gamma_{t=0}\) must be non-zero to ensure that the initial state \(z_0\) is not deterministic.
- interpolant_fn(embedded_pairs, t)¶
Compute the deterministic interpolant using only the embedded target.
- Parameters:
embedded_pairs (EmbeddedSourceTargetPair) – Embedded source-target pair \((z_{\mathrm{src}}, z_{\mathrm{tgt}})\). Only
targetis used.t (Time) – Interpolation time.
- Returns:
The deterministic interpolation at time \(t\).
- Return type:
Discrete interpolants¶
- class stix.core.interpolant.DiscreteInterpolant(num_categories, num_states, num_components, kappa_fn)¶
Bases:
InterpolantA base class for discrete interpolants, following the Discrete Flow Matching framework of [Gat et al. 2024](https://proceedings.neurips.cc/paper_files/paper/2024/hash/f0d629a734b56a642701bba7bc8bb3ed-Abstract-Conference.html5).
These interpolants yield the mixture marginals
\[p_t(z \mid z_{\mathrm{src}}, z_{\mathrm{tgt}}) = \sum_{j=1}^{M} \kappa_t^j w_j(z \mid z_{\mathrm{src}}, z_{\mathrm{tgt}}),\]where \(\kappa_t = (\kappa_t^j)_{j=1}^{M} \in \Delta^{M-1}\) are the mixture weights schedule and \(w_j(z \mid z_{\mathrm{src}}, z_{\mathrm{tgt}})\) are the mixture conditionals.
The corresponding interpolant map is given by
\[I_t(z_{\mathrm{src}}, z_{\mathrm{tgt}}, \epsilon) = s_J, \qquad s_J \overset{\epsilon_s}{\sim} w_J(\,\cdot\mid z_{\mathrm{src}}, z_{\mathrm{tgt}}), \qquad J \sim \mathrm{Cat}(\kappa_t)\ \text{via}\ \epsilon_m,\]with \(\epsilon = (\epsilon_m, \epsilon_s)\) and \(\kappa_t = (\kappa_t^j)_{j=1}^{M}\). Here \(\epsilon_m\) is used to draw the mixture index \(J\) and \(\epsilon_s\) is used to sample the conditional state \(s_J\) from the mixture conditional \(w_J\).
Concrete subclasses implement
mixture_conditionals()(the laws \(w_j\)) and are initialised with a schedule \(\kappa_t\).- Parameters:
- generator_type¶
alias of
TransitionRates
- sample_noise(key, shape)¶
Sample packed \((\epsilon_m, \epsilon_s)\), both i.i.d. \(\mathrm{Unif}[0, 1)\).
- draw_mixing_index(epsilon_m, t)¶
Inverse-CDF draw \(J \in \{0,\dots,M-1\}\) with masses \(\kappa_t\) from uniform noise \(\epsilon_m\sim \mathrm{Unif}[0, 1)\).
- abstractmethod mixture_conditionals(embedded_pairs)¶
Mixture conditionals \(w_j(\cdot \mid z_{\mathrm{src}}, z_{\mathrm{tgt}})\).
- Parameters:
embedded_pairs (EmbeddedSourceTargetPair) – Embedded pair \((z_{\mathrm{src}}, z_{\mathrm{tgt}})\).
- Returns:
Stacked categorical laws of shape
(..., M, num_states).- Return type:
- interpolate(embedded_pairs, t, epsilon)¶
The interpolant \(I_t = s_J\).
Draws \(J\) from \(\kappa_t\) with \(\epsilon_m\), then samples the selected conditional \(w_J\) with \(\epsilon_s\).
- Parameters:
embedded_pairs (EmbeddedSourceTargetPair) – Embedded pair \((z_{\mathrm{src}}, z_{\mathrm{tgt}})\).
t (Time) – Interpolation time.
epsilon (NoiseVar) – Packed \((\epsilon_m, \epsilon_s)\).
- Returns:
The interpolated one-hot state \(z_t\).
- Return type:
- rates_from_mixture_distributions(distributions, z_t, t, *, backward=False)¶
Computes the CTMC rates from stacked mixture distributions.
\[u_t = \sum_{j=1}^{M} a^j_t\, w^j + b_t\,\delta_{z_t},\]with \((a_t, b_t)\) from
_rate_coefficients. The same formula is the conditional rates whendistributionsare the mixture conditionals \(w_j(\cdot\mid z_{\mathrm{src}}, z_{\mathrm{tgt}})\), and the unconditional rates when they are the posteriors \(\hat w^j_t\).backward=False(default) is the generative forward rates (t:0→1).backward=Truebuilds the time-reversed CTMC that runs the same mixture path backwards in \(t\) (t:1→0): substitute \(\dot\kappa \to -\dot\kappa\) in \((a_t, b_t)\). Either branch gives valid transition rates, non-negative away from \(z_t\) and summing to zero.- Parameters:
distributions (Var) – Stacked mixture laws of shape
(..., M, num_states).z_t (EmbeddedVar) – Embedded current state. Shape:
(..., num_states).t (Time) – Interpolation time.
backward (bool) – If
True, return the time-reversed rates. Defaults toFalse.
- Returns:
Transition rates of shape
(..., num_states).- Return type:
- get_conditional_rates(embedded_pairs, z_t, t, *, backward=False)¶
Conditional CTMC rates given the source-target pair.
Applies
rates_from_mixture_distributions()tomixture_conditionals().- Parameters:
embedded_pairs (EmbeddedSourceTargetPair) – Embedded pair \((z_{\mathrm{src}}, z_{\mathrm{tgt}})\).
z_t (EmbeddedVar) – Embedded current state. Shape:
(..., num_states).t (Time) – Interpolation time.
backward (bool) – If
True, return the time-reversed rates. Defaults toFalse.
- Returns:
Conditional rates of shape
(..., num_states).- Return type:
- class stix.core.interpolant.MaskDiscreteInterpolant(num_categories, kappa_fn)¶
Bases:
DiscreteInterpolant,OneSidedInterpolantDiscrete interpolant for masked diffusion.
The interpolant path is given by
\[p_t(z \mid z_{\mathrm{tgt}}) = \kappa_t \delta_{z_{\mathrm{tgt}}} + (1-\kappa_t) \delta_{m},\]where \(\kappa_t\) is the schedule and \(m\) is the mask symbol.
The state space has
K+1states (Kdata categories + 1 mask state), and the mask index isK.Initial state \(z_0\) is the mask state.
- property mask_index: int¶
The dedicated mask symbol’s index,
num_categories.- Returns:
The mask symbol index
num_categories.
- mixture_conditionals(embedded_pairs)¶
Computes the mixture conditional distributions from the embedded pair.
The mixture conditional distributions are given by
\[w_1(\cdot \mid z_{\mathrm{tgt}}) = \delta_{z_{\mathrm{tgt}}} \quad \text{and} \quad w_2(\cdot \mid z_{\mathrm{tgt}}) = \delta_{m}.\]- Parameters:
embedded_pairs (EmbeddedSourceTargetPair) – Embedded pair; only \(z_{\mathrm{tgt}}\) is used.
- Returns:
Stacked laws of shape
(..., 2, num_states).- Return type:
- mixture_distributions_from_target_posterior(target_posterior, z_t)¶
Computes the mixture distributions from the target posterior \(\hat{p}_{\mathrm{tgt}\mid t}(z_{\mathrm{tgt}} \mid z_{t})\).
The mixture distributions are given by
\[\begin{split}\hat{w}_1 = \begin{cases} \hat{p}_{\mathrm{tgt}\mid t}(z_{\mathrm{tgt}} \mid z_{t}) & \text{if } z_{t} = m, \\ \delta_{m} & \text{otherwise}. \end{cases} \quad \text{and} \quad \hat{w}_2 = \delta_{m},\end{split}\]where the posterior probability \(\hat{p}_{\mathrm{tgt}\mid t}(z_{\mathrm{tgt}} \mid z_{t})\) is extended with a zero probability for the mask state.
- Parameters:
target_posterior (Var) – Target denoising posterior \(\hat{p}_{\mathrm{tgt}\mid t}(z_{\mathrm{tgt}} \mid z_{t})\) over the
Kdata categories, at \((t, z_{t})\). Shape:(..., num_states - 1).z_t (EmbeddedVar) – Embedded current state. Shape:
(..., num_states).
- Returns:
Stacked posteriors of shape
(..., 2, num_states).- Return type:
- rates_from_target_posterior(target_posterior, z_t, t, *, backward=False)¶
Computes the unconditional CTMC rates from a learned target posterior.
Applies
rates_from_mixture_distributions()tomixture_distributions_from_target_posterior(). For the two-point schedule \(\kappa_t = (\kappa_t, 1-\kappa_t)\) the forward rates are \(\hat u_t = \dot\kappa_t/(1-\kappa_t)\,(p_{\mathrm{tgt}\mid t} - \delta_z)\).- Parameters:
target_posterior (Var) – Target denoising posterior \(p_{\mathrm{tgt}\mid t}\) over the
Kdata categories. Shape:(..., num_states - 1).z_t (EmbeddedVar) – Embedded current state. Shape:
(..., num_states).t (Time) – Interpolation time.
backward (bool) – If
True, return the time-reversed rates. Defaults toFalse.
- Returns:
Transition rates of shape
(..., num_states).- Return type:
- sample_initial_state(z_src, epsilon)¶
Draw \(z_0\): every position starts on the mask symbol.
- Parameters:
z_src (EmbeddedVar | None) – Unused (one-sided).
epsilon (NoiseVar) – Packed \((\epsilon_m, \epsilon_s)\); supplies shape and dtype.
- Returns:
One-hot initial state with every position at the mask index.
- Return type:
- class stix.core.interpolant.UniformDiscreteInterpolant(num_categories, kappa_fn)¶
Bases:
DiscreteInterpolant,OneSidedInterpolantDiscrete interpolant for uniform diffusion.
The interpolant path is given by
\[p_t(z \mid z_{\mathrm{tgt}}) = \kappa_t \delta_{z_{\mathrm{tgt}}} + (1-\kappa_t)\,\mathrm{Unif}(\{0,\dots,K-1\}),\]where \(\kappa_t\) is the schedule.
The state space has exactly
Kstates (one state per data category).Initial state \(z_0\) is drawn uniformly from the data categories.
- mixture_conditionals(embedded_pairs)¶
Computes the mixture conditional distributions from the embedded pair.
The mixture conditional distributions are given by
\[w_1(\cdot \mid z_{\mathrm{tgt}}) = \delta_{z_{\mathrm{tgt}}} \quad \text{and} \quad w_2(\cdot \mid z_{\mathrm{tgt}}) = \mathrm{Unif}(\{0,\dots,K-1\}).\]- Parameters:
embedded_pairs (EmbeddedSourceTargetPair) – Embedded pair; only \(z_{\mathrm{tgt}}\) is used.
- Returns:
Stacked laws of shape
(..., 2, num_states).- Return type:
- mixture_distributions_from_target_posterior(target_posterior, z_t)¶
Computes the mixture distributions from the target posterior \(\hat{p}_{\mathrm{tgt}\mid t}(z_{\mathrm{tgt}} \mid z_{t})\).
The mixture distributions are given by
\[\hat{w}_1 = \hat{p}_{\mathrm{tgt}\mid t}(z_{\mathrm{tgt}} \mid z_{t}) \quad \text{and} \quad \hat{w}_2 = \mathrm{Unif}(\{0,\dots,K-1\}),\]where the current state \(z_t\) does not enter the closed-form source atom.
- Parameters:
target_posterior (Var) – Target denoising posterior \(\hat{p}_{\mathrm{tgt}\mid t}(z_{\mathrm{tgt}} \mid z_{t})\) over the
Kdata categories, at \((t, z_{t})\). Shape:(..., num_states).z_t (EmbeddedVar) – Embedded current state. Shape:
(..., num_states).
- Returns:
Stacked posteriors of shape
(..., 2, num_states).- Return type:
- rates_from_target_posterior(target_posterior, z_t, t, *, backward=False)¶
Computes the unconditional CTMC rates from a learned target posterior.
Applies
rates_from_mixture_distributions()tomixture_distributions_from_target_posterior(). For the two-point schedule \(\kappa_t = (\kappa_t, 1-\kappa_t)\) the forward rates are \(\hat u_t = \dot\kappa_t/(1-\kappa_t)\,(p_{\mathrm{tgt}\mid t} - \delta_z)\).- Parameters:
target_posterior (Var) – Target denoising posterior \(p_{\mathrm{tgt}\mid t}\) over the
Kdata categories. Shape:(..., num_states).z_t (EmbeddedVar) – Embedded current state. Shape:
(..., num_states).t (Time) – Interpolation time.
backward (bool) – If
True, return the time-reversed rates. Defaults toFalse.
- Returns:
Transition rates of shape
(..., num_states).- Return type:
- sample_initial_state(z_src, epsilon)¶
Draw \(z_0\) by sampling the uniform source atom.
At \(t=0\) the schedule selects the uniform component; categories are drawn with \(\epsilon_s\).
- Parameters:
z_src (EmbeddedVar | None) – Unused (one-sided).
epsilon (NoiseVar) – Packed \((\epsilon_m, \epsilon_s)\).
- Returns:
One-hot initial state over
Kcategories.- Return type:
Linear interpolants¶
- class stix.core.interpolant.LinearInterpolant(gamma_fn, alpha_fn, beta_fn)¶
Bases:
ContinuousInterpolantA base class for the generic linear interpolant.
\[z_t = \alpha_t z_{\mathrm{src}} + \beta_t z_{\mathrm{tgt}} + \gamma_t \epsilon.\]- Parameters:
- interpolant_fn(embedded_pairs, t)¶
Compute the deterministic part of the interpolant.
- Parameters:
embedded_pairs (EmbeddedSourceTargetPair) – Embedded source-target pair \((z_{\mathrm{src}}, z_{\mathrm{tgt}})\).
t (Time) – Interpolation time.
- Returns:
The deterministic interpolation at time
t.- Raises:
ValueError – If
alpha_fnis set but the pair has no source.- Return type:
- sample_initial_state(z_src, epsilon)¶
Return \(\alpha(0)\,z_{\mathrm{src}} + \gamma(0)\,\epsilon\).
Drops any leftover \(\beta(0)\,z_{\mathrm{tgt}}\) term. One-sided interpolants (
alpha_fn is None) ignorez_srcand return \(\gamma(0)\,\epsilon\).- Parameters:
z_src (EmbeddedVar | None) – Embedded source \(z_{\mathrm{src}}\), or
Nonewhen the interpolant is one-sided.epsilon (NoiseVar) – The interpolant’s noise variable \(\epsilon\).
- Returns:
The initial embedded state \(z_0\).
- Raises:
ValueError – If
alpha_fnis set butz_srcisNone.- Return type:
- class stix.core.interpolant.LinearDeterministicInterpolant(alpha_fn, beta_fn)¶
Bases:
LinearInterpolant,ContinuousDeterministicInterpolantLinear two-sided deterministic interpolant.
\[z_t = \alpha_t z_{\mathrm{src}} + \beta_t z_{\mathrm{tgt}}.\]- velocity_from_velocity_field(v_t, z_t, t)¶
Compute the velocity from the velocity field, \(b_t = v_t\).
- Parameters:
v_t (Var) – The velocity field at time \(t\).
z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The velocity \(b_t\).
- Return type:
- velocity_field_from_velocity(b_t, z_t, t)¶
Compute the velocity field from the velocity, \(v_t = b_t\).
- Parameters:
b_t (Var) – The velocity at time \(t\).
z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The velocity field \(v_t\).
- Return type:
- class stix.core.interpolant.LinearStochasticInterpolant(alpha_fn, beta_fn, gamma_fn)¶
Bases:
LinearInterpolant,ContinuousStochasticInterpolantLinear stochastic interpolant.
\[z_t = \alpha_t z_{\mathrm{src}} + \beta_t z_{\mathrm{tgt}} + \gamma_t \epsilon\]with \(\gamma_t\) non-zero for \(t\) in \((0, 1)\).
- Parameters:
- score_from_noise(epsilon, z_t, t)¶
Compute the score from the noise.
\[s_t = -\epsilon / \gamma_t.\]- Parameters:
epsilon (Var) – The noise variable \(\epsilon\).
z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The score \(s_t\).
- Return type:
- noise_from_score(s_t, z_t, t)¶
Compute the noise from the score.
\[\epsilon = -\gamma_t s_t.\]- Parameters:
s_t (Var) – The score at time \(t\).
z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The noise variable \(\epsilon\).
- Return type:
- velocity_from_velocity_field_and_noise(v_t, epsilon, z_t, t)¶
Compute the velocity from the velocity field and noise.
\[b_t = v_t + \dot{\gamma}_t \epsilon.\]- Parameters:
v_t (Var) – The velocity field at time
t.epsilon (Var) – The noise variable \(\epsilon\).
z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The velocity \(b_t\).
- Return type:
- velocity_field_from_velocity_and_noise(b_t, epsilon, z_t, t)¶
Compute the velocity field from the velocity and noise.
\[v_t = b_t - \dot{\gamma}_t \epsilon.\]- Parameters:
b_t (Var) – The velocity at time
t.epsilon (Var) – The noise variable \(\epsilon\).
z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The velocity field \(v_t\).
- Return type:
- velocity_from_velocity_field_and_score(v_t, s_t, z_t, t)¶
Compute the velocity from the velocity field and score.
\[b_t = v_t - \dot{\gamma}_t \gamma_t s_t.\]- Parameters:
v_t (Var) – The velocity field at time
t.s_t (Var) – The score at time
t.z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The velocity \(b_t\).
- Return type:
- velocity_field_from_velocity_and_score(b_t, s_t, z_t, t)¶
Compute the velocity field from the velocity and score.
\[v_t = b_t + \dot{\gamma}_t \gamma_t s_t.\]- Parameters:
b_t (Var) – The velocity at time
t.s_t (Var) – The score at time
t.z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The velocity field \(v_t\).
- Return type:
- class stix.core.interpolant.OneSidedLinearStochasticInterpolant(gamma_fn, beta_fn)¶
Bases:
LinearStochasticInterpolant,ContinuousOneSidedStochasticInterpolantLinear one-sided interpolant.
\[z_t = \beta_t z_{\mathrm{tgt}} + \gamma_t\,\epsilon,\]with \(\gamma_t\) non-zero for \(t\) in \([0, 1)\). Note that \(\gamma_{t=0}\) must be non-zero to ensure that the initial state \(z_0\) is not deterministic.
This interpolant class implements built-in conversion methods to convert between the interpolated variable \(z_t\), the embedded target variable \(z_{\mathrm{tgt}}\), the noise variable \(\epsilon\), the velocity field \(v_t\), the velocity \(b_t\), and the score \(s_t\).
- score_from_target(embedded_target, z_t, t)¶
Compute the score from the embedded target \(z_{\mathrm{tgt}}\).
\[s_t = (\beta_t z_{\mathrm{tgt}} - z_t) / \gamma_t^2\]- Parameters:
embedded_target (Var) – The embedded target \(z_{\mathrm{tgt}}\).
z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The score \(s_t\).
- Return type:
- score_from_velocity_field(v_t, z_t, t)¶
Compute the score from the velocity field.
\[s_t = ((\beta_t / \dot{\beta}_t) * v_t - z_t) / \gamma_t^2\]- Parameters:
v_t (Var) – The velocity field at time
t.z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The score \(s_t\).
- Return type:
- score_from_velocity(b_t, z_t, t)¶
Compute the score from the velocity.
\[s_t = 1 / ( (\dot{\beta}_t / \beta_t)*\gamma_t^2 - \dot{\gamma}_t*\gamma_t) * (b_t - (\dot{\beta}_t / \beta_t) * z_t)\]- Parameters:
b_t (Var) – The velocity at time
t.z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The score \(s_t\).
- Return type:
- velocity_from_noise(epsilon, z_t, t)¶
Compute the velocity from the noise.
\[b_t = \dot{\beta}_t (z_t - \gamma_t \epsilon) / \beta_t + \dot{\gamma}_t \epsilon.\]- Parameters:
epsilon (Var) – The noise variable \(\epsilon\).
z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The velocity \(b_t\).
- Return type:
- velocity_from_score(s_t, z_t, t)¶
Compute the velocity from the score.
\[b_t = (\dot{\beta}_t / \beta_t) * (z_t + \gamma_t^2 * s_t) - \dot{\gamma}_t * \gamma_t * s_t\]- Parameters:
s_t (Var) – The score at time
t.z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The velocity \(b_t\).
- Return type:
- velocity_from_target(embedded_target, z_t, t)¶
Compute the velocity from the embedded target \(z_{\mathrm{tgt}}\).
\[b_t = \dot{\beta}_t z_{\mathrm{tgt}} + (\dot{\gamma}_t / \gamma_t) * (z_t - \beta_t z_{\mathrm{tgt}})\]- Parameters:
embedded_target (Var) – The embedded target \(z_{\mathrm{tgt}}\).
z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The velocity \(b_t\).
- Return type:
- velocity_from_velocity_field(v_t, z_t, t)¶
Compute the velocity from the velocity field.
\[b_t = v_t + (\dot{\gamma}_t / \gamma_t) * (z_t - (\beta_t / \dot{\beta}_t) * v_t)\]- Parameters:
v_t (Var) – The velocity field at time
t.z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The velocity \(b_t\).
- Return type:
- velocity_field_from_noise(epsilon, z_t, t)¶
Compute the velocity field from the noise.
\[v_t = \dot{\beta}_t (z_t - \gamma_t \epsilon) / \beta_t\]- Parameters:
epsilon (Var) – The noise variable \(\epsilon\).
z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The velocity field \(v_t\).
- Return type:
- velocity_field_from_score(s_t, z_t, t)¶
Compute the velocity field from the score.
\[v_t = (\dot{\beta}_t / \beta_t) * (z_t + \gamma_t^2 * s_t)\]- Parameters:
s_t (Var) – The score at time
t.z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The velocity field \(v_t\).
- Return type:
- velocity_field_from_velocity(b_t, z_t, t)¶
Compute the velocity field from the velocity.
\[v_t =(1/(1-(\dot{\gamma}_t/\dot{\beta}_t)*(\beta_t/\gamma_t))*(b_t - (\dot{\gamma}_t / \gamma_t) * z_t)\]- Parameters:
b_t (Var) – The velocity at time
t.z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The velocity field \(v_t\).
- Return type:
- velocity_field_from_target(embedded_target, z_t, t)¶
Compute the velocity field from the embedded target \(z_{\mathrm{tgt}}\).
\[v_t = \dot{\beta}_t z_{\mathrm{tgt}}\]- Parameters:
embedded_target (Var) – The embedded target \(z_{\mathrm{tgt}}\).
z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The velocity field \(v_t\).
- Return type:
- noise_from_target(embedded_target, z_t, t)¶
Compute the noise from the embedded target \(z_{\mathrm{tgt}}\).
\[\epsilon = (z_t - \beta_t z_{\mathrm{tgt}}) / \gamma_t\]- Parameters:
embedded_target (Var) – The embedded target \(z_{\mathrm{tgt}}\).
z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The noise variable \(\epsilon\).
- Return type:
- noise_from_velocity_field(v_t, z_t, t)¶
Compute the noise from the velocity field.
\[\epsilon = (z_t - (\beta_t / \dot{\beta}_t) * v_t) / \gamma_t\]- Parameters:
v_t (Var) – The velocity field at time
t.z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The noise variable \(\epsilon\).
- Return type:
- noise_from_velocity(b_t, z_t, t)¶
Compute the noise from the velocity.
\[\epsilon = -\gamma_t / ( (\dot{\beta}_t / \beta_t)*\gamma_t^2 - \dot{\gamma}_t*\gamma_t) * (b_t - (\dot{\beta}_t / \beta_t) * z_t)\]- Parameters:
b_t (Var) – The velocity at time
t.z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The noise variable \(\epsilon\).
- Return type:
- target_from_noise(epsilon, z_t, t)¶
Compute the embedded target \(z_{\mathrm{tgt}}\) from the noise.
\[z_{\mathrm{tgt}} = (z_t - \gamma_t \epsilon) / \beta_t\]- Parameters:
epsilon (Var) – The noise variable \(\epsilon\).
z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The embedded target \(z_{\mathrm{tgt}}\).
- Return type:
- target_from_score(s_t, z_t, t)¶
Compute the embedded target \(z_{\mathrm{tgt}}\) from the score.
\[z_{\mathrm{tgt}} = (z_t + \gamma_t^2 * s_t) / \beta_t\]- Parameters:
s_t (Var) – The score at time
t.z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The embedded target \(z_{\mathrm{tgt}}\).
- Return type:
- target_from_velocity_field(v_t, z_t, t)¶
Compute the embedded target \(z_{\mathrm{tgt}}\) from the velocity field.
\[z_{\mathrm{tgt}} = v_t / \dot{\beta}_t\]- Parameters:
v_t (Var) – The velocity field at time
t.z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The embedded target \(z_{\mathrm{tgt}}\).
- Return type:
- target_from_velocity(b_t, z_t, t)¶
Compute the embedded target \(z_{\mathrm{tgt}}\) from the velocity.
\[z_{\mathrm{tgt}} = b_t/\dot{\beta}_t + 1/((\dot{\beta}_t / \dot{\gamma}_t)*(\gamma_t / \beta_t) - 1) * (b_t / \dot{\beta}_t - z_t / \beta_t)\]- Parameters:
b_t (Var) – The velocity at time
t.z_t (EmbeddedVar) – The interpolated state.
t (Time) – Interpolation time.
- Returns:
The embedded target \(z_{\mathrm{tgt}}\).
- Return type:
Standard interpolants¶
- class stix.core.interpolant.FlowMatchingOneSidedInterpolant¶
Bases:
OneSidedLinearStochasticInterpolantOne-sided flow-matching interpolant: \(\beta_t = t\), \(\gamma_t = 1 - t\).
Source is Gaussian noise; target is data.
- class stix.core.interpolant.FlowMatchingTwoSidedInterpolant¶
Bases:
LinearDeterministicInterpolantTwo-sided deterministic flow-matching interpolant.
\(\alpha_t = 1 - t\), \(\beta_t = t\). Source and target are both data; no noise term.
- class stix.core.interpolant.StochasticFlowMatchingTwoSidedInterpolant¶
Bases:
LinearStochasticInterpolantTwo-sided stochastic flow-matching interpolant.
\(\alpha_t = 1 - t\), \(\beta_t = t\), \(\gamma_t = \sqrt{2t(1-t)}\). Source and target are both data.
- class stix.core.interpolant.VarianceExplodingDiffusionOneSidedInterpolant(sigma_min=0.002, sigma_max=80.0)¶
Bases:
OneSidedLinearStochasticInterpolantOne-sided variance exploding diffusion interpolant.
\(\beta_t = 1\); \(\gamma_t\) interpolates linearly between
sigma_maxat \(t = 0\) andsigma_minat \(t = 1\).
- class stix.core.interpolant.ContinuousBFNOneSidedInterpolant(sigma_1=0.01, gamma_min=1e-06)¶
Bases:
OneSidedLinearStochasticInterpolantOne-sided Bayesian Flow Network schedule (continuous data).
With final noise level
sigma_1:\[\beta_t = 1 - \sigma_1^{2t}, \qquad \gamma_t = \sqrt{\beta_t (1 - \beta_t)}.\]\(\beta_t\) is clipped to
gamma_minso that \(\gamma_t > 0\) at \(t=0\), matching the reference BFN implementation.
- class stix.core.interpolant.DiscreteBFNOneSidedInterpolant(num_classes, beta_1=0.1, beta_min=1e-06)¶
Bases:
OneSidedLinearStochasticInterpolantOne-sided Bayesian Flow Network schedule (discrete data).
Derived from the BFN discrete noise kernel \(y \sim \mathcal{N}(\beta(t) z_x, \beta(t)\, K\, I)\) (Graves et al. 2023), with
num_classesclasses (\(K\)) and final accuracybeta_1:\[\beta_t = t^2 \beta_1, \qquad \gamma_t = \sqrt{K\,\beta_t}.\]\(\beta_t\) is clipped to
beta_minso that \(\gamma_t > 0\) at \(t=0\), matching the reference BFN implementation.