Interpolants

Base classes

class stix.core.interpolant.Interpolant

Bases: ABC

A 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\).

Parameters:
  • key (Array) – PRNG key used to draw the noise.

  • shape (Shape) – Shape of the noise variable to draw (typically the embedded target shape).

Returns:

A noise sample.

Return type:

NoiseVar

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 always None, so the initial state must be computed from \(\epsilon\) only.

Parameters:
  • z_src (EmbeddedVar | None) – Embedded source \(z_{\mathrm{src}}\), or None for a one-sided interpolant.

  • epsilon (NoiseVar) – The interpolant’s noise variable \(\epsilon\).

Returns:

The initial embedded state \(z_0\).

Return type:

Var

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:

EmbeddedVar

class stix.core.interpolant.OneSidedInterpolant

Bases: Interpolant

A 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: Interpolant

A 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\).

Parameters:

gamma_fn (Callable[[Time], Scalar] | None)

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:

EmbeddedVar

sample_noise(key, shape)

Sample standard Gaussian noise \(\epsilon \sim \mathcal{N}(0, I)\).

Parameters:
  • key (Array) – PRNG key used to draw the noise.

  • shape (Shape) – Shape of the noise variable to draw.

Returns:

A Gaussian noise sample of shape shape.

Return type:

NoiseVar

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:

EmbeddedVar

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:
Returns:

The time derivative of the deterministic interpolation at \(t\).

Return type:

Var

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:

Var

class stix.core.interpolant.ContinuousStochasticInterpolant(gamma_fn)

Bases: ContinuousInterpolant, Interpolant

Continuous-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)\).

Parameters:

gamma_fn (Callable[[Time], Scalar])

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:

Var

class stix.core.interpolant.ContinuousDeterministicInterpolant

Bases: ContinuousInterpolant

A base class for the generic deterministic interpolant.

\[z_t = J_t(z_{\mathrm{src}}, z_{\mathrm{tgt}}).\]
generator_type

alias of Velocity

class stix.core.interpolant.ContinuousOneSidedStochasticInterpolant(gamma_fn)

Bases: ContinuousStochasticInterpolant, OneSidedInterpolant

Continuous 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.

Parameters:

gamma_fn (Callable[[Time], Scalar])

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 target is used.

  • t (Time) – Interpolation time.

Returns:

The deterministic interpolation at time \(t\).

Return type:

EmbeddedVar

Discrete interpolants

class stix.core.interpolant.DiscreteInterpolant(num_categories, num_states, num_components, kappa_fn)

Bases: Interpolant

A 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)\).

Parameters:
  • key (Array) – PRNG key used to draw \(\epsilon = (\epsilon_m, \epsilon_s)\).

  • shape (Shape) – Embedded state shape (..., num_states).

Returns:

Noise of shape (*shape[:-1], 2).

Return type:

NoiseVar

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)\).

Parameters:
  • epsilon_m (Var) – Per-position uniforms \(\epsilon_m\sim \mathrm{Unif}[0, 1)\).

  • t (Time) – Interpolation time.

Returns:

Integer indices of shape epsilon_m.shape.

Return type:

Var

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:

Var

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:

EmbeddedVar

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 when distributions are 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=True builds 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 to False.

Returns:

Transition rates of shape (..., num_states).

Return type:

Var

get_conditional_rates(embedded_pairs, z_t, t, *, backward=False)

Conditional CTMC rates given the source-target pair.

Applies rates_from_mixture_distributions() to mixture_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 to False.

Returns:

Conditional rates of shape (..., num_states).

Return type:

Var

class stix.core.interpolant.MaskDiscreteInterpolant(num_categories, kappa_fn)

Bases: DiscreteInterpolant, OneSidedInterpolant

Discrete 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+1 states (K data categories + 1 mask state), and the mask index is K.

Initial state \(z_0\) is the mask state.

Parameters:
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:

Var

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 K data 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:

Var

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() to mixture_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 K data 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 to False.

Returns:

