Introduction¶
stix is a JAX / Flax NNX library for generative models based on the
Stochastic Interpolants (Albergo, Boffi, Vanden-Eijnden,
JMLR 2025) and the
Discrete Flow Matching (Gat et al.,
NeurIPS 2024)
frameworks. Designed to be flexible and explicitly tailored for multi-modal
applications, it enables the implementation of a wide variety of models,
including flow matching, diffusion, Bayesian flow networks, and discrete flow
matching.
This page introduces the core mathematical concepts on which the library is
built, and maps each of them onto the corresponding objects in stix. We
recommend reviewing this foundation before diving into the
tutorials, as it will make navigating the API much
more intuitive.
The stochastic interpolants framework¶
Generative modeling can be described as a transport problem, where one seeks to map samples from a source distribution \(p_{\mathrm{src}}\) to samples of a target distribution \(p_{\mathrm{tgt}}\).
Modern generative models achieve this by evolving samples from the source according to a Markovian dynamical process, whose generators are fitted so that the distribution of the evolved samples matches the target. Stochastic interpolants provide a generic framework for constructing such a process. Here we use a slightly extended version of the original construction, so that the same picture also encompasses discrete flow matching, and interpolants whose endpoints might not match the source and target exactly.
Given a source and a target distribution \(p_{\mathrm{src}}\) and \(p_{\mathrm{tgt}}\), an interpolant is a process \(z_t\) on the time interval \(t \in [0, 1]\), defined as
where \((z_{\mathrm{src}}, z_{\mathrm{tgt}})\) is drawn from a joint distribution \(\pi\) whose marginals match \(p_{\mathrm{src}}\) and \(p_{\mathrm{tgt}}\) respectively, and \(\epsilon\) is a random noise variable independent of \((z_{\mathrm{src}}, z_{\mathrm{tgt}})\).
Typically the interpolant is built so that the path endpoints match the pair variables,
The marginal laws of this process then form a continuous path of probability distributions \((p_t)_{0 \leq t \leq 1}\) with \(p_0 = p_{\mathrm{src}}\) and \(p_1 = p_{\mathrm{tgt}}\).
Having defined an interpolant, one then fits the generator of the considered dynamical process so as to reproduce its path of marginals, and uses this learned generator to evolve samples of the source distribution into samples of the target one.
In stix, interpolants are represented by
Interpolant objects, whose
interpolate() method implements the
map \(I_t\). The generators of the process are encapsulated in subclasses
of a Generator data class. Two types of Markovian
processes are implemented: continuous processes, defined by a stochastic or an
ordinary differential equation (S/ODE), and discrete processes defined by a
continuous-time Markov Chain (CTMC). Each type of Markovian process is
associated with a given Generator subclass, which is a property of the
corresponding Interpolant.
A generative model is completely defined by a
GenerativeModel object. Importantly, this object
is responsible for providing the Generators that are required to evolve
the dynamical process at sampling time. It encapsulates the learnable
components that allow to compute the Generators, as well as everything
that defines the strategy for training these components.
Multi-modal setups and embeddings¶
stix is designed to support multi-modal generative models. We distinguish
two data spaces: a raw data space, which describes data in the form in which
they are observed or stored, and an embedded space in which the interpolation
takes place. These spaces may coincide in practice, but keeping them separate
allows the model to choose a representation that is more convenient for the
generative task, for example, by carrying out the dynamics in a latent space
rather than directly on the observed data. We generally denote variables in the
raw data space by \(x\) and their counterparts in the embedded space by
\(z\). The maps between these spaces are encapsulated in
Embedder objects.
In the embedded space, variables are decomposed into a collection of
modalities of different shapes and nature (discrete or continuous), each one
associated with a specific interpolant. Modalities are described by
Modality objects containing, among other things,
the corresponding Embedder and Interpolant. The collection of
modalities is managed by a single
ModalityRegistry, which contains a generic pytree
of Modality objects. The registry itself is owned by the
GenerativeModel, tying the representation of the data, the interpolation
strategy, and the corresponding generators together into a single object.
One- and two-sided interpolants¶
Interpolants can be partitioned into two categories: the one- and two-sided interpolants.
One-sided interpolants, implemented by
OneSidedInterpolant, do not depend on any
source variable \(z_{\mathrm{src}}\), so their interpolated variables read
The source variable has been forgotten, so the joint \(\pi\) reduces to the target distribution, and in principle \(z_0\) depends only on \(\epsilon\).
Interpolants that explicitly depend on both a source and a target variable are said to be two-sided.
Continuous interpolants¶
We call continuous interpolant an interpolant of the form
These are the generic interpolants considered in (Albergo, Boffi, Vanden-Eijnden,
JMLR 2025). They are
implemented by the ContinuousInterpolant class.
Sampling: ODE, SDE, and the Fokker–Planck equation¶
The marginals of an interpolant of the form above satisfy the Fokker–Planck equation
where \(b_t(z)\) is the marginal velocity of the interpolant, defined as
Define the score of the time-\(t\) marginal \(s_t(z)\) and the velocity field \(v_t(z)\) as
We have \(s_t(z) = - \mathbb{E}\bigl[ \epsilon \mid z_t = z\bigr] / \gamma_t\), so that the velocity can be written
One can show that, for any non-negative stochasticity schedule \(\lambda_t\), the process defined by the stochastic differential equation (SDE)
yields the same Fokker–Planck equation as the previous continuous interpolant. As a result, assuming it starts from the same initial distribution \(p_{0}\), this process reproduces the same marginals \(p_t\) as the continuous interpolant.
To produce samples from the target distribution, one can thus fit \(b_t\) and \(s_t\) and evolve samples from \(p_{0}\) using the above SDE. Taking \(\lambda \equiv 0\) turns the SDE into an ordinary differential equation (ODE) that only involves \(b_t\), so one may choose not to fit \(s_t\) and sample with the ODE alone.
Training: conditional velocity and score¶
The marginal fields \(b_t\) and \(s_t\) are not available in closed form, but their conditional counterparts given a pair \((z_{\mathrm{src}}, z_{\mathrm{tgt}})\) and a noise sample \(\epsilon\) are. In fact, we have
One can rely on these conditional quantities to fit the associated unconditional ones. For instance, to fit the velocity, one can minimize
where the expectation is taken over \((z_{\mathrm{src}}, z_{\mathrm{tgt}})\sim\pi\), \(\epsilon\sim\mathcal{N}(0,I)\) and \(t\sim\mathcal{U}(0,1)\). Expanding and using the property of conditional expectation shows that
which is the desired unconditional objective. This method is also valid for \(s_t\) and \(v_t\). Note that this approach is not unique, and other objectives might exist that allow one to estimate the quantities required for sampling, depending on the exact continuous interpolant considered.
Linear continuous interpolants¶
An important subclass of continuous interpolants are the linear continuous
interpolants (LinearInterpolant) of the form
or, for the one-sided case
(OneSidedLinearStochasticInterpolant),
This class of interpolants encompasses many of the modern generative models,
including Flow Matching
(FlowMatchingOneSidedInterpolant,
FlowMatchingTwoSidedInterpolant,
StochasticFlowMatchingTwoSidedInterpolant),
diffusion models in the EDM framework
(VarianceExplodingDiffusionOneSidedInterpolant),
and Bayesian Flow Networks
(ContinuousBFNOneSidedInterpolant,
DiscreteBFNOneSidedInterpolant). The one-sided
linear path identifies \(z_0\) with \(\epsilon\), which will matter when
choosing a coupling.
Discrete interpolants¶
The interpolants above yield continuous interpolated states. We also implement
a class of interpolants (DiscreteInterpolant)
taking values in a discrete space, following the
Discrete Flow Matching framework of (Gat et al.,
NeurIPS 2024).
These interpolants produce interpolated variables in a discrete state space
\(\Sigma\) with \(K\) elements, with conditional marginal probability
paths of the form
where \(w^j(\cdot \mid z_{\mathrm{src}}, z_{\mathrm{tgt}})\) are time-independent mixture conditional distributions on \(\Sigma\) and \(\kappa_t = (\kappa^j_t)_{j=1}^{M}\) is a schedule on the simplex \(\Delta^{M-1}\), namely \(\kappa^j_t\ge 0\), \(\sum_j\kappa^j_t=1\).
To obtain these marginals, we define the interpolant map as
where \(J\sim\mathrm{Cat}(\kappa_t)\). Here \(\epsilon\) is used as a random seed to first sample \(J\), and then sample \(s_J\) from \(w_J\). This is achieved by splitting \(\epsilon\) into \(\epsilon = (\epsilon_m, \epsilon_s)\) with \(\epsilon_m,\epsilon_s\sim\mathcal{U}[0,1)\), and using \(\epsilon_m\), \(\epsilon_s\) to draw \(J\) and \(s_J\) by inverse CDF sampling.
Conditionally on \(\epsilon\), \(z_t\) is a point mass at \(I_t\). Integrating \(\epsilon\) recovers the marginal paths above.
Sampling: CTMC and probability velocity¶
Akin to the marginals of an SDE, which satisfy the Fokker–Planck equation (1), the marginals \(p_t(z)\) associated with a CTMC satisfy a continuity equation that reads
where \(u_t\) is the probability velocity. It is defined by
for all \(z, y \in \Sigma\), where
are the transition probabilities, and \(\mathrm{div}_{z}\) is the discrete divergence operator that reads
Using the fact that the transition probabilities sum to 1, we have \(\sum_{y\in\Sigma} p_{t+h\mid t}(y\mid z) = 1\) for all \(z \in \Sigma\), so that the continuity equation reduces to
Similar to the continuous case above, assuming the initial distributions \(p_0\) are the same, we can thus build a CTMC process reproducing the desired marginals by learning an approximation to the probability velocity \(u_t\).
Then, at sampling time we can simulate the process by taking stochastic jumps between times \(t\) and \(t+\delta t\) according to the Euler discretisation
i.e.
Note that for the right-hand side to be a proper probability distribution for a small enough \(\delta t\), the probability velocity must satisfy the conditions
The first condition ensures that the probabilities sum to 1, and the second that they are non-negative.
Training: conditional probability velocity and posterior¶
As previously, the probability velocity is not available in closed form, but one can use its conditional counterpart given \((z_{\mathrm{src}}, z_{\mathrm{tgt}})\). Following (Gat et al., NeurIPS 2024), for the previous interpolant one can write the conditional probability velocity as
with \(a^j_t = \dot\kappa^j_t - \kappa^j_t\,\dot\kappa^\ell_t / \kappa^\ell_t\), \(b_t = \dot\kappa^\ell_t / \kappa^\ell_t\), and \(\ell = \arg\min_j \dot\kappa^j_t / \kappa^j_t\).
In fact, injecting this conditional probability velocity into the continuity equation (2) yields the correct time derivative for the conditional marginals \(p_t(z\mid z_{\mathrm{src}}, z_{\mathrm{tgt}})\). Note that the terms involving the \(\ell\) index in the coefficients \(a_t\) and \(b_t\) are chosen so as to ensure the constraints in (3).
One can show that
For the considered interpolants, this yields
where
are the posteriors of the \(w^j\) given \(z_t = y\).
As in the continuous case, one can thus rely on the conditional velocity to learn \(u_t\). In practice, since \(\Sigma\) is finite, one can also directly fit the posterior distribution \(p_t(z_{\mathrm{src}}, z_{\mathrm{tgt}}\mid z)\) and build the \(\hat w^j_t\) to derive an estimator of \(u_t(z \mid y)\).
Time reversal and corrector sampling¶
The continuity equation (2) constrains the divergence \(\mathrm{div}_{z}(p_t u_t)\), not the probability velocity itself. Just as the continuous case admits a whole family of processes sharing the same marginals (indexed by the stochasticity schedule \(\lambda_t\)), for any \(\eta_t\) such that \(\mathrm{div}_{z}(p_t\eta_t) = 0\), \(u_t + \eta_t\) also reproduces the marginals \(p_t\).
A canonical example is given by the time-reversed velocity \(\check u_t\), which walks the same path of marginals with decreasing \(t\). For the mixture paths above it is obtained by substituting \(\dot\kappa_t \to -\dot\kappa_t\) in the coefficients of the previous section,
where \((\tilde a_t, \tilde b_t)\) are the coefficients \((a_t, b_t)\) evaluated with \(\dot\kappa_t\) replaced by \(-\dot\kappa_t\). The substitution generally moves the index \(\ell = \arg\min_j \dot\kappa^j_t / \kappa^j_t\), which is precisely what keeps the constraints (3) satisfied in the reversed direction. By construction we have
so \(\hat u_t + \check u_t\) is divergence-free against \(p_t\), and the whole family
reproduces the marginals \(p_t\). Here \(\lambda_t\) plays exactly the same role it plays for the SDE: \(\lambda_t = 0\) is the plain generative chain, while \(\lambda_t > 0\) adds a corrector that undoes part of the transport and lets the forward term decide again. For a masked diffusion path, for instance, \(\check u_t\) sends already-revealed states back onto the mask symbol, so \(\lambda_t > 0\) is the familiar re-masking corrector. Swapping the two weights, i.e. \(\bar u_t = \lambda_t\,\hat u_t + (1 + \lambda_t)\,\check u_t\), runs the chain backwards instead, transporting the target back onto the source.
In stix, the reversed velocity is produced by the backward=True
argument of
rates_from_mixture_distributions(),
and of the rates_from_target_posterior methods of the standard discrete
interpolants. Either direction comes out as valid transition rates, and both
are carried by the TransitionRates generator: its
forward_rates hold \(\hat u_t\) and its backward_rates hold
\(\check u_t\), so a solver assembles (4) as written.
The weight is the stochasticity_scale field of
ManualSolverConfig — the very field that
carries \(\lambda_t\) for a continuous SDE — which defaults to
\(\lambda_t \equiv 0\).
Finally, both the continuous and the discrete processes can be integrated in
either direction, as selected by the Direction
field of the solver configurations. On
REVERSE, an S/ODE is integrated from
\(t=1\) down to \(t=0\) with the sign of its score term flipped, and a
CTMC swaps the two weights of (4), so that the default
\(\lambda_t \equiv 0\) becomes the pure reversed chain
\(\bar u_t = \check u_t\).
Approximate endpoints and sampling time¶
For two-sided interpolants, the identification \(z_0 = z_{\mathrm{src}}\), \(z_1 = z_{\mathrm{tgt}}\) is the intended design, but in a generic implementation it may hold only approximately. For instance, one can consider a linear two-sided interpolant of the form
with \(\gamma_0\) and \(\gamma_1\) very small but non-zero, leaving a small residual noise at both endpoints. In order to avoid confusion, we keep distinct notations in the library, using \(z_{\mathrm{src}}\) and \(z_{\mathrm{tgt}}\) for the source and target variables, and \(z_0\) and \(z_1\) for the initial and final variables.
One-sided interpolants use no source distribution, so their initial state
\(z_0\) should typically be determined by the noise \(\epsilon\). Again,
this might not always be true in practice. As an example, the one-sided
interpolant corresponding to a variance exploding (VE) diffusion model
(VarianceExplodingDiffusionOneSidedInterpolant)
reads
so that the initial state \(z_{0}\) depends on \(z_{\mathrm{tgt}}\), although it should be largely dominated by the noise term \(\sigma_{\mathrm{max}}\epsilon\).
This kind of residual dependency on the target variable is problematic at sampling time. Because \(z_{\mathrm{tgt}}\) is unavailable, evaluating \(z_{0}\) using \(I_{t=0}\) would require passing an arbitrary value for \(z_{\mathrm{tgt}}\).
To avoid this arbitrary choice, every
Interpolant implements a
sample_initial_state() method that
maps the quantities available then — \(z_{\mathrm{src}}\) (or none, for a
one-sided path) and \(\epsilon\) — onto the sampling-time law of
\(z_0\), without using \(z_{\mathrm{tgt}}\).
Source and target distributions, coupling¶
Since the target (and sometimes the source) distribution is not available in closed form, training relies on sampling \((x_{\mathrm{src}}, x_{\mathrm{tgt}})\) pairs from a dataset and embedding them as \((z_{\mathrm{src}}, z_{\mathrm{tgt}})\).
In the interpolant framework one is free to choose the joint distribution of the pair, \(\pi\), subject only to the correct marginals. Different couplings change the geometry of the learned trajectories and can substantially reduce the number of solver steps needed at sampling time, without changing \(p_{\mathrm{src}}\) or \(p_{\mathrm{tgt}}\).
In practice, the sampling from a joint distribution \(\pi\) is simulated by coupling the source and target variables to introduce specific correlations. This can be done at the level of the dataset itself, in which case the pairs drawn from the dataset are already correlated. We call this approach an offline coupling.
Another approach is to introduce the correlations on the fly at the level of
each batch, a method we call online coupling. In stix, online coupling
can be implemented using a Coupling object. At
training time, the Coupling receives a batch of raw and embedded pairs and
returns correlated re-coupled pairs.
Couplings act on source and target only. In particular, they cannot be used with one-sided interpolants, for which no source variable is provided. Following the discussion of the previous paragraph, one could consider using couplings acting on the target variable \(z_\mathrm{tgt}\) and the noise \(\epsilon\). However, this would break the central assumption that \(\epsilon\) is independent of \(z_\mathrm{tgt}\). This assumption is required to derive the formula
that relates the continuous conditional score and the noise. Correlating \(\epsilon\) and \(z_\mathrm{tgt}\) would thus break SDE sampling (although ODE sampling would be fine).
Instead of coupling the target variable with the noise, one must thus turn the one-sided interpolant into an equivalent two-sided one. For instance, a linear one-sided interpolant of the form
can be replaced by a two-sided one of the form
with \(\gamma_t^2 = \alpha_t^2+\sigma_t^2\) and \(z_{\mathrm{src}}\sim\mathcal{N}(0,1)\). A specific choice of the \(\alpha_t\) and \(\sigma_t\) schedules then corresponds to a specific decomposition of the initial variable \(z_0\) into a part \(\alpha_0 z_{\mathrm{src}}\) that can be correlated to \(z_{\mathrm{tgt}}\) and an independent part \(\sigma_0 \epsilon\) that is used to build the score \(s_t = -\epsilon/\sigma_t\).
Implementing this requires being able to specify the marginal distribution of
\(z_{\mathrm{src}}\) directly in the embedded space. When the embedding is
non-trivial, defining this distribution as the image of a distribution of the
raw source \(x_{\mathrm{src}}\) can be difficult. To avoid this, we allow
generating source variables in embedded space on the fly. This is done by
attaching an embedded_source_prior method to the
Modality object. After embedding of the available
raw source variables \(x_{\mathrm{src}}\), missing embedded source
\(z_{\mathrm{src}}\) are generated using the provided
embedded_source_prior. At training time, this occurs before any coupling,
so the generated source can be used for online coupling.