Multisample Flow Matching: Straightening Flows with Minibatch Couplings

Aram-Alexandre Pooladian, Heli Ben-Hamu, Carles Domingo-Enrich, Brandon Amos, Yaron Lipman, Ricky T. Q. Chen

Introduction

Deep generative models offer an attractive family of paradigms that can approximate a data distribution and produce high quality samples, with impressive results in recent years (Ramesh et al., 2022; Saharia et al., 2022; Gafni et al., 2022). In particular, these works have made use of simulation-free training methods for diffusion models (Ho et al., 2020; Song et al., 2021b). A number of works have also adopted and generalized these simulation-free methods (Lipman et al., 2023; Albergo & Vanden-Eijnden, 2023; Liu et al., 2022; Neklyudov et al., 2022) for continuous normalizing flows (CNF; Chen et al. (2018)), a family of continuous-time deep generative models that parameterizes a vector field which flows noise samples into data samples.

Recently, Lipman et al. (2023) proposed Flow Matching (FM), a method to train CNFs based on constructing explicit conditional probability paths between the noise distribution (at time t=0t=0) and each data sample (at time t=1t=1). Furthermore, they showed that these conditional probability paths can be taken to be the optimal transport path when the noise distribution is a standard Gaussian, a typical assumption in generative modeling. However, this does not imply that the marginal probability path (marginalized over the data distribution) is anywhere close to the optimal transport path between the noise and data distributions.

Most existing works, including diffusion models and Flow Matching, have only considered conditional sample paths where the endpoints (a noise sample and a data sample) are sampled independently. However, this results in non-zero gradient variances even at convergence, slow training times, and in particular limits the design of probability paths. In turn, it becomes difficult to create paths that are fast to simulate, a desirable property for both likelihood evaluation and sampling.

Contributions: We present a tractable instance of Flow Matching with joint distributions, which we call Multisample Flow Matching. Our proposed method generalizes the construction of probability paths by considering non-independent couplings of kk-sample empirical distributions.

Among other theoretical results, we show that if an appropriate optimal transport (OT) inspired coupling is chosen, then sample paths become straight as the batch size k→∞k\to\infty, leading to more efficient simulation. In practice, we observe both improved sample quality on ImageNet using adaptive ODE solvers and using simple Euler discretizations with a low budget number of function evaluations. Empirically, we find that on ImageNet, we can reduce the required sampling cost by 30% to 60% for achieving a low Fréchet Inception Distance (FID) compared to a baseline Flow Matching model, while introducing only 4% more training time. This improvement in sample efficiency comes at no degradation in performance, e.g. log-likelihood and sample quality.

Within the deep generative modeling paradigm, this allows us to regularize towards the optimal vector field in a completely simulation-free manner (unlike e.g. Finlay et al. (2020b); Liu et al. (2022)), and avoids adversarial formulations (unlike e.g. Makkuva et al. (2020); Albergo & Vanden-Eijnden (2023)). In particular, we are the first work to be able to make use of solutions from optimal solutions on minibatches while preserving the correct marginal distributions, whereas prior works would only fit to the barycentric average (see detailed discussion in Section 5.1). Beyond generative modeling, we also show how our method can be seen as a new way to compute approximately optimal transport maps between arbitrary distributions in settings where the cost function is completely unknown and only minibatch optimal transport solutions are provided.

Preliminaries

To create a deep generative model, Chen et al. (2018) suggested modeling the vector field utu_{t} with a neural network, leading to a deep parametric model of the flow ψt\psi_{t}, referred to as a Continuous Normalizing Flow (CNF). A CNF is often used to transform a density p0p_{0} to a different one, p1p_{1}, via the push-forward equation

where the second equality defines the push-forward (or change of variables) operator ♯\sharp. A vector field utu_{t} is said to generate a probability path ptp_{t} if its flow ψt\psi_{t} satisfies (2).

2 Flow Matching

A simple simulation-free method for training CNFs is the Flow Matching algorithm (Lipman et al., 2023), which regresses onto an (implicitly-defined) target vector field that generates the desired probability density path ptp_{t}. Given two marginal distributions q0(x0)q_{0}(x_{0}) and q1(x1)q_{1}(x_{1}) for which we would like to learn a CNF to transport between, Flow Matching seeks to optimize the simple regression objective,

where vt(x;θ)v_{t}(x;\theta) is the parametric vector field for the CNF, and ut(x)u_{t}(x) is a vector field that generates a probability path ptp_{t} under the two marginal constraints that pt=0=q0p_{t=0}=q_{0} and pt=1=q1p_{t=1}=q_{1}. While Equation 3 is the ideal objective function to optimize, not knowing (pt,ut)(p_{t},u_{t}) makes this computationally intractable.

Lipman et al. (2023) proposed a tractable method of optimizing (3), which first defines conditional probability paths and vector fields, such that when marginalized over q0(x0)q_{0}(x_{0}) and q1(x1)q_{1}(x_{1}), provide both pt(x)p_{t}(x) and ut(x)u_{t}(x). When targeted towards generative modeling, q0(x0)q_{0}(x_{0}) is a simple noise distribution and easy to directly enforce, leading to a one-sided construction:

where the conditional probability path is chosen such that

Lipman et al. (2023) shows that if ut(x∣x1)u_{t}(x|x_{1}) generates pt(x∣x1)p_{t}(x|x_{1}), then the marginalized ut(x)u_{t}(x) generates pt(x)p_{t}(x), and furthermore, one can train using the much simpler objective of Conditional Flow Matching (CFM):

with xt=ψt(x0∣x1)x_{t}=\psi_{t}(x_{0}|x_{1}); see 2.2.1 for more details. Note that this objective has the same gradient with respect to the model parameters θ\theta as Eq. (3) (Lipman et al., 2023, Theorem 2).