Transition rates of shape (..., num_states).

Return type:

Var

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:

Var

class stix.core.interpolant.UniformDiscreteInterpolant(num_categories, kappa_fn)

Bases: DiscreteInterpolant, OneSidedInterpolant

Discrete 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 K states (one state per data category).

Initial state \(z_0\) is drawn uniformly from the data categories.

Parameters:
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:

Var

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 K data 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:

Var

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() to mixture_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 K data 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 to False.

Returns:

Transition rates of shape (..., num_states).

Return type:

Var

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 K categories.

Return type:

Var

Linear interpolants

class stix.core.interpolant.LinearInterpolant(gamma_fn, alpha_fn, beta_fn)

Bases: ContinuousInterpolant

A 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_fn is set but the pair has no source.

Return type:

EmbeddedVar

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) ignore z_src and return \(\gamma(0)\,\epsilon\).

Parameters:
  • z_src (EmbeddedVar | None) – Embedded source \(z_{\mathrm{src}}\), or None when the interpolant is one-sided.

  • epsilon (NoiseVar) – The interpolant’s noise variable \(\epsilon\).

Returns:

The initial embedded state \(z_0\).

Raises:

ValueError – If alpha_fn is set but z_src is None.

Return type:

Var

class stix.core.interpolant.LinearDeterministicInterpolant(alpha_fn, beta_fn)

Bases: LinearInterpolant, ContinuousDeterministicInterpolant

Linear two-sided deterministic interpolant.

\[z_t = \alpha_t z_{\mathrm{src}} + \beta_t z_{\mathrm{tgt}}.\]
Parameters:
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:

Var

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:

Var

class stix.core.interpolant.LinearStochasticInterpolant(alpha_fn, beta_fn, gamma_fn)

Bases: LinearInterpolant, ContinuousStochasticInterpolant

Linear 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:

Var

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:

Var

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:

Var

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:

Var

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:

Var

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:

Var

class stix.core.interpolant.OneSidedLinearStochasticInterpolant(gamma_fn, beta_fn)

Bases: LinearStochasticInterpolant, ContinuousOneSidedStochasticInterpolant

Linear 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\).

Parameters:
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:

Var

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:

Var

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:

Var

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:

Var

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:

Var

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:

Var

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:

Var

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:

Var

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:

Var

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:

Var

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:

Var

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:

Var

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:

Var

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:

Var

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:

Var

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:

Var

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:

Var

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:

Var

Standard interpolants

class stix.core.interpolant.FlowMatchingOneSidedInterpolant

Bases: OneSidedLinearStochasticInterpolant

One-sided flow-matching interpolant: \(\beta_t = t\), \(\gamma_t = 1 - t\).

Source is Gaussian noise; target is data.

class stix.core.interpolant.FlowMatchingTwoSidedInterpolant

Bases: LinearDeterministicInterpolant

Two-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: LinearStochasticInterpolant

Two-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: OneSidedLinearStochasticInterpolant

One-sided variance exploding diffusion interpolant.

\(\beta_t = 1\); \(\gamma_t\) interpolates linearly between sigma_max at \(t = 0\) and sigma_min at \(t = 1\).

Parameters:
class stix.core.interpolant.ContinuousBFNOneSidedInterpolant(sigma_1=0.01, gamma_min=1e-06)

Bases: OneSidedLinearStochasticInterpolant

One-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_min so that \(\gamma_t > 0\) at \(t=0\), matching the reference BFN implementation.

Parameters:
class stix.core.interpolant.DiscreteBFNOneSidedInterpolant(num_classes, beta_1=0.1, beta_min=1e-06)

Bases: OneSidedLinearStochasticInterpolant

One-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_classes classes (\(K\)) and final accuracy beta_1:

\[\beta_t = t^2 \beta_1, \qquad \gamma_t = \sqrt{K\,\beta_t}.\]

\(\beta_t\) is clipped to beta_min so that \(\gamma_t > 0\) at \(t=0\), matching the reference BFN implementation.

Parameters: