ConceptioArchivearXiv CS
arXiv CSopen access

Fixed-Point Neural Optimal Transport without Implicit Differentiation

2026 · arxiv_cs
arXiv CS · Papers · License: Open Access · 2026
Open Source ↗Direct PDF ↓
neural-networks
machine learning, deep learning, neural networks

Fixed-Point Neural Optimal Transport without Implicit Differentiation∗

arXiv:2605.10792v1 [math.OC] 11 May 2026

Yesom Park† , Eric Gelphman ‡ , Stanley Osher† , and Samy Wu Fung‡

Abstract. We propose an implicit neural formulation of optimal transport that eliminates adversarial min– max optimization and multi-network architectures commonly used in existing approaches. Our key idea is to parameterize a single potential in the Kantorovich dual and reformulate the associated c-transform as a proximal fixed-point problem. This yields a stable single-network framework in which dual feasibility is enforced exactly through proximal optimality conditions rather than adversarial training. Despite the inner fixed-point computation, gradients can be computed without differentiating through the fixed-point iterations, enabling efficient training without requiring implicit differentiation. We further establish convergence of stochastic gradient descent. The resulting framework is efficient, scalable, and broadly applicable: it simultaneously recovers forward and backward transport maps and naturally extends to class-conditional settings. Experiments on high-dimensional Gaussian benchmarks, physical datasets, and image translation tasks demonstrate strong transport accuracy together with improved training stability and favorable computational and memory efficiency. Key words. Optimal Transport, Fixed-point Iteration, Deep Learning, Kantorovich Dual, Convergence MSC codes. 49Q22, 68T07, 65K10

1. Introduction. Optimal transport (OT) is a fundamental problem of finding a mapping between probability distributions that minimizes a prescribed transportation cost. Owing to its solid theoretical foundation and wide applicability, OT has been successfully employed in diverse fields such as traffic control [13, 20, 9], biomedical data analysis [62, 45, 12], generative modeling [69, 53, 74, 48], and domain adaptation [16, 15, 18, 8]. Recent advances in deep learning have led to a surge of interest in scalable OT solvers based on neural parameterizations. Early approaches are rooted in the Monge formulation [51, 72] and its relaxation to the Kantorovich framework [52]. Despite their theoretical rigor, these formulations often entail significant computational complexity, particularly in high-dimensional settings. A predominant line of work is based on the Kantorovich dual formulation, which recasts the OT problem as a saddle-point optimization over transport maps and dual potentials [48, 65, 42, 49, 14]. While this perspective enables scalable algorithmic implementations, it typically necessitates adversarial training of multiple neural networks, thereby introducing optimization instability, sensitivity to hyperparameter selection, and convergence difficulties. These challenges are further exacerbated in WGAN-based methods [5, 48], particularly as the dimensionality of the problem increases. To leverage additional structural properties, several studies focus on the quadratic-cost setting, where Brenier’s theorem [10] guarantees that the optimal transport map can be ex∗

Yesom Park and Eric Gelphman contributed to this work equally. Funding: This work was supported by DARPA HR00112590074, DoE DE-SC0026262, NSF 2208272, ARO W911NF241015, and NSF DMS 2309810. † Department of Mathematics, University of California, Los Angeles ([email protected]), [email protected]. ‡ Department of Applied Mathematics and Statistics, Colorado School of Mines (eric [email protected]), [email protected]. 1

2

Y. PARK, E. GELPHMAN, S. OSHER AND S. WU FUNG

pressed as the gradient of a convex potential. This observation has motivated the use of input convex neural networks (ICNNs) [4] and related architectures [65, 42]. In addition, weak formulations have been explored to directly parameterize transport maps [7, 44, 6]. Nevertheless, many of these approaches continue to rely on auxiliary networks, alternating optimization procedures, or adversarial training, and thus inherit the limitations associated with min–max optimization. An alternative line of research formulates OT as a dynamical system via continuous flows [73, 67, 53, 39]. These methods typically require solving ordinary differential equations (ODEs) or stochastic differential equations (SDEs), which imposes considerable computational overhead during both training and inference. Regularized variants, including entropic and f divergence-based formulations [32, 63, 19, 34], can improve numerical stability but generally introduce bias, leading to transport maps that deviate from the true OT solution. Various efforts have been made to mitigate these challenges through improved optimization strategies, such as natural gradient methods [64, 50], as well as regularization techniques including L2 penalties [32, 61] and cycle-consistency constraints [32, 61, 41, 43]. However, these approaches do not fundamentally resolve the reliance on complex optimization schemes or multiple model components [42, 26]. More recently, formulations based on Hamilton–Jacobi–Bellman (HJB) equations [57] have been proposed to eliminate adversarial training by casting OT as a single-objective optimization problem [56]. In particular, characteristicbased representations enable the recovery of both forward and backward transport maps within a unified framework. However, such methods require solving the HJB equation over the entire computational domain, which significantly limits their scalability in high-dimensional settings. In this work, we propose a method that recovers both forward and backward optimal transport maps through a single minimization problem defined over one neural network. Specifically, we show that quadratic optimal transport admits a significantly simpler and more direct formulation than those employed in existing neural approaches. Starting from the Kantorovich dual problem, we parameterize the dual potential as the negative of a convex function represented by a neural network gθ . Under this parameterization, the associated c-transform reduces to a strongly convex proximal optimization problem, whose minimizer can be computed efficiently via standard fixed-point iterations. This perspective leads to a single-objective optimization problem in which dual feasibility is enforced exactly through the optimality conditions of the proximal operator, thereby eliminating the need for saddle-point formulations, auxiliary conjugate networks, and alternating optimization procedures. As a result, the entire model is defined by a single convex potential learned through a standard minimization framework. A key observation is that the gradient of the resulting objective does not require differentiation through the inner minimization. By exploiting the first-order optimality condition, the dependence of the objective on the proximal minimizer simplifies, yielding a closed-form expression for the gradient. Consequently, training can be carried out using standard backpropagation applied solely to the potential network, without implicit differentiation or unrolled optimization. These properties lead to a method that is both computationally efficient and stable in practice: removing adversarial training mitigates instability associated with min–max optimization, while avoiding differential equation solvers eliminates the overhead inherent in dynamical formulations. Furthermore, fixed-point iterations enable fast and scalable compu-

FIXED-POINT NEURAL OPTIMAL TRANSPORT

3

tation across dimensions. The proposed framework relies on a single network architecture to recover both forward and backward optimal transport maps, reducing model complexity and improving computational efficiency. In addition, the formulation naturally extends to class-conditional optimal transport problems. Empirically, we demonstrate that the proposed method accurately recovers transport maps on high-dimensional Gaussian benchmarks and remains effective on physics-based datasets, capturing complex and realistic distributions. The method also performs competitively on class-conditional transport tasks, including experiments on image data, while exhibiting stable training dynamics and favorable computational and memory efficiency as the dimensionality increases. Contributions. Our main contributions are summarized as follows: • We derive a single-network reformulation of quadratic optimal transport from the Kantorovich dual, in which evaluating the objective reduces to solving a fixed-point problem. • We show that this objective can be differentiated without implicit differentiation, due to an exact cancellation induced by the first-order optimality conditions of the proximal operator. • We establish convergence of the resulting training procedure to a stationary point under inexact fixed-point evaluations, accounting for the fact that the fixed point is not solved exactly in practice. • We propose a single-network method that simultaneously recovers both forward and backward optimal transport maps. • We validate our approach empirically, demonstrating strong performance and favorable scaling in high-dimensional settings, including class-conditional scenarios. 2. Background: Quadratic Optimal Transport and Convex Potentials. 2.1. Wasserstein–2 Distance and Kantorovich Duality. Let µ and ν be probability measures on Rd with finite second moments. For the quadratic cost 1 c(x, z) = ∥x − z∥2 , 2

(2.1)

the squared 2-Wasserstein distance is defined as Z 2 (2.2) W2 (µ, ν) = inf γ∈Π(µ,ν)

1 ∥x − z∥2 dγ(x, z), 2 d d R ×R

where Π(µ, ν) denotes the set of couplings with marginals µ and ν. The Kantorovich dual formulation states that (2.3)

W22 (µ, ν) = sup {Ex∼µ [φ(x)] + Ez∼ν [φc (z)]} , φ

where the c-transform is defined by (2.4)



c

φ (z) = inf

y∈Rd

 1 2 ∥y − z∥ − φ(y) . 2

4

Y. PARK, E. GELPHMAN, S. OSHER AND S. WU FUNG

2.2. Brenier’s Theorem and Convex Potentials. For the quadratic cost, optimal transport admits additional structure. Brenier’s theorem states that if µ is absolutely continuous, then there exists a convex function u : Rd → R such that the optimal transport map from µ to ν is given by (2.5)

Tµ→ν (x) = ∇u(x).

Equivalently, writing u(x) = 21 ∥x∥2 + g(x) with g convex, the map can be expressed as (2.6)

Tµ→ν (x) = x + ∇g(x),

which is a result of the following theorem. Theorem 2.1. Suppose µ is absolutely continuous with respect to the Lebesgue measure on Then, ∃ a unique optimal transport map T : Rd → Rd with respect to the quadratic cost (2.1) where T# µ = ν and T is given by (2.6). Rd .

This is a standard result from [68, Theorem 10.28] and we provide a proof in the appendix for completeness. This characterization motivates parameterizing the dual potential as (2.7)

φ(x) = −g(x),

where g is convex. Under this choice, the c-transform becomes   1 (2.8) (−g)c (z) = inf ∥y − z∥2 + g(y) . y∈Rd 2 This expression is precisely the Moreau envelope (or quadratic inf-convolution) of g [66, 54] and corresponds to the solution of a Hamilton-Jacobi equation [22, 21, 37, 23]. 2.3. Proximal Characterization and Transport Maps. For convex and differentiable g, the minimizer   1 2 ∗ (2.9) y (z) = arg min ∥y − z∥ + g(y) y 2 is uniquely defined [55]. The forward and backward transport maps admit the representations

(2.10)

Tµ→ν (x) = x + ∇g(x),

Tν→µ (z) = y ∗ (z).

Therefore, the quadratic optimal transport problem can be formulated entirely in terms of a single convex potential g, whose c-transform is evaluated by solving a strongly convex minimization problem. This proximal characterization forms the basis of the implicit neural formulation developed in the next section. 3. Implicit Neural Optimal Transport. We now develop an implicit neural formulation of quadratic optimal transport based on the proximal characterization of the c-transform.

FIXED-POINT NEURAL OPTIMAL TRANSPORT

5

3.1. Parameterized Dual Formulation. Recall the Kantorovich dual formulation for the quadratic cost: (3.1)

W22 (µ, ν) = sup {Ex∼µ [φ(x)] + Ez∼ν [φc (z)]} . φ

We parameterize the dual potential as (3.2)

φθ (x) = −gθ (x),

where gθ : Rd → R is assumed to be convex in its argument. Under this parameterization, the c-transform becomes   1 c 2 (3.3) (−gθ ) (z) = inf ∥y − z∥ + gθ (y) . y∈Rd 2 Substituting into the dual objective yields     1 2 2 ∥y − z∥ + gθ (y) . (3.4) W2 (µ, ν) = sup −Ex∼µ [gθ (x)] + Ez∼ν inf y 2 θ Equivalently, we minimize the negative dual objective:    1 2 (3.5) L(θ) = Ex∼µ [gθ (x)] − Ez∼ν inf ∥y − z∥ + gθ (y) . y 2 3.2. Constrained Implicit Formulation. For each z ∈ Rd , define the proximal minimizer   1 2 ⋆ ∥y − z∥ + gθ (y) . (3.6) yθ (z) = arg min y∈Rd 2 Since the objective is strongly convex in y, the minimizer exists and is unique. The c-transform can therefore be written explicitly as (3.7)

 1 (−gθ )c (z) = ∥yθ⋆ (z) − z∥2 + gθ yθ⋆ (z) . 2

Substituting this expression into (3.5), the training problem becomes the constrained training problem given by    1 ⋆ min Ex∼µ [gθ (x)] − Ez∼ν ∥yθ (z) − z∥2 + gθ yθ⋆ (z) θ 2   (3.8) 1 ⋆ 2 s.t. yθ (z) = arg min ∥y − z∥ + gθ (y) for z ∼ ν. y 2 Thus, learning reduces to optimizing a single convex potential gθ , with the model output implicitly defined by the proximal optimality condition that characterizes yθ⋆ (z). Importantly, dual feasibility is enforced by construction through (3.6), eliminating the need for adversarial min–max formulations with auxiliary networks [41, 44, 36, 47], as well as time-stepping approaches that require learning entire trajectories [53, 60].

6

Y. PARK, E. GELPHMAN, S. OSHER AND S. WU FUNG

Algorithm 3.1 Implicit Neural Optimal Transport via Fixed-Point Proximal Updates Require: Datasets Dµ , Dν , initial parameters θ, step size η, fixed-point steps K 1: for each training iteration do B 2: Sample minibatches {xi }B i=1 ∼ Dµ , {zi }i=1 ∼ Dν 3: // Fixed-point iteration for proximal operator 4: for each zi in minibatch do 5: Initialize y (0) ← zi 6: for k = 0, . . . , K − 1 do  7: y (k+1) ← y (k) − α ∇y gθ (y (k) ) + y (k) − zi 8: end for 9: ỹi ← stop gradient(y (K) ) // Detach computational graph 10: end for 11: // Compute Loss P P ∇ gθ (xi ) − B1 B 12: g ← B1 B θ i=1 i=1 ∇θ gθ (ỹi ) 13: // Parameter update 14: θ ← θ−ηg 15: end for This perspective places our approach within the framework of implicit deep learning [24, 70], where the model output is defined implicitly via yθ⋆ . However, as we show in Section 3.3, the structure induced by the proximal operator allows us to avoid implicit differentiation, even though computing yθ⋆ still requires solving a fixed-point problem. Once gθ is trained, the forward and backward transport maps are given by (3.9)

Tµ→ν (x) = x + ∇gθ (x),

Tν→µ (z) = yθ⋆ (z).

Both transport directions are therefore represented using a single convex potential. The backward map is computed as the unique minimizer of a strongly convex objective, while the forward map follows directly from Brenier’s theorem. The overall training procedure, including the gradient computation described in the following subsection, is summarized in Algorithm 3.1. Remark 3.1 (PDE interpretation). The proposed formulation admits an alternative interpretation from the perspective of partial differential equations (PDEs). In particular, the optimality condition of quadratic optimal transport can be expressed in terms of the Hamilton– Jacobi (HJ) equation with a quadratic Hamiltonian. Owing to convexity, its solution can be characterized by the Hopf–Lax formula, which coincides with the c-transform in (3.6). From this viewpoint, our method can be interpreted as solving the HJ equation via fixed-point iterations applied to the Hopf–Lax operator [58]. In contrast to prior works, which solve the HJ equation over the entire spatio-temporal computational domain and are therefore computationally demanding in high dimensions, we instead evaluate a local fixed-point map that enforces the corresponding optimality condition. This provides a causality-free mechanism for enforcing the PDE optimality condition without requiring the propagation of full PDE dynamics. Consequently, the proposed approach yields a

FIXED-POINT NEURAL OPTIMAL TRANSPORT

7

significantly more efficient procedure while still enforcing the underlying optimality condition associated with the optimal transport problem. 3.3. Gradient Computation Without Implicit Differentiation. A key property of the formulation (3.8) is that differentiation does not require backpropagating through the inner minimization defining yθ⋆ (z). The gradient of L admits a remarkably simple expression. Lemma 3.2. Assume 0 < γ < 1, gθ (y) is γ-weakly convex, continuously differentiable in y, and L-smooth in θ, and that yθ⋆ (z) is the unique minimizer of   1 2 (3.10) min ∥y − z∥ + gθ (y) . y 2 Then the gradient of L is given by (3.11)

" #   ∂gθ yθ⋆ (z) ∂gθ (x) dL − Ez∼ν . = Ex∼µ dθ ∂θ ∂θ

Proof. The forward pass of the network can be characterized as finding the fixed point of the operator (3.12)

Fθ (y) = y − α(∇g + y − z)

Assume that we can interchange Ex [·] and Ez [·] for any x, z. Then,     dL dg d 1 ∗ 2 ∗ (3.13) = Ex∼µ − Ez∼ν ( ∥y − z∥ + g(y )) dθ dθ dθ 2     dg dy ∗ ∂g dy ∗ ∂g ∗ = Ex∼µ (3.14) − Ez∼ν (y − z) + + dθ dθ ∂y dθ ∂θ  ∗     ∂g ∂g dy ∂g = Ex∼µ (3.15) − Ez∼ν + . (y ∗ − z) + ∂θ ∂y dθ ∂θ If y ∗ is a minimizer of 21 ∥y − z∥2 + gθ (y), then     ∂g(y ∗ ) 1 ∂g(y) 1 2 ∇y ∥y − z∥ + gθ (y) = 2( (y − z)) + = y∗ − z + = 0. 2 2 ∂y y=y∗ ∂y y=y ∗ Thus, if the fixed point problem is solved exactly,     ∂g(x) ∂g(y ∗ ) (3.16) ∇θ L = Ex∼µ − Ez∼ν . ∂θ ∂θ While one might think that using a network to parameterize gθ would require training an implicit network [24, 29, 30, 31], Lemma 3.2 shows that the gradient of the dual objective depends only on explicit derivatives of gθ , evaluated at samples from µ and at the proximal ⋆ points yθ⋆ (z). In particular, no implicit differentiation (computation of dy dθ ) through the inner optimization problem is required. The cancellation follows directly from first-order optimality of the proximal operator and can be interpreted as an instance of the envelope theorem.

8

Y. PARK, E. GELPHMAN, S. OSHER AND S. WU FUNG

3.4. Convergence. Thus, provided that the fixed point problem is solved exactly, JFB is not needed to compute the gradient with respect to θ of the loss function. This does not occur in practice, but it is proved in this section that stochastic gradient descent (SGD) converges to a local minimum of the loss function if the fixed point problem is solved within a specified error tolerance. For ease of presentation, we provide the proofs in Appendix A. 3.4.1. Approximate Gradient. Suppose the fixed point problem is not solved exactly, i.e. instead of the true fixed point y ∗ an approximate fixed point ỹ is computed. Then, the stochastic gradient d˜θ is given by     ∂g(x) ∂g(ỹ) ˜ (3.17) dθ = Ex∼µ − Ez∼ν . ∂θ ∂θ 3.4.2. Preliminary Assumptions and Fixed Point Computation Lemma. Assumption 3.3. The objective function L(θ) is bounded from below by Linf in some open subset of its domain.. The function g is C 1 with respect to all variables and the gradients of g with respect to θ and y, are Lθ - and Ly -Lipschitz, respectively. Furthermore, assume the 2nd order partial derivatives of g with respect to θ exist. From this point onwards, denote x ∼ µ and z ∼ ν as just x and z, respectively. Lemma 3.4. Under Assumption 3.3, suppose the approximate fixed point ỹ satisfies ∥(ỹ − z) + ∇y gθ (ỹ)∥ < ϵp .

(3.18) Then, ∃ϵ > 0 such that

∥y ∗ − ỹ∥ < ϵ.

(3.19)

3.4.3. Upper Bound on 2nd Moments of Exact and Approximate Gradients. R 3.5. Under the assumptions of Lemma 3.4 along with ∥x∥2 dµ(x) < ∞ and R Lemma ∥z∥2 dν(z) < ∞, " # " # " # ∂g(yθ∗ (z)) 2 ∂g(xθ ) 2 ∂g(y˜θ (z)) 2 < ∞, Ez < ∞, and Ez < ∞. (3.20) Ex ∂θ ∂θ ∂θ Theorem 3.6. Under the assumptions of Lemma 3.5 The norm of the true gradient ∇θ L and the stochastic gradient d˜θ squared is bounded above, i.e. ∃0 < M̃ < +∞ and ∃0 < ML < +∞ such that ∀θ     2 ∂g(ỹ (z)) ∂g(x ) θ θ 2 ≤ M̃ (3.21) ∥d˜θ ∥ = Ex − Ez ∂θ ∂θ and (3.22)

∥∇θ L∥2 = Ex



   ∂g(xθ ) ∂g(ỹθ (z)) 2 ≤ ML . − Ez ∂θ ∂θ

FIXED-POINT NEURAL OPTIMAL TRANSPORT

9

3.4.4. Approximate Gradient is a Descent Direction. Theorem 3.7. Under the assumptions of Theorem 3.6, suppose the norm of the difference between the exact and computed fixed point satisfies, ∀θ and ∀z, ∥ỹθ (z) − yθ∗ (z)∥ < ϵ ≤ V ∥∇θ L∥2 √ √ , for some 0 < V < 1. Then, ∃U = 1 − V > 0 such that ∀θ Lθ ( Mx + Mz ) ⟨∇θ L, d˜θ ⟩ ≥ U ∥∇θ L∥2 .

(3.23)

3.4.5. Convergence Results. We begin by introducing notation to formalize the stochasticity arising in the training process. Let {ξj }j≥0 denote a sequence of independent random variables representing the sampling procedure used to construct the stochastic gradient d˜ξj (θ). In particular, ξj corresponds to the random draw of initial conditions x ∼ ρ used to compute the stochastic gradient update at iteration j. We analyze the convergence of SGD when d˜θ is used as a stochastic gradient surrogate. Specifically, we consider the iterative scheme θj+1 = θj − αj d˜ξj (θj ),

(3.24)

j ≥ 0,

for minimizing the loss function over θ ∈ Rp . Here, d˜ξj (θj ) denotes the JFB update computed using either a single sample or a minibatch of samples with corresponding learning rate αj at iteration j. Following the notation of [11], we use Eξj [·] to denote the conditional expectation with respect to the randomness at iteration j, given the current iterate θj . Since θj depends on the sequence of random variables {ξ0 , ξ1 , . . . , ξj−1 }, we also consider the total expectation of the objective with respect to all prior randomness, which we write as h    i (3.25) E[Ex [Jx (θj )]] = Eξ0 Eξ1 · · · Eξj−1 Ex [Jx (θj )] · · · . With this notation in place, we establish the following Lemma, which is used to prove the main result. Lemma 3.8. Under the assumptions of Theorem 3.7, the SGD iterations (3.24) satisfy Eξj [L(θj+1 )] − L(θj ) ≤ −αj U ∥∇θ L(θj )]∥2 + αj2 Lθ M̃ .

(3.26)

If 0 < αj ≤ ULMM̃L , then it follows that Eξj [L(θj+1 )] − L(θj ) ≤ 0. θ

With this result, the main results of this paper can be proven. Theorem 3.9. Suppose the sequence of learning rates {αj }∞ j=0 is monotonically decreasing P∞ P∞ 2 P U ML and satisfies j=0 αj = ∞, j=0 αj < ∞, and 0 < α0 ≤ L M̃ . Let AK = K−1 j=0 αj . Then, θ under the assumptions of Lemma 3.8 the SGD iteration (3.24) satisfies   K X 1 αj ∥∇θ L(θj )∥2  = 0. lim E  K→∞ AK j=0

10

Y. PARK, E. GELPHMAN, S. OSHER AND S. WU FUNG

In other words, the weighted Cesaro sum of the sequence {∥∇θ L(θj )∥2 }∞ j=0 converges in (total) expectation to 0. Using Theorem 3.9, one can then use standard SGD analysis to show the following theorem and corollary. Theorem 3.10. Under the assumptions of Theorem 3.9, the SGD iteration (3.24) satisfies h i lim inf E ∥∇θ L(θj )∥2 = 0. j→∞

Using Theorem 3.9, we can also prove convergence in probability to a critical point. Corollary 3.11. Suppose the assumptions of Theorem 3.9 hold. For any K ∈ N let j(K) ∈ {0, 1, ..., K} represent a random index chosen with probabilities proportional to {αj }K j=0 . Then, K {∥∇θ L(θj )∥}j=0 → 0 as K → ∞ in probability. 3.5. Class-Conditional Optimal Transport. In many applications, the source and target measures admit class-wise decompositions (3.27)

µ=

K X

π k µk ,

k=1

ν=

K X

πk νk ,

k=1

where µk and νk denote the distributions restricted to class k. Class-conditional optimal transport (CC-OT) enforces that mass is transported only within corresponding classes by solving, for each k, i h (3.28) Tk = arg min Ex∼µk 21 ∥x − T (x)∥2 . T# µk =νk

Rather than defining separate potentials {gk }K k=1 or assuming well-separated supports, we parameterize a single class-conditional potential gθ (x, k) : Rd × {0, . . . , K − 1} → R,

(3.29)

where k is a class label. The network receives the input x concatenated with a one-hot encoding of k. Then, the class-conditional dual objective becomes (3.30)

" # K−1 i h1   1 X 2 ⋆ ⋆ LCC (θ) = Ex∼µk gθ (x, k) − Ez∼νk ∥z − yθ (z, k)∥ + gθ (yθ (z, k), k) , K 2 k=0

where the backward OT map for class k is defined as the fixed-point solution (3.31)

yθ⋆ (z, k) = arg min y

n1 2

o ∥y − z∥2 + gθ (y, k) ,

z ∼ νk .

Hence, conditioning the potential on class labels allows a single network to naturally represent class-conditional transport maps while maintaining a unified convex potential.

FIXED-POINT NEURAL OPTIMAL TRANSPORT

11

4. Experiments. We evaluate the proposed method across a variety of datasets and experimental settings to assess its performance and efficiency. We compare against existing baseline methods and further conduct ablation studies to analyze the contribution of each component. All experiments are implemented in PyTorch and detailed implementation details are provided in Appendix B. The experiments in Section 4.4 are conducted on an NVIDIA RTX Blackwell 6000 GPU, while all remaining experiments are performed on a single NVIDIA TITAN V GPU (12GB). The implementation of our method is publicly available at: https: //github.com/Yebbi/ImplicitOT 4.1. Evaluation on High-dimensional Gaussian Distributions. Quantitative evaluation of OT methods for general distributions is often challenging due to the lack of closed-form solutions. To address this, we focus on Gaussian distributions, µ = N (0, Σµ ) and ν = N (0, Σν ), for which the OT map admits an analytical solution: (4.1)

−1

1

1

1

−1

Tµν∗ (x) = Σµ 2 (Σµ2 Σν Σµ2 ) 2 Σµ 2 x.

Following the protocol of [42], we consider dimensions d ∈ [2, 64], constructing Σµ and Σν with random orthonormal eigenvectors and eigenvalues whose logarithms are sampled uniformly from [−2, 2]. Baseline Methods. We compare our method with several established OT approaches. These include NOT [44], which directly parameterizes the transport map and is trained using the weak formulation; MM-v1 [65, 42], a min–max framework based on input-convex neural networks (ICNNs) that alternates between optimizing a potential function and its convex conjugate; and MM:R [42], which also follows a min–max formulation but does not impose convexity, instead learning separate networks for the forward and backward maps using a negative Wasserstein loss with an additional conjugacy regularization. We additionally include NCF [56], which constructs transport maps based on the characteristics of the Hamilton–Jacobi equation. Furthermore, we evaluate LS [63], a dual OT solver with entropic regularization, and WGAN-QC [48], which adopts a WGAN architecture with a quadratic cost. Except for NOT, all methods, including ours, employ a shared network architecture for modeling the potential functions. Evaluation Metrics. Performance is evaluated using the unexplained variance percentage (UVP) [41]. Given a predicted transport map T̂ : µ → ν and the ground-truth optimal transport map T ∗ , the UVP is defined as

(4.2)

T̂ − T ∗ 2   L (µ) L2 -UVP T̂ := 100 (%). Var(ν)

We further report computational efficiency in terms of training and inference time, peak memory consumption, and storage requirements for bidirectional transport maps. Results. Table 1 reports the UVP scores across different models and dimensionalities, while Figure 1 summarizes the corresponding computational costs. Compared to all baselines, our method consistently learns substantially more accurate OT maps across all tested dimensions. While baseline approaches exhibit a noticeable degradation in accuracy as the dimensionality

12

Y. PARK, E. GELPHMAN, S. OSHER AND S. WU FUNG

Table 1: Quantitative evaluation on Gaussian distributions. UVP (↓) is measured across different OT methods as the data dimension d increases. Method

d=2

d=4

d=8

d = 16

d = 32

d = 64

NOT WGAN-QC LS MM-v1 MM:R NCF Ours

77.248 1.596 5.806 0.161 0.012 0.010 0.013

125.419 5.897 9.781 0.172 0.048 0.021 0.016

114.056 31.0367 15.963 0.173 0.117 0.086 0.046

176.086 59.314 25.232 0.210 0.202 0.146 0.053

182.287 113.237 41.445 0.374 0.354 0.436 0.054

196.831 141.407 55.360 0.415 0.604 0.858 0.0822

Training Time

Max Memory Allocation

0.6

140

0.4 0.3 0.2

120 100 80 60

0.1

40

0.0

20 2

4

8

16

Dimension NOT

32

WGAN-QC

Memory (MB)

Memory (MB)

0.5

Time (s)

Bidirectional OT Map Storage 0.5

64

0.3 0.2 0.1

2

LS

0.4

4

8

16

Dimension

MM-v1

32

MM:R

64

2

HJ-PINN

4

8

16

Dimension NCF

32

64

Ours

Figure 1: Computational comparison. Training time (s/epoch), peak memory (MB) during training, and memory (MB) for storing bidirectional OT maps are reported across models and dimensions.

increases, our model maintains a stable error profile even in high-dimensional settings. Notably, a single fixed experimental configuration was used for our model across all dimensions, without any dimension-specific tuning. These results highlight the accuracy, robustness, and scalability of our approach for high-dimensional OT. 4.2. Real-World Physics Data Distributions. To evaluate the extent to which the proposed method can handle complex and practically relevant distributions, we consider realworld data sets with diverse and heterogeneous characteristics. Data. We conduct experiments on several benchmark data sets from the University of California Irvine (UCI) machine learning data repository [40], including POWER, GAS, HEPMASS, and MINIBOONE, as well as the BSDS300 data set consisting of natural image patches. These data sets arise from a variety of domains, including high-energy physics experiments, household power consumption, and chemical sensor measurements of gas mixtures. As a result, they exhibit a wide range of statistical properties, dimensionalities, and structural features. In these settings, the data are not provided as paired samples from two distributions. Instead, we formulate an optimal transport problem by taking a standard Gaussian distri-

FIXED-POINT NEURAL OPTIMAL TRANSPORT

Dataset

Metric MMD #Params MMD #Params MMD #Params MMD #Params MMD #Params

Power (d = 6) Gas (d = 8) HEPMASS (d = 21) MINIBOONE (d = 43) BSDS300 (d = 63)

FFJORD 4.34×10−5 43K 1.02×10−4 279K 1.58×10−5 547K 2.84×10−4 821K 6.52×10−3 6.7M

13

RNODE 5.64×10−5 43K 8.03×10−5 279K 1.58×10−5 547K 2.84×10−4 821K 1.64×10−2 6.7M

OT-Flow 4.68×10−5 18K 2.47×10−4 127K 1.58×10−5 72K 2.84×10−4 78K 4.24×10−4 297K

NCF 2.56×10−4 8.25K 2.24×10−4 8.90K 6.84×10−5 13.06K 2.78×10−3 29.84K 6.28×10−4 156K

Ours 1.74×10−5 8.25K 7.34×10−5 8.90K 1.79×10−5 13.06K 1.19×10−3 29.84K 3.30×10−4 156K

Table 2: Comparison of distribution matching performance (MMD) and model size (#parameters) across real-world datasets. Lower MMD indicates better distribution matching.

bution as the source and the empirical data distribution as the target. The model therefore learns a transport map that pushes forward samples from the Gaussian reference distribution to the target data distribution. This setup provides a practical test of whether the learned transport map can adapt to heterogeneous and high-dimensional distributions in the absence of ground-truth correspondences. Baseline Methods. We compare against NCF, the most recent model that achieves the strongest performance in the Gaussian setting among OT-based approaches. In addition, we include continuous flow-based methods, namely FFJORD [33], RNODE [28], and OT-Flow [53], as representative approaches for learning continuous transport dynamics under likelihoodbased or optimal transport-inspired objectives. These methods provide strong and standard baselines for evaluating both the quality of learned transport maps and their computational efficiency. Our study focuses on optimal transport formulations for distribution matching rather than general-purpose generative modeling. Accordingly, we restrict our comparisons to methods operating within a comparable transport or continuous-flow framework. Evaluation Metrics. For evaluation, we adopt the same setup as [53], using the unbiased Maximum Mean Discrepancy (MMD) estimator with Gaussian kernels to measure the distance between generated and target distributions. Specifically, we use a kernel of the form k(x, x′ ) =

K X

exp(−αj ∥x − x′ ∥2 ),

j=1

where the bandwidth parameters {αj } are selected as in prior work. Given two sets of samples 1 2 {xi }ni=1 and {yj }nj=1 , the estimator is given by MMD2 =

1 n1 (n1 − 1)

X i̸=j

k(xi , xj ) +

1 n2 (n2 − 1)

X i̸=j

k(yi , yj ) −

2 X k(xi , yj ). n1 n2 i,j

To assess computational efficiency, we additionally measure the number of trainable parameters in each model.

Y. PARK, E. GELPHMAN, S. OSHER AND S. WU FUNG

Predicted

True

14

(5, 1)-slice

(5, 2)-slice

(5, 4)-slice

(5, 6)-slice

Figure 2: Each column shows a different slice of the UCI Physics Gas dataset, comparing real (top) and predicted (bottom) samples.

Results. The results are summarized in Table 2, while quantitative comparisons for Gas, Hepmass, and Miniboone are provided in Figures 2, 3, and 5. Since these datasets are highdimensional, we visualize two-dimensional slices along informative coordinate axes to better illustrate the various structure of the distributions. As shown in the figures, the target distributions exhibit highly complex and heterogeneous structures, yet the proposed method is able to accurately capture their geometric characteristics across different datasets. Despite using substantially smaller network architectures, our method achieves on-par or better performance in several benchmarks against continuous flowbased models. While we observe slightly higher MMD in some settings, this comparison should be interpreted in light of the fact that our models are significantly more compact. Moreover, continuous flow-based approaches are primarily developed within likelihood-based generative modeling frameworks and are not directly designed to address general optimal transport problems between arbitrary distributions. In addition, they rely on solving continuous-time dynamics via numerical ODE integration at inference time, which introduces additional computational overhead. In contrast, our formulation directly targets the optimal transport problem in a more general setting and does not require iterative ODE solving during sampling, resulting in significantly more efficient inference. Furthermore, compared to NCF, a representative OT-based method, our approach demonstrates consistently better performance on complex, high-dimensional distributions, indicating improved robustness in challenging non-Gaussian settings. Overall, these results highlight the generality, accuracy, and computational efficiency of the proposed method. 4.3. 2D Class-Conditional Transport. We present experimental results on a two-dimensional (@D) synthetic dataset consisting of class-labeled samples, designed to evaluate class-conditional

15

Predicted

True

FIXED-POINT NEURAL OPTIMAL TRANSPORT

(3, 21)-slice

(6, 21)-slice

(18, 21)-slice

(19, 21)-slice

Target Source

Figure 3: Each column shows a different slice of the UCI Physics Hepmass dataset, comparing real (top) and predicted (bottom) samples.

8

1

9

5

4

0

9

4

1

8

1

9

1

1

9

6

Figure 4: Sample image generation using our conditional optimal transport map, trained to transport the distribution of Fashion-MNIST images to that of MNIST. The top row shows input images from Fashion-MNIST, while the bottom row displays the corresponding outputs generated by the model in the MNIST domain. Each column is conditioned on the integer label indicated above the top row.

optimal transport. Data. The dataset is constructed as a mixture of Gaussian components, where each component corresponds to a distinct class label. This setting allows us to assess whether the learned transport respects both inter-class alignment and intra-class structure preservation. Results. Figure 6 illustrates the learned transport behavior across three 2D Gaussian mixture datasets. Each data point is associated with a class label, and transport is performed in a class-conditional manner. The results show that our method successfully aligns corresponding classes between source and target distributions while maintaining clear separation between different classes. In addition to global distribution matching, the learned transport preserves the local geometry within each class, indicating that the model captures both class-level correspondence

Y. PARK, E. GELPHMAN, S. OSHER AND S. WU FUNG

Predicted

True

16

(1, 7)-slice

(1, 14)-slice

(1, 16)-slice

(1, 41)-slice

Figure 5: Each column shows a different slice of the UCI Physics Miniboone dataset, comparing real (top) and predicted (bottom) samples.

and fine-grained structural consistency. This suggests that the proposed formulation can naturally extend to conditional settings without requiring architectural modifications or additional networks. 4.4. Class-Conditional Image Translation. We further evaluate our method on classconditional optimal transport for images. In this setting, the goal is to transport samples between image distributions while preserving semantic class information. This setting allows us to evaluate whether the learned transport map respects both global distribution alignment and class-level consistency in real image data. In particular, it tests whether the model can disentangle class structure while performing meaningful distributional alignment, which is a key challenge in conditional generative modeling and optimal transport. Data. We use two standard benchmark datasets: MNIST [46] and Fashion-MNIST (FMNIST) [71]. MNIST consists of grayscale images of handwritten digits, while FMNIST contains images of fashion products such as clothing and shoes. Both datasets contain 60,000 training images and 10,000 test images, with each dataset composed of 10 classes corresponding to their semantic categories. Direct application of optimal transport in the pixel space is not well-defined due to the fact that the intrinsic dimension of image datasets is substantially lower than the ambient dimension [59]. To address this issue, we perform optimal transport in a latent space learned via a variational autoencoder (VAE). Specifically, both FMNIST and MNIST images are encoded into a shared latent representation, and optimal transport is performed in this latent space. We use a latent dimension of 15, which provides a compact representation while preserving the essential semantic structure of the data. The transported latent codes are subsequently decoded back into the image space using the VAE decoder.

FIXED-POINT NEURAL OPTIMAL TRANSPORT

Forward OT

Backward OT

Four-Modes

Horizontal Swapped

Cross Ring

Data

17

Figure 6: Class-conditional optimal transport on Gaussian mixtures. Each row corresponds to a different class-structured problem. The first column shows the empirical data distributions, the second column visualizes the learned forward transport map, and the third column shows the learned backward transport. Colors indicate different distributions, while marker shapes differentiate individual classes within each distribution.

Baseline Models. Following prior works [6, 56], we compare against a wide range of representative baselines spanning generative models and optimal transport methods. Pixel-level adaptation methods include one-to-many translation models such as MUNIT [38] and AugCycleGAN [1, 75]. For semi-supervised or label-aware transport, we include OTDD flow [2, 3], which leverages gradient flows to preserve class structure during transport. We further evaluate General Discrete Optimal Transport (DOT), implemented using the Sinkhorn algorithm [17] with Laplacian cost regularization [16] We also consider neural optimal transport methods [44, 25, 6], evaluated under the qua-

18

Y. PARK, E. GELPHMAN, S. OSHER AND S. WU FUNG

dratic cost 12 ∥x − y∥2 (denoted W2) the γ-weak quadratic cost (W2,γ with γ = 1), which allows one-to-many mappings, and class-conditional quadratic cost (FG). Finally, we include HJB characteristic-based NCF [56], which solve class-conditional transport via Hamilton–Jacobi dynamics. Evaluation Metric. To assess the quality of transported images, we evaluate the Fréchet Inception Distance (FID). FID measures the distance between feature distributions of real and generated samples extracted from a pretrained Inception network. Assuming Gaussian approximations of the feature embeddings with means and covariances (µr , Σr ) and (µg , Σg ) for real and generated data, respectively, FID is defined as:   FID = ∥µr − µg ∥2 + Tr Σr + Σg − 2(Σr Σg )1/2 . FID is well-suited for our setting as it captures both sample fidelity and diversity, and is sensitive to distributional shifts across classes, making it a standard metric for evaluating image translation quality in class-conditional optimal transport tasks. To further assess the class-preserving property of the learned transport map, we follow the evaluation protocol of [6] and measure the classiciation accuracy using a pretrained ResNet-18 classifier [35], which achieves over 95% accuracy on the target domain. For each transported sample T (x, z), we obtain the predicted class label and compare it with the ground-truth label of the source sample x. A transport is considered correct if the predicted label of the transported sample T (x) matches the label of x, thereby measuring class consistency under the learned mapping. Results. FID scores and classification accuracies are summarized in Tables 4 and 3, respectively. Moreover, qualitative results are presented in Figure 4. As shown in Table 4, our method achieves the lowest FID among all competing baselines. In particular, it significantly outperforms both pixel-level generative models (AugCycleGAN and MUNIT) and classical optimal transport approaches (W2, W2-γ, DOT, and OTDD flow). Compared to neural OT methods such as GNOT and NCF, our approach consistently yields improved perceptual quality. As shown in Table 3, our method achieves the highest classification accuracy among all baselines, indicating that it not only aligns marginal distributions but also preserves class-level structure more effectively than competing approaches. This is further illustrated in Figure 4, where the proposed method successfully performs class-consistent transport, aligning samples according to their respective class labels while preserving semantic structure across domains. These results suggest that the proposed singlenetwork minimization framework effectively captures both distributional alignment and classconditional structure in the latent space, leading to more faithful class-conditional transport compared to adversarial and flow-based neural OT methods. 4.5. Ablation Studies. 4.5.1. Ablation Study on Fixed-Point Tolerance. In practical implementations, the proximal subproblem in (3.6) is solved using an iterative fixed-point scheme, which is terminated once a prescribed tolerance ϵp is reached. While the theoretical analysis assumes exact computation of yθ⋆ (z), training is carried out using an approximate solution ỹθ (z). It is therefore important to quantify the effect of the stopping tolerance on both accuracy and computational efficiency.

FIXED-POINT NEURAL OPTIMAL TRANSPORT

19

Table 3: Accuracy ↑ of transported samples. Pixel-level Datasets

Discrete / OTDD

MUNIT AugCG OTDD SinkLpL1

FMNIST → MNIST

8.93

12.03

10.28

10.67

Neural OT W2

W2,γ

Ours / Neural Flow

FG

NCF FG(no z) Ours

10.96 8.02 83.72 83.42

82.79

96.00

Table 4: FID ↓ of transported samples. Pixel-level Datasets

Discrete / OTDD

Neural OT

Ours / Neural Flow

MUNIT AugCG OTDD SinkLpL1 W2 W2,γ FG NCF FG(no z) Ours

FMNIST → MNIST

7.91

26.35

>100

>100

7.51 7.02 5.26 18.27

7.14

4.57

To this end, we perform an ablation study on the 32-dimensional Gaussian optimal transport task presented in Section 4.1, systematically varying the tolerance ϵp . For each configuration, we report the forward and backward transport errors (UVP), along with the average training time per epoch. Tolerance ϵp 1 × 10−4 5 × 10−4 1 × 10−3 5 × 10−3 1 × 10−2 5 × 10−2 1 × 10−1 5 × 10−1

Forward UVP 4.909 × 10−2 5.125 × 10−2 5.435 × 10−2 5.492 × 10−2 5.449 × 10−2 5.622 × 10−2 6.255 × 10−2 2.229 × 10−1

Backward UVP 4.968 × 10−2 5.906 × 10−2 5.777 × 10−2 5.861 × 10−2 5.925 × 10−2 6.051 × 10−2 6.557 × 10−2 1.724 × 10−1

Time (s) 0.986 0.917 0.856 0.739 0.640 0.516 0.479 0.277

Table 5: Sensitivity of transport accuracy and computational cost to the fixed-point stopping tolerance. The results exhibit a consistent trade-off between computational cost and estimation accuracy, while simultaneously demonstrating that the method remains robust with respect to the choice of fixed-point tolerance. As predicted by the theory, increasing the tolerance ϵp introduces a discrepancy between the approximate solution ỹθ (z) and the exact minimizer yθ⋆ (z), which leads to some degradation in accuracy. At the same time, a larger tolerance reduces the number of iterations required for convergence of the fixed-point procedure, thereby improving computational efficiency. Notably, across the range ϵp ∈ [10−5 , 10−2 ], the degradation in transport accuracy is marginal, indicating that the method is not highly sensitive to moderate inexactness in the proximal computation. This suggests that the learned transport maps are relatively insensitive to approximation errors in the fixed-point solution within this regime. This observation

20

Y. PARK, E. GELPHMAN, S. OSHER AND S. WU FUNG

Table 6: Comparison of network parameterization. We report forward and backward UVP values. Convex Activation Tanh SoftPlus CeLU

Fwd UVP −1

7.173 × 10 6.857 × 100 6.826 × 100

Non-convex

Bwd UVP −1

7.565 × 10 5.945 × 100 5.794 × 100

Fwd UVP −1

1.730 × 10 6.996 × 10−2 5.4 × 10−2

Bwd UVP 2.184 × 10−1 7.551 × 10−2 5.67 × 10−2

is consistent with the analysis in Section 3.4, which guarantees that sufficiently accurate approximate fixed points remain close to the true minimizer. From a practical perspective, this empirical behavior indicates that, within a reasonable tolerance regime, the approximate fixed-point ỹθ (z) provides a sufficiently accurate surrogate for yθ⋆ (z) and does not significantly distort the gradient signal used during training. In contrast, when the tolerance becomes excessively large (e.g., ϵp ≥ 10−1 ), the approximation error becomes non-negligible and leads to a clear deterioration in performance. Overall, these findings suggest that solving the inner optimization problem to very high precision is unnecessary in practice. Instead, moderately accurate fixed-point solutions achieve a favorable balance, providing substantial computational savings while maintaining transport accuracy. 4.5.2. Ablation Study on Potential Parameterization. To investigate how different neural parameterizations of the Kantorovich potential g affect the behavior of our method, we conduct an ablation study focusing on architectural convexity. Specifically, we consider the following variants: • Convex Network: an input-convex neural network (ICNN) where convexity is enforced via weight projection at each iteration. • Nonconvex Network: an ICNN trained without enforcing convexity constraints. This setting preserves the same underlying architecture, allowing us to isolate the effect of convexity enforcement without confounding architectural differences. In addition, to examine the effect of nonlinearities, we evaluate each architecture using three different activation functions: Tanh, SoftPlus, and CeLU. This allows us to disentangle the impact of convexity constraints from that of the activation choice on optimization stability and transport performance. We evaluate each model under 32-dimensional case. The results are summarized in Table 6, while the evolution of the fixed-point iteration depth and residual y k+1 − z + ∇g(y k+1 )

are illustrated in Figure 7. Trade-off Perspective.. The results in Table 6 and Figure 7 highlight a critical trade-off in our method between the expressive power of the neural potential and the stability of the induced fixed-point iterations. Since our approach relies on solving a proximal fixed-

FIXED-POINT NEURAL OPTIMAL TRANSPORT

21

Figure 7: Ablation study on network parameterizations. We compare convex and nonconvex network with different activation functions across (a) OT map accuracy (UVP), (b) the number of fixed-point iteration to converge, and (c) convergence residual.

point problem at every step, good performance requires both (i) a sufficiently expressive parameterization of the Kantorovich potential g and (ii) stable and efficient convergence of the fixed-point dynamics. ICNNs with explicit convexification enforce convexity through architectural constraints, yielding monotone gradient fields and, in principle, well-behaved fixed-point updates. Indeed, we observe that this setting leads to the most stable and fastest convergence in terms of fixedpoint iterations. However, such strong enforcement of convexity degrades overall performance. While convexity is theoretically well-motivated by the Kantorovich formulation, strong convexification restricts the effective parameter space and introduces optimization bias. As a result, the learned potential lacks the flexibility needed to accurately capture the geometry of the target transport, leading to suboptimal values of the dual objective. In contrast, ICNNs without explicit convexification achieve the best overall performance. In this case, convexity is not explicitly enforced through hard projection. This provides a favorable balance: the model retains sufficient flexibility to approximate complex transport structures while still benefiting from a partially monotone gradient field. Empirically, this leads to maintain stable fixed-point iterations—while achieving substantially better optimization of the OT objective compared to convexified ICNNs. The choice of activation function further reveals the importance of this balance. In particular, Tanh consistently underperforms both in terms of objective value and fixed-point convergence. Although Tanh is Lipschitz and bounded, its saturation behavior leads to vanishing gradients, especially in deeper compositions. This results in poorly conditioned gradient fields for ∇gθ which in turn degrades the contraction properties required for stable fixed-point iterations. Consequently, the proximal updates become less reliable, often requiring more iterations and converging to inferior solutions. On the other hand, smoother and non-saturating activations such as SoftPlus and CeLU provide better gradient flow and improved conditioning of the optimization landscape. These activations enable richer functional representations while maintaining sufficiently stable dynamics, leading to both improved fixed-point convergence and superior transport performance.

22

Y. PARK, E. GELPHMAN, S. OSHER AND S. WU FUNG

5. Conclusion. We proposed a novel formulation of quadratic optimal transport based on a single neural potential, in which both forward and backward transport maps are recovered through a unified minimization problem. By exploiting the proximal structure of the dual formulation, the objective reduces to a fixed-point problem whose gradient can be computed without implicit differentiation. This yields a simple and stable training procedure that avoids saddle-point optimization, auxiliary networks, and unrolled dynamics, while maintaining strong empirical performance in high-dimensional and class-conditional settings. Future work includes extending the evaluation to more diverse and complex datasets, as well as establishing stronger theoretical guarantees, in particular convergence to the true optimal transport map under practical neural network parameterizations. Furthermore, generalizing the proposed framework beyond the quadratic cost to more general cost functions remains an important direction for broadening its applicability. Appendix A. Proofs. A.1. Proof of Theorem 2.1. Proof. Because µ, ν are probability measures on Rd with finite second moments, and µ is absolutely continuous with respect to the Lebesgue measure on Rd , all of the assumptions of Brenier’s Theorem [27, Theorem 2.5.10] are true. Then by [27, Corollary 2.5.12] ∃ a unique optimal transport map T : Rd → Rd such that T# µ = ν and T = ∇u for some convex u : Rd → R ∪ {+∞}. By [68, Theorem 10.28], T satisfies ∇g(x) + ∇x c(x, T (x)) = 0 µ-almost everywhere for some c-convex function g : Rd → R ∪ {+∞}. Thus, µ-a.e.,   1 2 ∇g(x) + ∇x (A.1) ∥x − T (x)∥ = 0 2 (A.2) ∇g(x) + x − T (x) = 0. . Thus, (A.3)

T (x) = x + ∇g(x).

Finally, since u and c(x, y) are convex, g must be 1-weakly convex. A.2. Proof of Lemma 3.4. Proof. Because ∇y gθ is assumed to be Ly Lipschitz in y it follows that ∇y gθ is continuous in y. Therefore, the function ỹ 7→ (ỹ − z) + ∇y gθ (ỹ) is continuous at all points in its domain. Then, by definition, ∀ϵ′ > 0, ∃δ ′ > 0 such that ∥(ỹ − z) + ∇y gθ (ỹ) − ((y ∗ − z) − ∇y gθ (y ∗ )) ∥ < ϵ′ whenever ∥ỹ − y ∗ ∥ < δ ′ . Because y ∗ is an exact minimizer of 21 ∥y − z∥2 + gθ (y), it follows that whenever ∥ỹ − y ∗ ∥ < δ ′ , ∥(ỹ−z)+∇y gθ (ỹ)−((y ∗ − z) − ∇y gθ (y ∗ )) ∥ = ∥(ỹ−z)+∇y gθ (ỹ)−0∥ = ∥(ỹ−z)+∇y gθ (ỹ)∥ < ϵ′ . Thus, the result follows by choosing ϵ = δ ′ corresponding to ϵ′ = ϵp .

FIXED-POINT NEURAL OPTIMAL TRANSPORT

23

A.3. Proof of Lemma 3.5. Proof. Because ∂g ∂θ is Lθ -Lipschitz in θ and by the result of Lemma 3.4, ∃ϵ > 0 such that ∥ỹθ (z) − yθ∗ (z)∥ < ϵ, it follows that " # h i ∂g(ỹθ (z)) ∂g(ỹθ (z)) 2 0 ≤ Ez − ≤ Ez Lg ∥ỹθ (z) − yθ∗ (z)∥2 ≤ ϵLg . ∂θ ∂θ   2 ∂g(ỹθ (z)) ∂g(ỹθ (z)) , the above inequality becomes Expanding Ez − ∂θ ∂θ " " " # # # ∂g(ỹθ (z)) 2 ∂g(ỹθ (z)) ⊤ ∂g(yθ∗ (z)) ∂g(ỹθ (z)) 2 0 ≤ Ez + Ez − 2Ez ≤ ϵLg . ∂θ ∂θ ∂θ ∂θ R R 2 dµ(x) and It follows from the above inequality and the assumption that ∥x∥ ∥z∥2 dν(z)       2 2 ⊤ ∂g(y ∗ (z)) θ , Ez ∂g(ỹ∂θθ (z)) , and Ez ∂g(ỹ∂θθ (z)) are all finite. that Ez ∂g(ỹ∂θθ (z)) ∂θ   2 θ) is finite by picking any two points xθ,1 , xθ,2 A near-identical argument shows Ex ∂g(x ∂θ θ) such that ∥xθ,1 − xθ,2 ∥ < ϵ and using the fact that ∂g(x is Lg -Lipschitz. ∂θ

A.4. Proof of Theorem 3.6. ˜ Proof. Expanding the definition of d,     ∂g(xθ ) ∂g(ỹθ (z)) 2 2 ˜ (A.4) ∥dθ ∥ = Ex − Ez ∂θ ∂θ  2         ∂g(ỹθ (z)) ∂g(xθ ) ∂g(ỹθ (z)) 2 ∂g(xθ ) , Ez + Ez (A.5) = Ex − 2 Ex ∂θ ∂θ ∂θ ∂θ # " # "      2 ∂g(ỹθ (z)) ∂g(ỹθ (z)) 2 ∂g(xθ ) ∂g(xθ ) , Ez + Ez (A.6) = Ex − 2 Ex , ∂θ ∂θ ∂θ ∂θ where (A.6) follows by Jensen’s inequality. By the result of Lemma     3.5, ∃0 < Mx < +∞, 0 < 2

2

θ) Mz < +∞ such that Ex ∂g(x < Mx and Ez ∂g(ỹ∂θθ (z)) < Mz . Applying Jensen’s ∂θ     2 2 θ) , Ez ∂g(ỹ∂θθ (z)) yields inequality to Ex ∂g(x ∂θ     p p ∂g(xθ ) ∂g(ỹθ (z)) Ex < Mx and Ex < Mz . ∂θ ∂θ

It then follows by the Cauchy-Schwarz inequality that (A.7)            ∂g(xθ ) ∂g(ỹθ (z)) ∂g(xθ ) ∂g(ỹθ (z)) ∂g(xθ ) Ex ≤ Ex , Ez ≤ Ex − Ex ∂θ ∂θ ∂θ ∂θ ∂θ Combining (A.6) and (A.7), 0 ≤ ∥d˜θ ∥2 ≤ Mx + Mz + 2

p

Mx Mz = M̃ < +∞.

 Ex

∂g(ỹθ (z)) ∂θ

 .

24

Y. PARK, E. GELPHMAN, S. OSHER AND S. WU FUNG

An identical argument shows the existence of ML such that 0 ≤ ∥∇θ L∥2 ≤ ML < +∞. A.5. Proof of Theorem 3.7. Proof. Expanding the inner product,          ∂g(yθ∗ (z)) ∂g(xθ ) ∂g(xθ ) ∂g(ỹθ (z)) ˜ ⟨∇θ L, dθ ⟩ = Ex − Ez , Ex − Ez ∂θ ∂θ ∂θ ∂θ             ∂g(yθ∗ (z)) ∂g(xθ ) ∂g(xθ ) ∂g(ỹθ (z)) ∂g(xθ ) ˜ , Ex − Ex , Ez + Ez + ⟨∇θ L, dθ ⟩ = Ex ∂θ ∂θ ∂θ ∂θ ∂θ      ∂g(yθ∗ (z)) ∂g(ỹθ (z)) Ez , Ez ∂θ ∂θ          ∂g(yθ∗ (z)) ∂g(ỹθ (z)) ∂g(xθ ) 2 ∂g(xθ ) ˜ ⟨∇θ L, dθ ⟩ = Ex , Ez + Ez + − Ex ∂θ ∂θ ∂θ ∂θ (A.8)      ∂g(yθ∗ (z)) ∂g(ỹθ (z)) Ez , Ez . ∂θ ∂θ ∀z, let ξθ (z) = ỹθ (z) − yθ∗ (z) so ∥ξθ (z)∥ = ϵ. Fix z. Because the 2nd order partial derivatives of g with respect to θ are assumed to exist, by Taylor’s Theorem, ∃t ∈ (0, 1) such that ∂g(yθ∗ (z)) ∂g(ỹθ (z)) = + ∇2θ g(yθ∗ (z) + tξθ (z))ξθ (z). ∂θ ∂θ Taking expectation with respect to z on both sides, by linearity of expectation,       ∂g(yθ∗ (z)) ∂g(ỹθ (z)) (A.10) Ez = Ez + Ez ∇2θ g(yθ∗ (z) + tξθ (z))ξθ (z) . ∂θ ∂θ

(A.9)

Therefore,      ∂g(yθ∗ (z)) ∂g(xθ ) ∂g(ỹθ (z)) Ex , Ez + Ez = ∂θ ∂θ ∂θ         2 ∗  ∂g(yθ∗ (z)) ∂g(yθ∗ (z)) ∂g(xθ ) Ex , Ez + Ez ∇θ g(yθ (z) + tξθ (z))ξθ (z) + Ez ∂θ ∂θ ∂θ



(A.11)             ∂g(yθ∗ (z)) ∂g(yθ∗ (z)) ∂g(xθ ) ∂g(ỹθ (z)) ∂g(xθ ) , Ez + Ez = 2 Ex , Ez + Ex ∂θ ∂θ ∂θ ∂θ ∂θ      2 ∗  ∂g(xθ ) , Ez ∇θ g(yθ (z) + tξθ (z))ξθ (z) Ex ∂θ and (A.12)            2 ∗  ∂g(yθ∗ (z)) ∂g(yθ∗ (z)) ∂g(yθ∗ (z)) ∂g(ỹθ (z)) Ez , Ez = Ez + Ez ∇θ g(yθ (z) + tξθ (z))ξθ (z) , Ez ∂θ ∂θ ∂θ ∂θ   2    ∗  2 ∗  ∂g(yθ (z)) ∂g(yθ∗ (z)) + Ez ∇θ g(yθ (z) + tξθ (z))ξθ (z) , Ez (A.13) = Ez . ∂θ ∂θ

FIXED-POINT NEURAL OPTIMAL TRANSPORT

25

Combining (A.8), (A.11), and (A.13) gives          ∂g(yθ∗ (z)) ∂g(yθ∗ (z)) 2 ∂g(xθ ) 2 ∂g(xθ ) ˜ , Ez + Ez ⟨∇θ L, dθ ⟩ = Ex − 2 Ex − ∂θ ∂θ ∂θ ∂θ        ∂g(yθ∗ (z)) ∂g(xθ ) Ez ∇2θ g(yθ∗ (z) + tξθ (z))ξθ (z) , Ex + Ez ∂θ ∂θ (A.14)

      2 ∗  ∂g(yθ∗ (z)) ∂g(xθ ) 2 ˜ + Ez . ⟨∇θ L, dθ ⟩ = ∥∇θ L∥ − Ez ∇θ g(yθ (z) + tξθ (z))ξθ (z) , Ex ∂θ ∂θ

By Cauchy-Schwarz and the triangle inequality       2 ∗  ∂g(yθ∗ (z)) ∂g(xθ ) Ez ∇θ g(yθ (z) + tξθ (z))ξθ (z) , Ex + Ez ≤ ∂θ ∂θ       ∂g(yθ∗ (z)) ∂g(xθ ) Ez ∇2θ g(yθ∗ (z) + tξθ (z))ξθ (z) Ex + Ez ∂θ ∂θ        2 ∗ ∂g(yθ∗ (z)) ∂g(xθ ) Ez ∇θ g(yθ (z) + tξθ (z))ξθ (z) , Ex + Ez ≤ ∂θ ∂θ         2 ∗ ∂g(yθ∗ (z)) ∂g(xθ ) Ex Ez ∇θ g(yθ (z) + tξθ (z))ξθ (z) + Ez ∂θ ∂θ      ∗   2 ∗ ∂g(yθ (z)) ∂g(xθ ) + Ez ≤ Ez ∇θ g(yθ (z) + tξθ (z))ξθ (z) , Ex ∂θ ∂θ p   p  Ez ∇2θ g(yθ∗ (z) + tξθ (z)) ∥ξθ (z)∥ Mx + M z  (A.15)

      2 ∗ ∂g(yθ∗ (z)) ∂g(xθ ) + Ez ≤ Ez ∇θ g(yθ (z) + tξθ (z))ξθ (z) , Ex ∂θ ∂θ   p p  sup ∇2θ g(yθ∗ (z) + tξθ (z)) sup (∥ξθ (z)∥) M x + Mz . z

z

 Because ∂g(·) ∇2θ g(yθ∗ (z) + tξθ (z)) ≤ Lθ . ∂θ is assumed to be Lθ -Lipschitz, it follows that supz Then, because for any z, ∥ξθ (z)∥ = ϵ (A.15) becomes (A.16)      p p    2 ∗ ∂g(yθ∗ (z)) ∂g(xθ ) Ez ∇θ g(yθ (z) + tξθ (z))ξθ (z) , Ex + Ez ≤ Lθ ϵ M x + Mz . ∂θ ∂θ Thus, because ϵ ≤ L

θ

V ∥∇θ L∥2 √ √

( Mx + Mz )

for some 0 < V < 1,

(A.17)       2 ∗  ∂g(yθ∗ (z)) ∂g(xθ ) Ez ∇θ g(yθ (z) + tξθ (z))ξθ (z) , Ex + Ez ≤ Lθ ∂θ ∂θ (A.18)

V ∥∇ L∥2 √ θ √  Lθ Mx + Mz

≤ V ∥∇θ L∥2 .

!

p p  Mx + Mz

26

Y. PARK, E. GELPHMAN, S. OSHER AND S. WU FUNG

Hence, combining (A.14) and (A.18), (A.19)

D

E ∇θ L, d˜θ ≥ ∥∇θ L∥2 − V ∥∇θ L∥2

(A.20)

≥ (1 − V )∥∇θ L∥2

(A.21)

≥ U ∥∇θ L∥2 .

A.6. Proof of Lemma 3.8. Proof. First, it needs to be shown that ∇θ L is Lipschitz with respect to θ. Using the fact p that ∂g ∂θ is Lθ -Lipschitz with respect to θ, for any θj+1 , θj ∈ R " " # #!    ∂g(yθ∗j+1 (z)) ∂g(yθ∗j (z)) ∂g(xθj ) ∂g(xθj+1 ) − Ez − Ex − Ez ∥∇θ L(θj+1 ) − ∇θ L(θj )∥ = Ex ∂θ ∂θ ∂θ ∂θ " #   ∂g(yθ∗j (z)) ∂g(yθ∗j+1 (z)) ∂g(xθj+1 ) ∂g(xθj ) = Ex − + Ez − ∂θ ∂θ ∂θ ∂θ # "   ∂g(yθ∗j (z)) ∂g(yθ∗j+1 (z)) ∂g(xθj+1 ) ∂g(xθj ) − − ≤ Ex + Ez ∂θ ∂θ ∂θ ∂θ 

≤ Ex [Lθ ∥θj+1 − θj ∥] + Ez [Lθ ∥θj+1 − θj ∥] ≤ Lθ ∥θj+1 − θj ∥ + Lθ ∥θj+1 − θj ∥ ≤ 2Lθ ∥θj+1 − θj ∥. Therefore, ∇θ L is 2Lθ -Lipschitz with respect to θ. Because the gradient with respect to θ of L, the second order Taylor series expansion of L centered at θ = θj satisfies, using θj+1 = θj − αj d˜ϵj (θj ) 1 L(θj+1 ) ≤ L(θj ) + ∇θ L(θj )⊤ (θj+1 − θj ) + (2Lθ )∥θj+1 − θj ∥2 2 L(θj+1 ) ≤ L(θj ) − αj ∇θ L(θj )⊤ d˜ξj (θj ) + αj2 Lθ ∥d˜ξj (θj )∥2 . Taking expectation of both sides above with respect to ξj yields (A.22)

i h i h Eξj [L(θj+1 )] ≤ Eξj [L(θj )] − αj Eξj ∇θ L(θj )⊤ d˜ξj (θj ) + αj2 Lθ Eξj ∥d˜ξj (θj )∥2 .

Because θj depends only on ξj , ξj−1 , ..., ξ0 , taking Eξj [·] only affects the LHS of (A.22). Thus, (A.22) becomes (A.23)

Eξj [L(θj+1 )] ≤ L(θj ) − αj ∇θ L(θj )⊤ d˜ξj (θj ) + αj2 Lθ ∥d˜ξj (θj )∥2 .

Applying the result of Theorems 3.7 and 3.6 to the above inequality, (A.24)

Eξj [L(θj+1 )] − L(θj ) ≤ −αj U ∥∇θ L(θj )∥2 + αj2 Lθ M̃ .

FIXED-POINT NEURAL OPTIMAL TRANSPORT

27

It can be derived that the RHS above is ≤ 0 when αj satisfies 0 < αj ≤ ULMM̃L : θ

−αj U ∥∇θ L(θj )∥2 + αj2 Lθ M̃ ≤ 0 αj2 Lθ M̃ ≤ αj U ∥∇θ L(θj )∥2 U ∥∇θ L(θj )∥2 Lθ M̃ U ML α≤ . Lθ M̃

αj ≤

Proof of Theorem 3.9: Suppose the sequence of learning rates {αj }∞ j=0 is monotonically P∞ P∞ 2 U ML decreasing and satisfies j=0 αj = ∞, j=0 αj < ∞, and 0 < α0 ≤ Lθ M̃ . Let AK = PK−1 j=0 αj . Then, under the assumptions of Lemma 3.8 the SGD iteration (3.24) satisfies   K X 1 lim E  αj ∥∇θ L(θj )∥2  = 0. K→∞ AK j=0

Proof. Taking the total expectation of (A.24), we have E[L(θj+1 )] − E[L(θj )] ≤ −αj U E[∥∇θ L(θj )∥2 ] + αj2 Lθ M̃ .

(A.25)

Setting j = 0 in (A.25), we have (A.26)

E[L(θ1 )] − E[L(θ0 )] ≤ −αj U E[∥∇θ L(θ0 )∥2 ] + αj2 Lθ M̃

(A.27)

E[L(θ1 )] ≤ E[L(θ0 )] − αj U E[∥∇θ L(θ0 )∥2 ] + αj2 Lθ M̃ .

Setting j = 1 in (A.25) and applying (A.27) E[L(θ2 )] − E[L(θ1 )] ≤ −αj U E[∥∇θ L(θ1 )∥2 ] + αj2 Lθ M̃ E[L(θ2 )] ≤ E[L(θ1 )] − αj U E[∥∇θ L(θ1 )∥2 ] + αj2 Lθ M̃ E[L(θ2 )] ≤ E[L(θ0 )] − U

1 X

αj E[∥∇θ L(θj )∥2 ] + Lθ M̃

1 X

αj2

j=0

j=0

E[L(θ2 )] − E[L(θ0 )] ≤ −U

1 X

αj E[∥∇θ L(θj )∥2 ] + Lθ M̃

j=0

1 X

αj2

j=0

.. . E[L(θK )] − E[L(θ0 )] ≤ −U

K−1 X j=0

2

αj E[∥∇θ Jx (θj )∥ ] + Lθ M̃

K−1 X j=0

αj2 .

28

Y. PARK, E. GELPHMAN, S. OSHER AND S. WU FUNG

Since L(θ) is bounded from below by Linf , algebraically rearranging the last line above yields (A.28)

K−1 X

K−1

αj E[∥∇θ L(θj )∥2 ] ≤

j=0

(A.29)

1 AK

E[L(θ0 )] − Linf Lθ M̃ X 2 + αj U U j=0

K−1 X

K−1

αj E[∥∇θ L(θj )∥2 ] ≤

j=0

E[L(θ0 )] − Linf Lθ M̃ X 2 + αj . U AK U AK j=0

Using linearity of expectation, (A.29) becomes   K−1 K−1 X E[L(θ0 )] − Linf Lθ M̃ X 2 1 + αj ∥∇θ L(θj )∥2  ≤ αj . (A.30) E AK U AK U AK j=0

j=0

P∞

Hence, since limK→∞ AK = ∞, and j=0 αj2 converges, we have   K X 1 lim E  αj ∥∇θ Jx (θj )∥2  K→∞ AK j=0 # " P 2 E[L(θ0 )] − Linf + Lθ M̃ K−1 j=0 αj ≤ lim K→∞ U AK = 0. A.7. Proof of Theorem 3.10. Proof. Suppose, for contradiction that, for some a > 0 h i (A.31) lim inf E ∥∇θ L(θj )∥2 = a. j→∞

Let K ∈ N and AK = 

PK−1

k=0 αk . Then, using linearity of expectation

1 E AK

K−1 X

 2

αk ∥∇θ L(θj )∥

j=0

K−1 h i 1 X = αj E ∥∇θ L(θj )∥2 . AK j=0

Taking the liminf as K → ∞ on both sides of the above equation yields a contradiction, as the LHS converges to 0 but the RHS diverges. Hence, by contradiction the result follows. A.8. Proof of Corollary 3.11. Proof. This proof is similar to that of Theorem 4.11 in [11], but we will include it here to be complete. Let ϵ > 0 and let E[·] represent total expectation. By Markov’s inequality and the law of total expectation, also known as the tower property, (A.32) (A.33)

P (∥∇θ L(θj(K) )∥ ≥ ϵ) = P (∥∇θ L(θj(K) )∥2 ≥ ϵ2 ) 1 ≤ 2 E[Ej(K) [∇θ L(θj(K) )]]. ϵ

FIXED-POINT NEURAL OPTIMAL TRANSPORT

29

hP i K−1 2 < ∞. Therefore, we By the proof of Theorem 3.9 we have limK→∞ E α ∥∇ L(θ )∥ j j θ j=0   must have limj→∞ E αj ∥∇θ L(θj )∥2 = 0. Thus, by (A.33), (A.34)

lim P (∥∇θ L(θj(K) )∥ ≥ ϵ) ≤ lim

1

K→∞ ϵ2

K→∞

E[Ej(K) [∇θ L(θj(K) )]] = 0.

Since the choice of ϵ > 0 was arbitrary, this holds ∀ϵ > 0, proving convergence in probability. Appendix B. Implementation Details. B.1. Datasets. B.1.1. Synthetic Gaussians. For the synthetic high-dimensional Gaussian experiments in Section 4.1, we use the data generation procedure of [56], which follows the standard protocol of [42] for constructing Gaussian pairs with controlled covariance spectra. In this setup, each covariance matrix is defined as Σ = QΛQ⊤ , where Q is an orthogonal matrix obtained via QR decomposition of a Gaussian random matrix, and Λ contains eigenvalues sampled from a log-uniform distribution, λi ∼ exp(U(a, b)). This yields well-conditioned but heterogeneous Gaussian covariances. Given two Gaussian distributions N (0, Σ0 ) and N (0, Σ1 ), the optimal transport map is linear, T (x) = Γx, where −1/2

Γ = Σ0



1/2

 1/2 1/2

Σ0 Σ1 Σ0

−1/2

Σ0

.

Samples from N (0, Σ0 ) and N (0, Σ1 ) are drawn following the same procedure as in [56], with n = 105 samples for both training and evaluation. Independent test samples are used to evaluate transport accuracy, where the ground-truth map T (x) = Γx is applied to source samples, and Γ−1 is used for inverse-consistency checks. B.1.2. UCI Physics Data. The UCI physics dataset is used following the preprocessing protocol of [53], and we adopt the same train/validation/test split and normalization procedure as in the original work. GAS Dataset. For the GAS dataset, the raw data is loaded from ethylene CO.csv, and the variables Meth, Eth, and Time are removed. Highly correlated features are iteratively removed by discarding variables whose pairwise correlation exceeds 0.98. After this cleaning step, each feature is standardized using z-score normalization based on the training data statistics. The dataset is then split into training, validation, and test sets using a 80/10/10 split. POWER Dataset. For the POWER dataset, we use the preprocessed data from data.npy. Two features are removed as in the original preprocessing pipeline, and additional synthetic noise is added to several variables following the procedure of prior work. The dataset is then randomly shuffled and split into training, validation, and test sets using a 80/10/10 ratio.

30

Y. PARK, E. GELPHMAN, S. OSHER AND S. WU FUNG

HEPMASS Dataset. For the HEPMASS dataset, we follow the preprocessing pipeline of prior work on tabular density estimation. We use the split provided in 1000 train.csv and 1000 test.csv, and restrict the data to the positive class only by removing background noise samples (class label 0). In addition, the label column and one redundant feature column in the test set are removed. We consider three normalization schemes: static normalization, min-max scaling, and SVD-based whitening. In the experiments reported in the main paper, we use the min-max scaling variant (denoted as scale), where each feature is normalized as x′ =

x−µ , s

µ=

max(x) + min(x) , 2

s=

max(x) − min(x) . 2

Statistics are computed using the training set only. The resulting transformation maps each feature approximately to [−1, 1]. We further remove features with highly repetitive values following the procedure in the original preprocessing code. MINIBOONE Dataset. For the MINIBOONE dataset, we use the preprocessed version provided as a NumPy array. The dataset is split into training, validation, and test sets using an 80/10/10 ratio. We apply min-max scaling (denoted as scale) using statistics computed on the combined training and validation set: x′ =

x−µ , s

µ=

max(x) + min(x) , 2

s=

max(x) − min(x) . 2

This normalization is applied independently to each feature, resulting in an approximate [−1, 1] range per dimension. BSDS300 Dataset. The BSDS300 dataset consists of natural image patches extracted from the Berkeley Segmentation Dataset. We use the standard train/validation/test split provided in the preprocessed HDF5 file BSDS300.hdf5, following prior work on image-based optimal transport and density modeling. Each sample corresponds to a vectorized image patch, and no additional feature engineering is applied. We directly load the precomputed splits for training, validation, and testing. Since the data consists of image patches, we additionally record the corresponding spatial resolution, given by h√ i √ image size = d + 1, d + 1 , where d denotes the dimensionality of the vectorized patch. No further normalization beyond the provided preprocessing is applied. B.1.3. Synthetic Class-Conditional OT Datasets. We construct three synthetic benchmarks for class-conditional optimal transport using mixtures of two-dimensional Gaussian clusters arranged on a circular manifold. All datasets are generated on a ring of radius r = 0.9 with isotropic Gaussian noise ϵ = 0.07. Each class corresponds to a subset of Gaussian components placed at fixed angular positions. Crossed Ring Gaussians. In the Crossed Ring Gaussian dataset, each distribution consists of two Gaussian modes placed at distinct angular positions on the unit circle. Specifically, the

FIXED-POINT NEURAL OPTIMAL TRANSPORT

31

source distribution is defined by two clusters centered at angles 0 and π/2, while the target distribution is defined by clusters at angles π/4 and 3π/4. This creates a crossed transport structure between classes. The total number of samples is N = 50,000, with equal allocation per mode. Horizontal Swapped Gaussians. We construct an additional synthetic benchmark consisting of two one-dimensional Gaussian clusters embedded in R2 . The dataset is designed to evaluate class-conditional optimal transport under near-degenerate geometric structure. The source distribution consists of two Gaussian components centered at (−0.23, 0) and (0.23, 0), while the target distribution consists of two components centered at (0.7, 0) and (−0.7, 0), respectively. This induces a horizontal swapping structure along the x-axis. Formally, each component is generated as x = µ + ϵ · N (0, I2 ), where µ ∈ R2 denotes the cluster center and ϵ = 0.05 controls the noise level. The total number of samples is N = 50,000, equally split across components. This dataset tests the ability of class-conditional OT methods to recover correct pairwise matching in a low-dimensional, highly structured setting. Four-Mode Ring Gaussians. The Four-Mode Ring Gaussian dataset consists of four Gaussian components per distribution placed uniformly on the unit circle. The source distribution uses angles {0, π/2, π, 3π/2}, while the target distribution uses {π/4, 3π/4, 5π/4, 7π/4}. This results in a dense multimodal matching problem with N = 5,000 total samples. For all datasets, each mode is sampled as x = r · (cos θ, sin θ) + ϵ · N (0, I2 ), where r = 0.9 and ϵ = 0.07. Samples are concatenated across modes to form the full empirical distributions. The resulting datasets are used to evaluate class-conditional optimal transport under structured multimodal alignments. B.1.4. Image Data. Normalization is performed using a min-max scaling transformation computed from the combined training and validation set: x′ =

x−µ , s

µ=

max(x) + min(x) , 2

s=

max(x) − min(x) . 2

This maps each feature approximately into the range [−1, 1]. All images are resized to 32 × 32 pixels and normalized to [0, 1]. For training the OT map, we construct a class-paired dataset by randomly drawing matched (xsrc , xtgt , k) triples where both images belong to the same class k ∈ {0, . . . , 9}, drawn uniformly at random from FashionMNIST (source) and MNIST (target) training sets (60,000 images each). B.2. Kantorovich Potential Network. For the high-dimensional Gaussian experiments in Section 4.1, we employ the DenseICNN architecture, a fully connected neural network augmented with input-quadratic skip connections, in order to ensure a fair comparison with the baseline models of [42]. When implementing our method and the NCF baseline, we remove the nonlinearity constraints that are typically imposed to enforce convexity. Following [42],

32

Y. PARK, E. GELPHMAN, S. OSHER AND S. WU FUNG

we adopt the network architecture DenseICNN[1; max(2d,64), max(2d,64), max(d,32)] for a ddimensional problem. The model is optimized using the Adam optimizer with a fixed learning rate of 10−4 , independent of the input dimension. The same network architecture is also used for the 2D class-conditional optimal transport experiments in Section 4.2, Section 4.3, and ablation studies in Section 4.5. For the image translation experiments in Section 4.4, the potential gθ (x, k) is parameterized by a scalar-valued network taking as input the concatenation [x; ck ] ∈ R25 , where ck ∈ {0, 1}10 is the class one-hot vector. The network decomposes as (B.1)

gθ (x, k) = 12 x⊤ (A⊤A) x + b⊤ x + c + w⊤ Φ(x, k), | {z } | {z } quadratic

nonlinear

where A ∈ R25×25 is Xavier-initialized, enforcing a positive semi-definite quadratic component. The nonlinear branch Φ is a residual network: a linear opening layer R25 → R512 , followed by four hidden layers of width 512 with residual connections weighted by h = 14 , and a linear readout. The activation function is softplus throughout. B.3. Implicit Fixed-Point Solver. The proximal fixed point y ⋆ (z, k) = proxgθ (z, k) is computed by iterating   (B.2) y (t+1) = y (t) − α y (t) − z + ∇y gθ (y (t) , k) ,  with step size α = min 0.1/L̂, 10−2 , where L̂ is the spectral norm of the Hessian of gθ estimated via 20 steps of power iteration. Iterations run until ∥∇y gθ (y ⋆ , k) + y ⋆ − z∥∞ < 10−3 or for at most 104 steps. The previous batch’s fixed point is used as a warm start. B.4. OT Map Training. The OT map is trained to minimize (B.3)

2

\ (T (x), z) , L(θ) = Eµk [gθ (x, k)] − Eνk [gθ (y ⋆ (z, k), k)] + λMMD MMD 2

\ is the unbiased multiwhere T (x) = x + ∇x gθ (x, k) is the forward transport map and MMD 2 scale RBF kernel estimator. We use five kernel bandwidths, σ ∈ {0.25, 0.5, 1.0, 2.0, 4.0} × 2 σmedian , which are only applied in the image-based experiments in Section 4.4. Consequently, we set λMMD = 0 for the experiments in Sections 4.1, 4.2, 4.3, and 4.5. B.5. Conditional VAE for Image Task. Architecture. Both the FashionMNIST and MNIST VAEs share identical architecture. The encoder consists of four strided convolutional layers (kernel 4 × 4, stride 2, padding 1), with channel widths 1 → 64 → 128 → 256 → 512, each followed by BatchNorm and LeakyReLU(0.2). The flattened 512 × 2 × 2 feature map is concatenated with the class onehot vector and projected to the mean µ and log-variance log σ 2 of the latent distribution via linear layers, yielding a latent dimension of d = 15. The decoder takes the concatenation of the latent code and class one-hot as input, projects it to a 512 × 2 × 2 feature map, and upsamples through four transposed convolutional layers (kernel 4 × 4, stride 2) with channel widths 512 → 256 → 128 → 64 → 1. Each upsampling block is followed by GroupNorm(16) and three residual blocks (two 3 × 3 convolutions with

FIXED-POINT NEURAL OPTIMAL TRANSPORT

Hyperparameter

33

Value

Epochs Batch size MMD subsample size λMMD Fixed-point tolerance τ Max fixed-point iterations Optimizer Learning rate LR schedule Minimum LR Floating-point precision

20 64 5000 10−3 10−3 104 Adam 10−4 ReduceLROnPlateau (factor 0.5, patience 5) 10−7 float64

Table 7: OT map training hyperparameters.

GroupNorm and a skip connection). The final output is passed through a Sigmoid activation. GroupNorm is used in the decoder (instead of BatchNorm) to avoid train/eval distribution mismatch. Hyperparameter

FashionMNIST VAE

MNIST VAE

Latent dimension d βmax KL warmup epochs Epochs Batch size Optimizer Learning rate LR schedule Minimum LR

15 15 1.0 0.1 20 (linear ramp) 0 1000 1000 128 128 Adam Adam 10−3 10−3 ReduceLROnPlateau (factor 0.5, patience 100) 10−7

Table 8: VAE training hyperparameters.

VAE Training Hyperparameters. Both VAEs are frozen after pre-training; their parameters are not updated during OT map training. REFERENCES [1] A. Almahairi, S. Rajeshwar, A. Sordoni, P. Bachman, and A. Courville, Augmented cyclegan: Learning many-to-many mappings from unpaired data, in International conference on machine learning, PMLR, 2018, pp. 195–204. [2] D. Alvarez-Melis and N. Fusi, Dataset dynamics via gradient flows in probability space, in International conference on machine learning, PMLR, 2021, pp. 219–230.

34

Y. PARK, E. GELPHMAN, S. OSHER AND S. WU FUNG

[3] D. Alvarez-Melis, Y. Schiff, and Y. Mroueh, Optimizing functionals on the space of probabilities with input convex neural networks, arXiv preprint arXiv:2106.00774, (2021). [4] B. Amos, L. Xu, and J. Z. Kolter, Input convex neural networks, in International conference on machine learning, PMLR, 2017, pp. 146–155. [5] M. Arjovsky, S. Chintala, and L. Bottou, Wasserstein generative adversarial networks, in International conference on machine learning, PMLR, 2017, pp. 214–223. [6] A. Asadulaev, A. Korotin, V. Egiazarian, P. Mokrov, and E. Burnaev, Neural optimal transport with general cost functionals, International conference on learning representations, (2024). [7] J. Backhoff-Veraguas, M. Beiglböck, and G. Pammer, Existence, duality, and cyclical monotonicity for weak transport costs, Calculus of Variations and Partial Differential Equations, 58 (2019), p. 203. [8] Y. Balaji, R. Chellappa, and S. Feizi, Robust optimal transport with applications in generative modeling and domain adaptation, Advances in Neural Information Processing Systems, 33 (2020), pp. 12934–12944. [9] M. Barthélemy and A. Flammini, Optimal traffic networks, Journal of Statistical Mechanics: Theory and Experiment, 2006 (2006), p. L07002. [10] J.-D. Benamou and Y. Brenier, A computational fluid mechanics solution to the monge-kantorovich mass transfer problem, Numerische Mathematik, 84 (2000), pp. 375–393. [11] L. Bottou, F. E. Curtis, and J. Nocedal, Optimization methods for large-scale machine learning, SIAM review, 60 (2018), pp. 223–311. [12] C. Bunne, S. G. Stark, G. Gut, J. S. Del Castillo, M. Levesque, K.-V. Lehmann, L. Pelkmans, A. Krause, and G. Rätsch, Learning single-cell perturbation responses using neural optimal transport, Nature methods, 20 (2023), pp. 1759–1768. [13] G. Carlier, C. Jimenez, and F. Santambrogio, Optimal transportation with traffic congestion and wardrop equilibria, SIAM Journal on Control and Optimization, 47 (2008), pp. 1330–1350. [14] J. Choi, Y. Chen, and J. Choi, Improving neural optimal transport via displacement interpolation, arXiv preprint arXiv:2410.03783, (2024). [15] N. Courty, R. Flamary, A. Habrard, and A. Rakotomamonjy, Joint distribution optimal transportation for domain adaptation, Advances in neural information processing systems, 30 (2017). [16] N. Courty, R. Flamary, D. Tuia, and A. Rakotomamonjy, Optimal transport for domain adaptation, IEEE transactions on pattern analysis and machine intelligence, 39 (2016), pp. 1853–1865. [17] M. Cuturi, Sinkhorn distances: Lightspeed computation of optimal transport, Advances in neural information processing systems, 26 (2013). [18] B. B. Damodaran, B. Kellenberger, R. Flamary, D. Tuia, and N. Courty, Deepjdot: Deep joint distribution optimal transport for unsupervised domain adaptation, in Proceedings of the European conference on computer vision (ECCV), 2018, pp. 447–463. [19] M. Daniels, T. Maunu, and P. Hand, Score-based generative neural networks for large-scale optimal transport, Advances in neural information processing systems, 34 (2021), pp. 12955–12965. [20] B. Danila, Y. Yu, J. A. Marsh, and K. E. Bassler, Optimal transport on complex networks, Physical Review E—Statistical, Nonlinear, and Soft Matter Physics, 74 (2006), p. 046106. [21] J. Darbon and G. P. Langlois, On bayesian posterior mean estimators in imaging sciences and hamilton–jacobi partial differential equations, Journal of Mathematical Imaging and Vision, 63 (2021), pp. 821–854. [22] J. Darbon, G. P. Langlois, and T. Meng, Connecting hamilton-jacobi partial differential equations with maximum a posteriori and posterior mean estimators for some non-convex priors, in Handbook of Mathematical Models and Algorithms in Computer Vision and Imaging: Mathematical Imaging and Vision, Springer, 2021, pp. 1–25. [23] N. Di, E. C. Chi, and S. W. Fung, Operator splitting with hamilton-jacobi-based proximals, arXiv preprint arXiv:2601.22370, (2026). [24] L. El Ghaoui, F. Gu, B. Travacca, A. Askari, and A. Tsai, Implicit deep learning, SIAM Journal on Mathematics of Data Science, 3 (2021), pp. 930–958. [25] J. Fan, S. Liu, S. Ma, H.-M. Zhou, and Y. Chen, Neural monge map estimation and its applications, Transactions on Machine Learning Research, (2023). [26] J. Fan, Q. Zhang, A. Taghvaei, and Y. Chen, Variational wasserstein gradient flow, International

FIXED-POINT NEURAL OPTIMAL TRANSPORT

35

Conference on Machine Learning, (2022). [27] A. Figalli and F. Glaudo, An Invitation to Optimal Transport, Wasserstein Distances, and Gradient Flows, European Mathematical Society, 2021. [28] C. Finlay, J.-H. Jacobsen, L. Nurbekyan, and A. Oberman, How to train your neural ode: the world of jacobian and kinetic regularization, in International conference on machine learning, PMLR, 2020, pp. 3154–3164. [29] S. W. Fung, H. Heaton, Q. Li, D. McKenzie, S. Osher, and W. Yin, Jfb: Jacobian-free backpropagation for implicit networks, in Proceedings of the AAAI Conference on Artificial Intelligence, vol. 36, 2022, pp. 6648–6656. [30] E. Gelphman, D. Verma, N. T. Yang, S. Osher, and S. W. Fung, End-to-end training of highdimensional optimal control with implicit hamiltonians via jacobian-free backpropagation, arXiv preprint arXiv:2510.00359, (2025). [31] E. Gelphman, D. Verma, N. T. Yang, S. Osher, and S. W. Fung, On the convergence of jacobian-free backpropagation for optimal control problems with implicit hamiltonians, arXiv preprint arXiv:2602.00921, (2026). [32] A. Genevay, M. Cuturi, G. Peyré, and F. Bach, Stochastic optimization for large-scale optimal transport, Advances in neural information processing systems, 29 (2016). [33] W. Grathwohl, R. T. Chen, J. Bettencourt, I. Sutskever, and D. Duvenaud, Ffjord: Free-form continuous dynamics for scalable reversible generative models, International Conference on Learning Representations, (2019). [34] N. Gushchin, A. Kolesov, A. Korotin, D. P. Vetrov, and E. Burnaev, Entropic neural optimal transport via diffusion processes, Advances in Neural Information Processing Systems, 36 (2023), pp. 75517–75544. [35] K. He, X. Zhang, S. Ren, and J. Sun, Deep residual learning for image recognition, in Proceedings of the IEEE Conference on Computer Vision and Pattern Recognition (CVPR), June 2016. [36] H. Heaton, S. W. Fung, A. T. Lin, S. Osher, and W. Yin, Wasserstein-based projections with applications to inverse problems, SIAM Journal on Mathematics of Data Science, 4 (2022), pp. 581– 603. [37] H. Heaton, S. Wu Fung, and S. Osher, Global solutions to nonconvex problems by evolution of hamilton-jacobi pdes, Communications on Applied Mathematics and Computation, 6 (2024), pp. 790– 810. [38] X. Huang, M.-Y. Liu, S. Belongie, and J. Kautz, Multimodal unsupervised image-to-image translation, in Proceedings of the European conference on computer vision (ECCV), 2018, pp. 172–189. [39] G. Huguet, D. S. Magruder, A. Tong, O. Fasina, M. Kuchroo, G. Wolf, and S. Krishnaswamy, Manifold interpolating optimal-transport flows for trajectory inference, Advances in neural information processing systems, 35 (2022), pp. 29705–29718. [40] M. Kelly, R. Longjohn, and K. Nottingham, The uci machine learning repository. https://archive. ics.uci.edu. [41] A. Korotin, V. Egiazarian, A. Asadulaev, A. Safin, and E. Burnaev, Wasserstein-2 generative networks, arXiv preprint arXiv:1909.13082, (2019). [42] A. Korotin, L. Li, A. Genevay, J. M. Solomon, A. Filippov, and E. Burnaev, Do neural optimal transport solvers work? a continuous wasserstein-2 benchmark, Advances in neural information processing systems, 34 (2021), pp. 14593–14605. [43] A. Korotin, L. Li, J. Solomon, and E. Burnaev, Continuous wasserstein-2 barycenter estimation without minimax optimization, arXiv preprint arXiv:2102.01752, (2021). [44] A. Korotin, D. Selikhanovych, and E. Burnaev, Neural optimal transport, International conference on learning representations, (2023). [45] T. Koshizuka and I. Sato, Neural lagrangian schr\” odinger bridge: Diffusion modeling for population dynamics, arXiv preprint arXiv:2204.04853, (2022). [46] Y. LeCun, The mnist database of handwritten digits, http://yann. lecun. com/exdb/mnist/, (1998). [47] A. T. Lin, S. W. Fung, W. Li, L. Nurbekyan, and S. J. Osher, Alternating the population and control neural networks to solve high-dimensional stochastic mean-field games, Proceedings of the National Academy of Sciences, 118 (2021), p. e2024713118. [48] H. Liu, X. Gu, and D. Samaras, Wasserstein gan with quadratic transport cost, in Proceedings of the

36

Y. PARK, E. GELPHMAN, S. OSHER AND S. WU FUNG

IEEE/CVF international conference on computer vision, 2019, pp. 4832–4841. [49] S. Liu, S. Ma, Y. Chen, H. Zha, and H. Zhou, Learning high dimensional wasserstein geodesics, arXiv preprint arXiv:2102.02992, (2021). [50] S. Liu, S. Osher, and W. Li, A natural primal-dual hybrid gradient method for adversarial neural network training on solving partial differential equations, arXiv preprint arXiv:2411.06278, (2024). [51] G. Lu, Z. Zhou, J. Shen, C. Chen, W. Zhang, and Y. Yu, Large-scale optimal transport via adversarial training with cycle-consistency, arXiv preprint arXiv:2003.06635, (2020). [52] A. Makkuva, A. Taghvaei, S. Oh, and J. Lee, Optimal transport mapping via input convex neural networks, in International Conference on Machine Learning, PMLR, 2020, pp. 6672–6681. [53] D. Onken, S. W. Fung, X. Li, and L. Ruthotto, Ot-flow: Fast and accurate continuous normalizing flows via optimal transport, in Proceedings of the AAAI Conference on Artificial Intelligence, vol. 35, 2021, pp. 9223–9232. [54] S. Osher, H. Heaton, and S. Wu Fung, A hamilton–jacobi-based proximal operator, Proceedings of the National Academy of Sciences, 120 (2023), p. e2220469120. [55] N. Parikh and S. Boyd, Proximal algorithms, Foundations and Trends in optimization, 1 (2014), pp. 127–239. [56] Y. Park, S. Liu, M. Zhou, and S. Osher, Neural hamilton–jacobi characteristic flows for optimal transport, International Conference on Learning Representations, (2026). [57] Y. Park and S. Osher, Neural implicit solution formula for efficiently solving hamilton-jacobi equations, SIAM Journal on Scientific Computing, 47 (2025), pp. C1223–C1263. [58] Y. Park and S. Osher, Scalable fixed-point framework for high-dimensional hamilton-jacobi equations, (2025). [59] P. Pope, C. Zhu, A. Abdelkader, M. Goldblum, and T. Goldstein, The intrinsic dimension of images and its impact on learning, arXiv preprint arXiv:2104.08894, (2021). [60] L. Ruthotto, S. J. Osher, W. Li, L. Nurbekyan, and S. W. Fung, A machine learning framework for solving high-dimensional mean field game and mean field control problems, Proceedings of the National Academy of Sciences, 117 (2020), pp. 9183–9193. [61] M. Sanjabi, J. Ba, M. Razaviyayn, and J. D. Lee, On the convergence and robustness of training gans with regularized optimal transport, Advances in Neural Information Processing Systems, 31 (2018). [62] G. Schiebinger, J. Shu, M. Tabaka, B. Cleary, V. Subramanian, A. Solomon, J. Gould, S. Liu, S. Lin, P. Berube, et al., Optimal-transport analysis of single-cell gene expression identifies developmental trajectories in reprogramming, Cell, 176 (2019), pp. 928–943. [63] V. Seguy, B. B. Damodaran, R. Flamary, N. Courty, A. Rolet, and M. Blondel, Large-scale optimal transport and mapping estimation, arXiv preprint arXiv:1711.02283, (2017). [64] Z. Shen, Z. Wang, A. Ribeiro, and H. Hassani, Sinkhorn natural gradient for generative models, Advances in Neural Information Processing Systems, 33 (2020), pp. 1646–1656. [65] A. Taghvaei and A. Jalali, 2-wasserstein approximation via restricted convex potentials with application to improved training for gans, arXiv preprint arXiv:1902.07197, (2019). [66] R. J. Tibshirani, S. W. Fung, H. Heaton, and S. Osher, Laplace meets moreau: Smooth approximation to infimal convolutions using laplace’s method, Journal of Machine Learning Research, 26 (2025), pp. 1–36. [67] A. Tong, J. Huang, G. Wolf, D. Van Dijk, and S. Krishnaswamy, Trajectorynet: A dynamic optimal transport network for modeling cellular dynamics, in International conference on machine learning, PMLR, 2020, pp. 9526–9536. [68] C. Villani et al., Optimal transport: old and new, vol. 338, Springer, 2008. [69] G. Wang, Y. Jiao, Q. Xu, Y. Wang, and C. Yang, Deep generative learning via schrödinger bridge, in International conference on machine learning, PMLR, 2021, pp. 10794–10804. [70] S. Wu Fung and B. Berkels, A generalization bound for a family of implicit networks, Neurocomputing, 678 (2026), p. 133136. [71] H. Xiao, K. Rasul, and R. Vollgraf, Fashion-mnist: a novel image dataset for benchmarking machine learning algorithms, arXiv preprint arXiv:1708.07747, (2017). [72] Y. Xie, M. Chen, H. Jiang, T. Zhao, and H. Zha, On scalable and efficient computation of large scale optimal transport, in International Conference on Machine Learning, PMLR, 2019, pp. 6882–6892. [73] L. Yang and G. E. Karniadakis, Potential flow generator with l2 optimal transport regularity for

FIXED-POINT NEURAL OPTIMAL TRANSPORT

37

generative models, IEEE Transactions on Neural Networks and Learning Systems, 33 (2020), pp. 528– 538. [74] B. J. Zhang and M. A. Katsoulakis, A mean-field games laboratory for generative modeling, arXiv preprint arXiv:2304.13534, (2023). [75] J.-Y. Zhu, T. Park, P. Isola, and A. A. Efros, Unpaired image-to-image translation using cycleconsistent adversarial networks, in Proceedings of the IEEE international conference on computer vision, 2017, pp. 2223–2232.

Record · ID 175274 · SHA-256 9cd8f3efcf73112f
Retrieved via Conceptio — every document is proof-bundled with source, license, and retrieval metadata.