One particular choice of conditional path pt(x∣x1)p_{t}(x|x_{1}) is to use the flow that corresponds to the optimal transport displacement interpolant (McCann, 1997) when q0(x0)q_{0}(x_{0}) is the standard Gaussian, a common convention in generative modeling. The vector field that corresponds to this is

Using this conditional vector field in (1), this gives the conditional flow

Substituting (9) into (8), one can also express the value of this vector field using a simpler expression,

It is evident that this results in conditional flows that (i) tranports all points x0x_{0} from t=0t=0 to x1x_{1} at exactly t=1t=1 and (ii) are straight paths between the samples x0x_{0} and x1x_{1}. This particular case of straight paths was also studied by Liu et al. (2022) and Albergo & Vanden-Eijnden (2023), where the conditional flow (9) is referred to as a stochastic interpolant. Lipman et al. (2023) additionally showed that the conditional construction can be applied to a large class of Gaussian conditional probability paths, namely when pt(x∣x1)=N(x∣μt(x1),σt(x1)2I)p_{t}(x|x_{1})=\mathcal{N}(x|\mu_{t}(x_{1}),\sigma_{t}(x_{1})^{2}I). This family of probability paths encompasses most prior diffusion models where probability paths are induced by simple diffusion processes with linear drift and constant diffusion (e.g. Ho et al. (2020); Song et al. (2021b)). However, existing works mostly consider settings where q0(x0)q_{0}(x_{0}) and q1(x1)q_{1}(x_{1}) are sampled independently when computing training objectives such as (7).

3 Optimal Transport: Static & Dynamic

where Γ(q0,q1)\Gamma(q_{0},q_{1}) is the set of joint measures with left marginal equal to q0q_{0} and right marginal equal to q1q_{1}, called the set of couplings. The minimizer to Equation 11 is called the optimal coupling, which we denote by qc∗q^{*}_{c}. In the case where c(x0,x1)≔∥x0−x1∥2c(x_{0},x_{1})\coloneqq\|x_{0}-x_{1}\|^{2}, the squared-Euclidean distance, Equation 11 amounts to the (squared) 22-Wasserstein distance W22(q0,q1)W_{2}^{2}(q_{0},q_{1}), and we simply write the optimal transport plan as q∗q^{*}.

where utu_{t} generates ptp_{t}, and ptp_{t} satisfies boundary conditions pt=0=q0p_{t=0}=q_{0} and pt=1=q1p_{t=1}=q_{1}. The optimality condition ensures that sample paths xtx_{t} are straight lines, i.e. minimize the length of the path, and leads to paths that are much easier to simulate. Some prior approaches have sought to regularize the model using this optimality objective (e.g. Tong et al. (2020); Finlay et al. (2020b)). In contrast, instead of directly minimizing (12), we will discuss an approach based on using solutions of the optimal coupling q∗q^{*} on minibatch problems, while leaving the marginal constraints intact.

Flow Matching with Joint Distributions

While Conditional Flow Matching in (7) leads to an unbiased gradient estimator for the Flow Matching objective, it was designed with independently sampled x0x_{0} and x1x_{1} in mind. We generalize the framework from Subsection 2.2 to a construction that uses arbitrary joint distributions of q(x0,x1)q(x_{0},x_{1}) which satisfy the correct marginal constraints, i.e.

We will show in Subsection 4 that this can potentially lead to lower gradient variance during training and allow us to design more optimal marginal vector fields ut(x)u_{t}(x) with desirable properties such as improved sample efficiency.

Building on top of Flow Matching, we propose modifying the conditional probability path construction (6) so that at t=0t=0, we define

where q(x0∣x1)q(x_{0}|x_{1}) is the conditional distribution q(x0,x1)q1(x1)\tfrac{q(x_{0},x_{1})}{q_{1}(x_{1})}. Using this construction, we still satisfy the marginal constraint,

i.e. pt=0(x)=∫q(x,x1)dx1=q0(x)p_{t=0}(x)=\int q(x,x_{1})dx_{1}=q_{0}(x) by the assumption made in (13). Then similar to Chen & Lipman (2023), we note that the conditional probability path pt(x∣x1)p_{t}(x|x_{1}) need not be explicitly formulated for training, and that only an appropriate conditional vector field ut(x∣x1)u_{t}(x|x_{1}) needs to be chosen such that all points arrive at x1x_{1} at t=1t=1, which ensures pt=1(x∣x1)=δ(x−x1)p_{t=1}(x|x_{1})=\delta(x-x_{1}). As such, we can make use of the same conditional vector field as prior works, e.g. the choice in Equations 8, 9 and 10.

We then propose the Joint CFM objective as

where xt=ψt(x0∣x1)x_{t}=\psi_{t}(x_{0}|x_{1}) is the conditional flow. Training only involves sampling from q(x0,x1)q(x_{0},x_{1}) and does not require explicitly knowing the densities of q(x0,x1)q(x_{0},x_{1}) or pt(x∣x1)p_{t}(x|x_{1}). Note that Equation (15) reduces to the original CFM objective (7) when q(x0,x1)=q0(x0)q1(x1)q(x_{0},x_{1})=q_{0}(x_{0})q_{1}(x_{1}).

A quick sanity check shows that this objective can be used with any choice of joint distribution q(x0,x1)q(x_{0},x_{1}).

The optimal vector field vt(⋅;θ)v_{t}(\cdot;\theta) in (15), which is the marginal vector field utu_{t}, maps between the marginal distributions q0(x0)q_{0}(x_{0}) and q1(x1)q_{1}(x_{1}).

In the remainder of the section, we highlight some motivations for using joint distributions q(x0,x1)q(x_{0},x_{1}) that are different from the independent distribution q0(x0)q1(x1)q_{0}(x_{0})q_{1}(x_{1}).

Choosing a good joint distribution can be seen as a way to reduce the variance of the gradient estimate, which improves and speeds up training. We develop the gradient covariance at a fixed xx and tt, and bound its total variance:

The total variance (i.e. the trace of the covariance) of the gradient at a fixed xx and tt is bounded as:

Ideally, the flow ψt\psi_{t} of the marginal vector field utu_{t} (and of the learned vθv_{\theta} by extension) should be close to a straight line. The reason is that ODEs with straight trajectories can be solved with high accuracy using fewer steps (i.e. function evaluations), which speeds up sample generation. The quantity

which we call the straightness of the flow and was also studied by Liu (2022), measures how straight the trajectories are. Namely, we can rewrite it as

which shows that S≥0S\geq 0 and only zero if ut(ψt(x0))u_{t}(\psi_{t}(x_{0})) is constant along tt, which is equivalent to ψt(x0)\psi_{t}(x_{0}) being a straight line.

When x0x_{0} and x1x_{1} are sampled independently, the straightness is in general far from zero. This can be seen in the CondOT plots in Figure 2 (right); if flows were close to straight lines, samples generated with one function evaluation (NFE=1) would be of high quality. In Section 4, we show that for certain joint distributions, the straightness of the flow is close to zero.

By Lemma 3.1, the flow ψt\psi_{t} corresponding to the optimal utu_{t} satisfies that ψ0(x0)=x0∼q0\psi_{0}(x_{0})=x_{0}\sim q_{0} and ψ1(x0)∼q1\psi_{1}(x_{0})\sim q_{1}. Hence, x0↦ψ1(x0)x_{0}\mapsto\psi_{1}(x_{0}) is a transport map between q0q_{0} and q1q_{1} with an associated transport cost

Multisample Flow Matching

Constructing a joint distribution satisfying the marginal constraints is difficult, especially since at least one of the marginal distributions is based on empirical data. We thus discuss a method to construct the joint distribution q(x0,x1)q(x_{0},x_{1}) implictly by designing a suitable sampling procedure that leaves the marginal distributions invariant. Note that training with (15) only requires sampling from q(x0,x1)q(x_{0},x_{1}).

We use a multisample construction for q(x0,x1)q(x_{0},x_{1}) in the following manner:

1. Sample {x0(i)}i=1k∼q0(x0)\smash{\{x_{0}^{(i)}\}_{i=1}^{k}\sim q_{0}(x_{0})} and {x1(i)}i=1k∼q1(x1)\smash{\{x_{1}^{(i)}\}_{i=1}^{k}\sim q_{1}(x_{1})}. 2. Construct a doubly-stochastic matrix with probabilities π(i,j)\pi(i,j) dependent on the samples {x0(i)}i=1k\smash{\{x_{0}^{(i)}\}_{i=1}^{k}} and {x1(i)}i=1k\smash{\{x_{1}^{(i)}\}_{i=1}^{k}}. 3. Sample from the discrete distribution, qk(x0,x1)=1k∑i,j=1kδ(x0−x0i)δ(x1−x1j)π(i,j)\smash{q^{k}(x_{0},x_{1})=\frac{1}{k}\sum_{i,j=1}^{k}\delta(x_{0}-x_{0}^{i})\delta(x_{1}-x_{1}^{j})\pi(i,j)}.

Marginalizing qk(x0,x1)q^{k}(x_{0},x_{1}) over samples from Step 1, we obtain the implicitly defined q(x0,x1)q(x_{0},x_{1}). By choosing different couplings π(i,j)\pi(i,j), we induce different joint distributions. In this work, we focus on couplings that induce joint distributions which approximates, or at least partially satisfies, the optimal transport joint distribution. The following result, proven in App. D.3, guarantees that qq has the right marginals.

The joint distribution q(x0,x1)q(x_{0},x_{1}) constructed in Steps has marginals q0(x0)q_{0}(x_{0}) and q1(x1)q_{1}(x_{1}).

That is, the marginal constraints (13) are satisfied and consequently we are allowed to use the framework of Section 3.

The aforementioned multisample construction subsumes the independent joint distribution used by prior works, when the joint coupling is taken to be uniformly distributed, i.e. π(i,j)=1k\pi(i,j)=\frac{1}{k}. This is precisely the coupling used by (Lipman et al., 2023) under our introduced notion of Multisample Flow Matching, and acts as a natural reference point.

2 Batch Optimal Transport (BatchOT) Couplings

The natural connections between optimal transport theory and optimal sampling paths in terms of straight-line interpolations, lead us to the following pseudo-deterministic coupling, which we call Batch Optimal Transport (BatchOT). While it is difficult to solve (11) at the population level, it can efficiently solved on the level of samples. Let {x0(i)}i=1k∼q0(x0)\smash{\{x_{0}^{(i)}\}_{i=1}^{k}}\sim q_{0}(x_{0}) and {x1(i)}i=1k∼q1(x1)\smash{\{x_{1}^{(i)}\}_{i=1}^{k}}\sim q_{1}(x_{1}). When defined on batches of samples, the OT problem (11) can be solved exactly and efficiently using standard solvers, as in POT (Flamary et al., 2021, Python Optimal Transport). On a batch of kk samples, the runtime complexity is well-understood via either the Hungarian algorithm or network simplex algorithm, with an overall complexity of O(k3)\mathcal{O}(k^{3}) (Peyré & Cuturi, 2019, Chapter 3). The resulting coupling πk,∗\pi^{k,*} from the algorithm is a permutation matrix, which is a type of doubly-stochastic matrix that we can incorporate into Step 3 of our procedure.

We consider the effect that the sample size kk has on the marginal vector field ut(x)u_{t}(x). The following theorem shows that in the limit of k→∞k\to\infty, BatchOT satisfies the three criteria that motivate Joint CFM: variance reduction, straight flows, and near-optimal transport cost.

Suppose that Multisample Flow Matching is run with BatchOT. Then, as k→∞k\to\infty,

The value of the Joint CFM objective (Equation (15)) for the optimal utu_{t} converges to 0.

The straightness SS for the optimal marginal vector field utu_{t} (Equation (18)) converges to zero.

As k→∞k\rightarrow\infty, result (i) implies that the gradient variance both during training and at convergence is reduced due to Equation 17; result (ii) implies the optimal model will be easier to simulate between tt=0 and tt=1; result (iii) implies that Multisample Flow Matching can be used as a simulation-free algorithm for approximating optimal transport maps.

The full version of Thm. 4.2 can be found in App. D, and it makes use of standard, weak technical assumptions which are common in the optimal transport literature. While Thm. 4.2 only analyzes asymptotic properties, we provide theoretical evidence that the transport cost decreases with kk, as summarized by a monotonicity result in Thm. D.8.

3 Batch Entropic OT (BatchEOT) Couplings

For kk sufficiently large, the cubic complexity of the BatchOT approach is not always desirable, and instead one may consider approximate methods that produce couplings sufficiently close to BatchOT at a lower computational cost. A popular surrogate, pioneered in (Cuturi, 2013), is to incorporate an entropic penalty parameter on the doubly stochastic matrix, pulling it closer to the independent coupling:

The output of performing Sinkhorn’s algorithm is a doubly-stochastic matrix. The two limiting regimes of the regularization parameter are well understood (c.f. Peyré & Cuturi (2019), Proposition 4.1, for instance): as ε→0\varepsilon\to 0, BatchEOT recovers the BatchOT permutation matrix from Section 4.2; as ε→∞\varepsilon\to\infty, BatchEOT recovers the independent coupling on the indices from Section 4.1.

4 Stable and Heuristic Couplings

An alternative approach is to consider faster algorithms that satisfy at least some desirable properties of an optimal coupling. In particular, an optimal coupling is stable. A permutation coupling is stable if no pair of x0(i)x_{0}^{(i)} and x1(j)x_{1}^{(j)} favor each other over their assigned pairs based on the coupling. Such a problem can be solved using the Gale-Shapeley algorithm (Gale & Shapley, 1962) which has a compute cost of O(k2)\mathcal{O}(k^{2}) given the cross set ranking of all samples. Starting from a random assignment, it is an iterative algorithm that reassigns pairs if they violate the stability property and can terminate very early in practice. Note that in a cost-based ranking, one has to sort the coupling costs of each sample with all samples in the opposing set, resulting in an overall O(k2log⁡(k))\mathcal{O}(k^{2}\log(k)) compute cost.

The Gale-Shapeley algorithm is agnostic to any particular costs, however, as stability is only defined in terms of relative rankings of individual samples. We design a modified version of this algorithm based on a heuristic for satisfying the cyclical monotonicity property of optimal transport, namely that should pairs be reassigned, the reassignment should not increase the total cost of already matched pairs. We refer to the output of this modified algorithm as a heuristic coupling and discuss the details in Appendix A.2.

Related Work

Generative modeling and optimal transport are inherently intertwined topics, both often aiming to learn a transport between two distributions but with very different goals. Optimal transport is widely recognized as a powerful tool for large-scale generative modeling as it can be used to stabilize training (Arjovsky et al., 2017). In the context of continuous-time generative modeling, optimal transport has been used to regularize continuous normalizing flows for easier simulation (Finlay et al., 2020b; Onken et al., 2021), and increase interpretability (Tong et al., 2020). However, the existing methods for encouraging optimality in a generative model generally require either solving a potentially unstable min-max optimization problem (e.g. (Arjovsky et al., 2017; Makkuva et al., 2020; Albergo & Vanden-Eijnden, 2023)) or require simulation of the learned vector field as part of training (e.g. Finlay et al. (2020b); Liu et al. (2022)). In contrast, the approach of using batch optimal couplings can be used to avoid the min-max optimization problem, but has not been successfully applied to generative modeling as they do not satisfy marginal constraints—we discuss this further in the following Section 5.1. On the other hand, neural optimal transport approaches are mainly centered around the quadratic cost (Makkuva et al., 2020; Amos, 2023; Finlay et al., 2020a) or rely heavily on knowing the exact cost function (Fan et al., 2021; Asadulaev et al., 2022). Being capable of using batch optimal couplings allows us to build generative models to approximate optimal maps under any cost function, and even when the cost function is unknown.

Among works that use optimal transport for training generative models are those that make use of batch optimal solutions and their gradients such as Li et al. (2017); Genevay et al. (2018); Fatras et al. (2019); Liu et al. (2019). However, naïvely using solutions to batches only produces, at best, the barycentric map, i.e. the map that fits to average of the batch couplings (Ferradans et al., 2014; Seguy et al., 2017; Pooladian & Niles-Weed, 2021), and does not correctly match the true marginal distribution. This is a well-known problem and while multiple works (e.g. Fatras et al. (2021); Nguyen et al. (2022)) have attempted to circumvent the issue through alternative formulations of optimality, the lack of marginal preservation has been a major downside of using batch couplings for generative modeling as they do not have the ability to match the target distribution for finite batch sizes. This is due to the use of building models within the static setting, where the map is parameterized directly with a neural network. In contrast, we have shown in Lemma 4.1 that in our dynamic setting, where we parameterize the map as the solution of a neural ODE, it is possible to preserve the marginal distribution exactly. Furthermore, we have shown in Proposition D.7 (App. D.5) that our method produces a map that is no higher cost than the joint distribution induced from BatchOT couplings.

Concurrently, Tong et al. (2023) motivates the use of BatchOT solutions within a similar framework as our Joint CFM, but from the perspective of obtaining accurate solutions to dynamic optimal transport problems. Similarly, Lee et al. (2023) propose to explicitly learn a joint distribution, parameterized with a neural network, with the aim of minimizing trajectory curvature; this is done using through an auxiliary VAE-style objective function. In contrast, we propose a family of couplings that all satisfy the marginal constraints, all of which are easy to implement and have negligible cost during training. Our construction allow us to focus on (i) fixing consistency issues within simulation-free generative models, and (ii) using Joint CFM to obtain more optimal solutions than the original BatchOT solutions.

Experiments

We empirically investigate Multisample Flow Matching on a suite of experiments. First, we show how different couplings affect the model on a 2D distribution. We then turn to benchmark, high-dimensional datasets, namely ImageNet (Deng et al., 2009). We use the official face-blurred ImageNet data and then downsample to 32×\times32 and 64×\times64 using the open source preprocessing scripts from Chrabaszcz et al. (2017). Finally, we explore the setting of unknown cost functions while only batch couplings are provided. Full details on the experimental setting can be found in Appendix E.2.

Figure 2 shows the proposed Multisample Flow Matching algorithm on fitting to a checkboard pattern distribution in 2D. We show the marginal probability paths induced by different coupling algorithms, as well as low-NFE samples of trained models on these probability paths.

The diffusion and CondOT probability paths do not capture intricate details of the data distribution until it is almost at the end of the trajectory, whereas Multisample Flow Matching approaches provide a gradual transition to the target distribution along the flow. We also see that with a fixed step solver, the BatchOT method is able to produce an accurate target distribution in just one Euler step in this low-dimensional setting, while the other coupling approaches also get pretty close. Finally, it is interesting that both Stable and Heuristic exhibit very similar probability paths to optimal transport despite only satisfying weaker conditions.

2 Image Datasets

We find that Multisample Flow Matching retains the performance of Flow Matching while improving on sample quality, compute cost, and variance. In Table 6 of Section B.1, we report sample quality using the standard Fréchet Inception Distance (FID), negative log-likelihood values using bits per dimension (BPD), and compute cost using number of function evaluations (NFE); these are all standard metrics throughout the literature. Additionally, we report the variance of ut(x∣x0,x1)u_{t}(x|x_{0},x_{1}), estimated using the Joint CFM loss (15) which is an upper bound on the variance. We do not observe any performance degradations while simulation efficiency improves significantly, even with small batch sizes.

Additionally, in Section B.5, we include runtime comparisons between Flow Matching and Multisample Flow Matching. On ImageNet32, we only observe a 0.8% relative increase in runtime compared to Flow Matching, and a 4% increase on ImageNet64.

We observe that with a fixed NFE, models trained using Multisample Flow Matching generally achieve better sample quality. For these experiments, we draw x0∼N(0,Id)x_{0}\sim\mathcal{N}(0,I_{d}) and simulate vt(⋅,θ)v_{t}(\cdot,\theta) up to time t=1t=1 using a fixed step solver with a fixed NFE. Figures 3 show that even on high dimensional data distributions, the sample quality of of multisample methods improves over the naïve CondOT approach as the number of function evaluations drops. We compare to the FID of diffusion baseline methods in Table 2, and provide additional results in Section B.4.

Interestingly, we find that the Stable coupling actually performs on par, and some times better than the BatchOT coupling, despite having a smaller asymptotic compute cost and only satisfying a weaker condition within each batch.

As FID is computed over a full set of samples, it does not show how varying NFE affects individual sample paths. We discuss a notion of consistency next, where we analyze the similarity between low-NFE and high-NFE samples.

In Figure 1 we show samples at different NFEs, where it can be qualitatively seen that BatchOT produces samples that are more consistent between high- and low-NFE solutions than CondOT, despite achieving similar FID values.

To evaluate this quantitatively, we define a metric for establishing the consistency of a model with respect to an integration scheme: let x(m)x^{(m)} be the output of a numerical solver initialized at xx using mm function evalutions to reach t=1t=1, and let x(∗)x^{(*)} be a near-exact sample solved using a high-cost solver starting from x0x_{0} as well. We define

where F(⋅)\mathcal{F}(\cdot) outputs the hidden units from a pretrained InceptionNetWe take the same layer as used in standard FID computation., and DD is the number of hidden units. These kinds of perceptual losses have been used before to check the content alignment between two image samples (e.g. Gatys et al. (2015); Johnson et al. (2016)). We find that Multisample Flow Matching has better consistency at all values of NFE, shown in Table 3.

Figure 4 shows the convergence of Multisample Flow Matching with BatchOT coupling compared to Flow Matching with CondOT and diffusion-based methods. We see that by choosing better joint distributions, we obtain faster training. This is in line with our variance estimates reported in Table 6 and supports our hypothesis that gradient variance is reduced by using non-trivial joint distributions.

3 Improved Batch Optimal Couplings

This can be viewed as learning the barycentric projection (Ferradans et al., 2014; Seguy et al., 2017), i.e. ψ∗(x0)=EqOT,ck(x1∣x0)[x1]\psi^{*}(x_{0})=E_{q_{OT,c}^{k}(x_{1}|x_{0})}\left[x_{1}\right], a well-studied quantity but is known to not preserve the marginal distribution (Fatras et al., 2019).

We experiment with 4 different cost functions on three synthetic datasets in dimensions {2,32,64}\{2,32,64\} where both q0q_{0} and q1q_{1} are chosen to be Gaussian mixture models. In Table 4 we report both the transport cost and the KL divergence between q1q_{1} and the distribution induced by the learned map, i.e. [ψ1]♯q0[\psi_{1}]_{\sharp}q_{0}. We observe that while B-ST always results in lower transport costs compared to B-FM, its KL divergence is always very high, meaning that the pushed-forward distribution by the learned static map poorly approximates q1q_{1}. Another interesting observation is that B-FM always reduces transport costs compared to B, providing experimental support to the theory (Theorem D.8).

Figure 6 shows the cost of the learned model as we vary the batch size for computing couplings, where the models are trained sufficiently to achieve the same KL values as reported in Table 4. We see that our approach decreases the cost compared to the BatchOT oracle for any fixed batch size, and furthermore, converges to the OT solution faster than the batchOT oracle. Thus, since Multisample Flow Matching retains the correct marginal distributions, it can be used to better approximate optimal transport solutions than simply relying on a minibatch solution.

Conclusion

We propose Multisample Flow Matching, building on top of recent works on simulation-free training of continuous normalizing flows. While most prior works make use of training algorithms where data and noise samples are sampled independently, Multisample Flow Matching allows the use of more complex joint distribution. This introduces a new approach to designing probability paths. Our framework increases sample efficiency and sample quality when using low-cost solvers. Unlike prior works, our training method does not rely on simulation of the learned vector field during training, and does not introduce any min-max formulations. Finally, we note that our method of fitting to batch optimal couplings is the first to also preserve the marginal distributions, an important property in both generative modeling and solving transport problems.

References

Appendix A Coupling algorithms

Multisample FM makes use of batch coupling algorithms to construct an implicit joint distribution satisfying the marginal constraints. While BatchOT coupling is motivated by approximating the OT map, we consider other lower complexity coupling algorithms which produce coupling that satisfy some desired property of optimal couplings. In Table 5 we summarize the runtime complexities for the different algorithms used in this work. We will now describe in detail the Stable and Heuristic coupling algorithms.

(Wolansky, 2020) surveys discrete optimal transport from a stable coupling perspective proving that stability is a necessary condition for OT couplings. Although stable couplings are not OT, they are cheaper to compute and are therefore an appealing approach to pursue. For completeness we formulate the Gale Shapely Algorithm in our setting in Algorithm 1. The rankings R0,R1R_{0},R_{1} hold the preferences of the samples in {x0(i)}i=1k\{x_{0}^{(i)}\}_{i=1}^{k} and {x1(i)}i=1k\{x_{1}^{(i)}\}_{i=1}^{k} respectively. Where R0(i,j)R_{0}(i,j) is the rank of x1(j)x_{1}^{(j)} in x0(i)x_{0}^{(i)}’s preferences and R1(i,j)R_{1}(i,j) is the rank of x0(j)x_{0}^{(j)} in x1(i)x_{1}^{(i)}’s preferences.

A.2 Heuristic couplings

The stable coupling is agnostic to the cost of pairing samples and only takes into account the ranks. Therefore, reassignments during the Gale Shapely algorithms might increase the total cost although the rankings of assigned samples are improved. We draw inspiration from the cyclic monotonicity of OT couplings (Villani, 2008) and from the marriage with sharing formulation in (Wolansky, 2020) and modify the reassignment condition in the Gale Shapely algorithm (see Algorithm 2). The modified condition encourages ”local” monotonicity between the reassigned pairs only, reassigning a pair only if the potentially newly assigned pairs have a lower cost.

Appendix B Additional tables and figures

B.2 How batch size affects the marginal probability paths on 2D checkerboard data

B.3 FID vs NFE using midpoint discretization scheme

B.4 Comparison of FID vs NFE for baseline methods DDPM and ScoreSDE

B.5 Runtime per iteration is not significantly affected by solving for couplings

B.6 Convergence improves when using larger coupling sizes

Appendix C Generated samples

Appendix D Theorems and proofs

We need only prove that the marginal probability path interpolates between q0q_{0} and q1q_{1}.

Theorems 1 and 2 of Lipman et al. (2023) can then be used to prove that (i) the marginal vector field ut(x)u_{t}(x) transports between p0=q0p_{0}=q_{0} and p1=q1p_{1}=q_{1}, and (ii) the Joint CFM objective has the same gradient in expectation as the Flow Matching objective and is uniquely minimized by vt(x;θ)=ut(x)v_{t}(x;\theta)=u_{t}(x).

D.2 Proof of Lemma 3.2

D.3 Proof of Lemma 4.1

For an arbitrary test function ff, by the construction of qq we write

which proves that the marginal of qq for x0x_{0} is q0q_{0}. The same argument works for the x1x_{1} marginal.

D.4 Proof of Theorem 4.2

The marginal vector field corresponding to sample size kk:

We made the dependency on kk explicit, and we used that ψt(x0∣x1)=tx1+(1−t)x0\psi_{t}(x_{0}|x_{1})=tx_{1}+(1-t)x_{0}. Note that equivalently, we can write utku_{t}^{k} as the solution of a simple variational problem.

The flow ψtk(x0)\psi_{t}^{k}(x_{0}) corresponding to utku_{t}^{k}, i.e. the solution of dxtdt=utk(xt)\frac{dx_{t}}{dt}=u_{t}^{k}(x_{t}) with initial condition x0x_{0}. We made the dependency on kk explicit.

The straightness of the flow ψtk\psi_{t}^{k}:

We will use the following three assumptions, which allow us to potentially extend our result beyond BatchOT:

(A2) q0q_{0} admits a density and the optimal transport map TT between q0q_{0} and q1q_{1} under the quadratic cost is continuous.

(A3) We assume that almost surely w.r.t. the draw of X0\bm{X}_{0} and X1\bm{X}_{1}, qkq^{k} converges weakly to qq as k→∞k\to\infty.

Some comments are in order as to when assumptions (A2), (A3) hold, since they are not directly verifiable. By the Caffarelli regularity theorem (see Villani (2008), Ch. 12, originally in Caffarelli (1992)), a sufficient condition for (A2) to hold is the following:

(A2’) q0q_{0} and q1q_{1} have a common support Ω\Omega which is compact and convex, have α\alpha-Hölder densities, and they satisfy the lower bound q0,q1>γq_{0},q_{1}>\gamma for some γ>0\gamma>0.

Assumption (A3) holds when the matching algorithm is BatchOT, that is, when qkq^{k} is the optimal transport plan between q0kq_{0}^{k} and q1kq_{1}^{k}, as shown by the following proposition, which is proven in App. D.4.3.

Let qkq^{k} be the optimal transport plan between q0kq_{0}^{k} and q1kq_{1}^{k} under the quadratic cost (i.e. the result of Steps under BatchOT). We have that almost surely w.r.t. the draws of X0\bm{X}_{0} and X1\bm{X}_{1}, the sequence (qk)k≥0(q_{k})_{k\geq 0} converges weakly to q∗q^{*}, i.e. assumption (A3) holds.

We split the proof of Theorem 4.2 into two parts: in Subsubsec. D.4.1 we prove that the optimal value of the Joint CFM objective (15) converges to zero as k→∞k\to\infty. In Subsubsec. D.4.2, we prove that the straightness converges to zero and the transport cost converges to the optimal transport cost as k→∞k\to\infty.

D.4.1 Convergence of the optimal value of the CFM objective

Suppose that assumptions (A1), (A2) and (A3) hold. We have that

where utku_{t}^{k} is the marginal vector field as defined in (31).

By Lemma D.3, which holds under (A1) and (A2), we have that ff is bounded and continuous. Assumption (A3) states that almost surely w.r.t. the draws of X0\bm{X}_{0} and X1\bm{X}_{1}, the measure qkq_{k} converges weakly to q∗q^{*}. We apply the definition of weak convergence of measures, which implies that almost surely,

To conclude the proof, we use the variational characterization of utku_{t}^{k} given in (32), which implies that

Let ff be the function defined in equation (38). Suppose that assumptions (A1) and (A2) hold. Then, ff is bounded and continuous.

To show that ut∗u^{*}_{t} is continuous, we use that q0q_{0} is absolutely continuous and that consequently a transport map TT exists. Moreover, we have that x1′=T(x0′)x_{1}^{\prime}=T(x_{0}^{\prime}). Consider the transport map TtT_{t} at time tt, defined as Tt(x)=tT(x)+(1−t)xT_{t}(x)=tT(x)+(1-t)x. Thus, we can write that ut∗(Tt(x0))=T(x0)−x0u_{t}^{*}(T_{t}(x_{0}))=T(x_{0})-x_{0}. The non-crossing paths property implies that TtT_{t} is invertible, which means that an inverse Tt−1T_{t}^{-1} exists. We can write

By assumption (A2), the transport map TT is continuous, and so is TtT_{t}. It is well-known fact that if E,E′E,E^{\prime} are metric spaces, EE is compact, and f:E→E′f:E\rightarrow E^{\prime} a continuous bijective function, then f−1:E′→Ef^{-1}:E^{\prime}\rightarrow E is continuous. Thus, Tt−1T_{t}^{-1} is also continuous. From equation (45), we conclude that ut∗u_{t}^{*} is continuous.

The rest of the proof is straightforward: (x1,x0)↦∥x1−x0−ut∗(tx1+(1−t)x0)∥2(x_{1},x_{0})\mapsto\|x_{1}-x_{0}-u_{t}^{*}(tx_{1}+(1-t)x_{0})\|^{2} is bounded and continuous on the bounded supports of q0q_{0} and q1q_{1} for all t∈t\in, and then ff is also continuous and bounded since it is an average of continuous bounded functions, applying the dominated convergence theorem. ∎

D.4.2 Convergence of the straightness and the transport cost

Suppose that assumptions (A1) and (A3) hold. Then,

We have that lim⁡k→∞Sk=0\lim_{k\to\infty}S^{k}=0, where SkS^{k} is the straightness defined in (33).

We begin with the proof of (i). We introduce some additional notation. We define the quantity S∗S^{*} in analogy with SkS^{k}:

Remark that the second factor in the right-hand side is bounded because ut∗u_{t}^{*} and utku_{t}^{k} are bounded. Using Lemma D.5, we obtain that the first factor in the right-hand side tends to zero as kk grows. Thus,

Since ψtk\psi_{t}^{k} is the flow of utku_{t}^{k} and by Jensen’s inequality, we have that

where the limit holds by . Putting together (51) and (54), we end up with Sk=∣S∗−Sk∣→k→∞0S^{k}=|S^{*}-S^{k}|\xrightarrow[]{k\to\infty}0, which proves (i).

For given instances of X0k\bm{X}_{0}^{k} and X1k\bm{X}_{1}^{k}, we can write

Assumption (A3) implies that almost surely, qkq^{k} converges to qq weakly. For distributions on a bounded domain, weak convergence is equivalent to convergence in the Wasserstein distance (Villani, 2008, Thm. 6.8), and this means that W22(q,qk)→k→∞0W_{2}^{2}(q,q^{k})\xrightarrow[]{k\to\infty}0 almost surely. Almost sure convergence implies convergence in probability, which means that

Note that W22(q,qk)W_{2}^{2}(q,q^{k}) is a bounded random variable because qq and qkq^{k} have bounded support as q0,q1,q0kq_{0},q_{1},q_{0}^{k} and q1kq_{1}^{k} have bounded support. Suppose that W22(q,qk)W_{2}^{2}(q,q^{k}) is bounded by the constant CC. Hence, we can write

We can take ϵ\epsilon arbitrarily small, and for a given ϵ\epsilon we can make the second term in the right-hand side arbitrarily small by taking kk large enough. The final result follows. ∎

D.4.3 Proof of Proposition D.1

We have that almost surely, the empirical distributions q0kq_{0}^{k}, resp. q1kq_{1}^{k}, converge weakly to q0q_{0}, resp. q1q_{1} (Varadarajan, 1958). Hence, we can apply Theorem D.6. Since convergence in distribution of random variables is equivalent to weak convergence of their laws, and the law of an optimal coupling is the optimal transport plan, we conclude that (qk)k≥0(q_{k})_{k\geq 0} converges weakly to q∗q^{*}.

D.5 Bounds on the transport cost and monotone convergence results

The following result shows that for an arbitrary joint distribution q(x0,x1)q(x_{0},x_{1}), we can upper-bound the transport cost associated to the marginal vector field utu_{t} to a quantity that depends only q(x0,x1)q(x_{0},x_{1}).

For an arbitrary joint distribution q(x0,x1)q(x_{0},x_{1}) with marginals q0(x0)q_{0}(x_{0}) and q1(x1)q_{1}(x_{1}), let ψt\psi_{t} be the flow corresponding to the marginal vector field utu_{t}. We have that

We make use of the notation introduced in App. D.4. We will rely on the fact that for all t∈t\in, the random variable tx1+(1−t)x0tx_{1}+(1-t)x_{0}, with (x0,x1)∼q(x_{0},x_{1})\sim q has the same distribution as the random variable ψt(x0)\psi_{t}(x_{0}), with x0∼q0x_{0}\sim q_{0}. This is a direct consequence of Lemma 3.1. Using that ψt\psi_{t} is the flow for utu_{t} and Jensen’s inequality twice, we have that

Note that that the statement and proof of this proposition is equivalent to Theorem 3.5 of (Liu et al., 2022), although the language and notation that we use is different, which is why we though convenient to include it.

For the case of BatchOT, the following theorem shows that the quantity in the upper bound of (61) is monotonically decreasing in kk. The combination of Proposition D.7 and Theorem D.8 provides a weak guarantee that for BatchOT, the transport cost should not get much higher when kk increases.

Suppose that Multisample Flow Matching is run with BatchOT. For clarity, we make the dependency on the sample size kk explicit and letNote that here q(k):=qq^{(k)}:=q is a marginalized distribution and is different from qkq^{k} defined in Step 3. q(k)(x0,x1):=q(x0,x1)q^{(k)}(x_{0},x_{1}):=q(x_{0},x_{1}), and ψtk(x0):=ψt(x0)\psi_{t}^{k}(x_{0}):=\psi_{t}(x_{0}). Then, for any k≥1k\geq 1, we have that

In the first equality, we used that the optimal transport map between the empirical distributions q0kq_{0}^{k} and q1kq_{1}^{k} can be encoded as a permutation, which we denote by σk+1\sigma_{k+1}. In the inequality, we introduced the notation σk−j\sigma_{k}^{-j} to denote the optimal permutation within {x0(i)}i∈[k+1]∖{j}\{x^{(i)}_{0}\}_{i\in[k+1]\setminus\{j\}}. The inequality holds because using the optimality of σk+1\sigma_{k+1}:

Appendix E Experimental & evaluation details

We report the hyper-parameters used in Table 10. We use the architecture from Dhariwal & Nichol (2021) but with much lower attention resolution. We use full 32 bit-precision for training ImageNet-32 and 16-bit mixed precision for training ImageNet-64. All models are trained using the Adam optimizer with the following parameters: β1=0.9\beta_{1}=0.9, β2=0.999\beta_{2}=0.999, weight decay = 0.0, and ϵ=1e−8\epsilon=1e{-8}. All methods we trained using identical architectures, with the same parameters for the the same number of epochs (see Table 10 for details), with the exception of Rectified Flow, which we trained for much longer starting from the fully trained CondOT model. We use either a constant learning rate schedule or a polynomial decay schedule (see Table 10). The polynomial decay learning rate schedule includes a warm-up phase for a specified number of training steps. In the warm-up phase, the learning rate is linearly increased from 1e−81e{-8} to the peak learning rate (specified in Table 10). Once the peak learning rate is achieved, it linearly decays the learning rate down to 1e−81e{-8} until the final training step.

When reporting negative log-likelihood, we dequantize using the standard uniform dequantization (Dinh et al., 2016). We report an importance-weighted estimate using

with xx is in {0,…,255}D\{0,\dots,255\}^{D}. We solve for ptp_{t} at exactly t=1t=1 with an adaptive step size solver dopri5 with atol=rtol=1e-5 using the torchdiffeq (Chen, 2018) library. We used KK=15 for ImageNet32 and KK=10 for ImageNet64.

When computing FID, we use the TensorFlow-GAN library https://github.com/tensorflow/gan.

We run coupling algorithms only within each GPU. We also ran coupling algorithms across all GPUs (using the “Effective Batch Size”) in preliminary experiments, but did not see noticeable gains in sample efficiency while obtaining slightly worse performance and sample quality, so we stuck to the smaller batch sizes for running our coupling algorithms.

For Rectified Flow, we use the finalized FM-CondOT model, generate 50000 noise and sample pairs, then train using the same FM-CondOT algorithm and hyperparameters on these sampled pairs. This is equivalent to their 2-Rectified Flow approach (Liu et al., 2022). For the rectification process, we train for 300 epochs.

E.2 Improved batch optimal couplings

Datasets. We experimented with 3 datasets in dimensions {2,32,64}\{2,32,64\} consisting of 50K50K samples. Both q0q_{0} and q1q_{1} were Gaussian mixtures with number of centers described in Table 11.

Neural Networks Architectures. For B-ST we used stacked blocks of Convex Potential Flows (Huang et al., 2020) as an invertible neural network parametrizing the map, which also allowed us to estimate KL divergence:

For B-FM we used a simple MLP with Swish activation. For each dataset we built architectures with roughly the same number of parameters.

Hyperparameter Search. For each dataset and each cost we swept over learning rates {0.005,0.001,0.0005}\{0.005,0.001,0.0005\} and chose the best setting.