A Neural JKO Scheme for Hellinger-Kantorovich Gradient Flows via Monge-Growth Pairs
Authors: Geuntaek Seo, Cheolhyeong Kim, Hwijae Son, Hyung Ju Hwang
Organizations: Department of Mathematics, Pohang University of Science and Technology · Samsung Electronics · Department of Mathematics, Konkuk University
We develop a mesh-free neural JKO scheme for advection-reaction-diffusion equations with a gradient-flow structure in the Hellinger-Kantorovich (HK) geometry of unbalanced optimal transport. Each update is parametrized by a spatial map and a mass-changing factor, allowing spatial redistribution and local mass creation or loss to be treated jointly within a single variational step. Their cone action bounds the squared HK distance from above, yielding a sufficient condition for discrete energy dissipation through comparison with the identity pair. Minimizing the pair objective over all admissible pairs recovers the exact JKO minimum when the source and a minimizer have positive densities. We establish existence and mass bounds for JKO minimizers and, under additional assumptions, obtain positivity and regularity together with a discrete Euler-Lagrange equation and a metric-dissipation identity. The self-consistent chemical potential is then nonincreasing along an optimal map. There exist parametric pairs whose endpoint densities and objective values converge to those of an exact JKO minimizer, provided a regular-pair approximation hypothesis holds. Finally, we show that a primal-dual gap controls objective suboptimality and, for Boltzmann entropy, the L1 density error, assuming exact-step regularity, positive-semidefinite interactions, and global dual feasibility. Numerical experiments examine pointwise agreement with the PDE, energy dissipation, and the roles of transport, reaction, and fully implicit interactions.
Figures & tables
Restricted dual value ( ≤HK2 )
Primal cone action ( ≥HK2 )
Cumulative state μk=(Sk)#(rk2μ0)
Density can be evaluated from a fixed reference law. A dual potential or c -transform is required, and the lower bound alone does not certify exact JKO descent.
A feasible pair can certify descent, but cumulative displacement and representation complexity may grow.
Restarted state μk+1=T#(q2μk)
Each update starts from the current measure. Dual optimization is still required, and the lower bound alone does not certify exact JKO descent.
This paper: one pair defines the next measure and an upper bound on its JKO objective. The unrestricted problem is exact; density and sampling information must be propagated between steps.
Table 1: Here (Sk,rk) is the cumulative Monge-growth pair representing the approximate state μk , with rk≥0 . Restarting parametrizes each update from the current state; it does not reset accumulated errors. The one-step descent criterion for the pair endpoint relies on the primal upper bound and comparison with the identity pair.
Experiment
d
pair-training rule
density family
Entropy–potential
1
Gauss–Legendre 257
mixed-activation FCNN
Entropy–potential
2
Gauss–Legendre 4225
mixed-activation FCNN
HK entropy flow
8
importance draws 2048
mixed-activation FCNN
Fokker–Planck type
2
Gauss–Legendre 1089
mixed-activation FCNN
Implicit interaction
2
periodic grid 961
mixed-activation FCNN
Table 2: Positive neural state families and the actual pair-training rules. Importance counts are draws before rejection outside the box.
Experiment
d
rel. L2
maxk rel. L2
Atr/A
reference
minJT
sup∥Dδ∥
Entropy–potential
1
2.006e-05
9.819e-03
0.150
0.145
0.981
0.539
Entropy–potential
2
5.319e-04
0.016
0.105
0.102
0.969
0.644
HK entropy flow
8
8.029e-03
8.029e-03
0.069
0.069
0.995
0.806
Fokker–Planck type
2
3.476e-03
3.476e-03
0.228
0.228
0.995
0.453
Implicit interaction
2
1.312e-03
2.434e-03
0.537
0.536
0.977
0.092
Table 3 : Trajectory diagnostics. Action shares are cone-action-weighted; maxℓκθ,ℓ<1 certifies the two blocks of each accepted map.
Figure 1 : Pointwise one-step PDE consistency under τ refinement. Left and center: Lτ and Fτ on a common scale. Right: Rτstr .
Figure 2 : Small-time pair scaling. Left: pair deviations; center: rescaled-field errors; right: profiles at the smallest τ .
τ
∥T−id∥L2(ρ)
∥q2−1∥L2(ρ)
velocity error
reaction error
minJT
0.020
0.014
0.050
0.234
0.146
0.958
0.010
8.169e-03
0.027
0.133
0.081
0.975
5.000e-03
4.370e-03
0.014
0.071
0.043
0.987
2.500e-03
2.265e-03
7.014e-03
0.037
0.022
0.993
1.250e-03
1.153e-03
3.542e-03
0.019
0.011
0.996
Table 4 : Small-time pair diagnostics on [−1,1] . Velocity and reaction errors are relative L2(ρ0dx) errors of (T−id)/τ and (q2−1)/τ against −g0′ and −4g0 .
Figure 3 : One-dimensional relaxation. Step labels count attempts; finite-volume references and diagnostics use the corresponding accepted physical time.
Figure 4 : Full pair versus pure-growth and transport-only restrictions.
Configuration
F
energy monotone?
rel. L2
Atr/A
reference
common t
reached t
full pair (Tθ,qθ)
-0.395
yes
1.308e-03
0.150
0.145
0.800
0.800
pure growth, Tθ=id
-0.395
yes
0.019
0
0.109
0.800
0.800
transport only, qθ=1
1.237
yes
3.726
1.000
0.316
0.800
0.800
Table 5 : Pair-restriction diagnostics at a common attained physical time. Transport shares use the prefix ending at that time. The energy-monotonicity flag uses Ek+1−Ek≤10−7(1+∣Ek∣) , a tighter tolerance than the acceptance guard.
Figure 5 : Two-dimensional relaxation. Top: reference, neural, and signed error; middle: axial sections; bottom: trajectory diagnostics.
Figure 6 : Eight-dimensional Gaussian entropy flow. Top: final coordinate sections; bottom: mass, energy, and error to the whole-space Gaussian reference.
Figure 7 : Two-dimensional confinement test: density and mean-path comparisons with mass, energy, density error, and mean error.
Figure 8 : Interaction benchmark: periodic reference, fully implicit update, and frozen-field ablation at matched times.
Figure 9 : Interaction diagnostics: mass, common fully implicit energy, trajectory error, transport share, and guard margin.
interaction treatment
rel. L2
mass
true energy
transport share
residual (neural/ref.)
common t
reached t
null
fallback
fully implicit
1.312e-03
0.844
-0.940
0.050
0.021/4.787e-03
0.300
0.300
0
0
frozen field
1.133e-03
0.844
-0.940
0.051
0.021/4.787e-03
0.300
0.300
0
0
Table 6 : Implicit/frozen comparison at the common attained physical time. Density errors use the bilinearly interpolated 642 reference; residuals use the same full-periodic operator for both trajectories on the refined 1282 grid. Because the reference is advanced using the same spatial operator, its residual primarily reflects temporal discretization. Null and fallback counts cover all attempted steps; the transport share is that of the last update reaching the common time.
Figure 10 : Finite-grid certificate across network widths, with depth fixed. Left: primal and dual values; center: gap, observed error, and certified radius; right: directly evaluated uθ and the numerical grid reference.
width
primal upper
dual lower
gap Rh
observed L1
certified radius
metric gap /2τ
feasibility violation
8
-0.627
-0.627
9.783e-05
8.534e-03
0.022
1.261e-09
-1.000e-10
16
-0.627
-0.627
9.330e-05
8.159e-03
0.022
1.378e-10
-1.000e-10
32
-0.627
-0.627
9.346e-05
8.371e-03
0.022
5.154e-10
-1.000e-10
64
-0.627
-0.627
9.360e-05
8.370e-03
0.022
1.779e-10
-1.000e-10
128
-0.627
-0.627
9.208e-05
8.485e-03
0.022
2.226e-10
-1.000e-10
Table 7 : Finite-grid certificate. Feasibility is checked over all grid pairs. The scaled metric gap measures coupling-solve accuracy with the candidate density fixed.
Appendix figures & tables4 assets
Supplementary material from the paper’s appendix.
Appendix
Experiment
d
domain
τ
steps
requested T=Nτ
widths (T,q,ρ)
base N
Small- τ pair
1
[−1,1]
0.00125–0.02
1
–
32,64,–
2400
Pointwise PDE diagnostic
1
[−1,1]
0.000625–0.02
1
–
64,64,64
3000
Finite-grid certificate
1
[−0.6,0.6]
0.020
1
0.020
varied,64,–
2400
Entropy–potential
1
[−1,1]
2.500e-03
1000
2.500
64,64,64
250
Entropy–potential
2
[−1,1]2
2.500e-03
480
1.200
64,64,64
250
HK entropy
8
[−3.5,3.5]8
6.250e-04
240
0.150
64,64,64
250
Appendix
Table 8 : Executed settings. N excludes conditioned and growth-polishing stages. Certificate map widths are 8,16,32,64,128 ; a missing density width denotes a raw pushforward.
Experiment
τ
attained T
rel. L2
rel. mass error
flow residual
reference residual
Entropy 8 d
1.250e-03
0.150
0.025
0.022
–
–
Entropy 8 d
6.250e-04
0.150
8.043e-03
0.012
–
–
Implicit interaction
2.500e-03
0.300
5.212e-04
–
0.021
4.787e-03
Frozen interaction
2.500e-03
0.300
4.211e-04
–
0.021
4.787e-03
Implicit interaction
5.000e-03
0.300
1.051e-03
–
0.021
9.533e-03
Frozen interaction
5.000e-03
0.300
9.422e-04
–
0.020
9.533e-03
Appendix
Table 9 : Fixed-horizon time refinement. Interaction errors and residuals use the refined reference grid and the same periodic operator for both trajectories.
Case
steps
fit rel. L2
mass error
pair evals./step
fit evals./step
Relaxation 1D
1000
1.72e-05
5.29e-06
381.5
2.1
Relaxation 2D
480
4.78e-05
7.79e-06
693.3
2.8
Entropy 8D
240
1.92e-04
3.91e-05
499.5
5.0
FP2D
960
1.89e-04
1.22e-04
635.9
7.2
Implicit interaction
120
1.26e-05
7.12e-07
800.0
2.0
Frozen interaction
120
1.25e-05
5.96e-07
823.7
2.0
Appendix
Table 10: Propagation and work diagnostics. Fit and mass errors are maxima over steps; work columns are mean objective-closure evaluations per step, including pair stages and density fitting. Pair and fit closures have different costs.
Treatment
1282
2562
5122
Implicit
0.02069
0.02370
0.06650
Frozen
0.02123
0.02446
0.06700
Appendix
Table 11: Interaction strong residuals at T=0.3 , τ=0.0025 . All grids evaluate the same final two learned states with the full-periodic operator, without retraining.
The space P2(Rd) of probability measures with finite second moment carries a natural geometry: the quadratic Wasserstein distance W_2 makes it a complete metric space and, following Otto, a (formal) Riemannian manifold whose geodesics are the optimal-transport interpolations. On this manifold, the gradient flow of the free energy F(rho) = KL(rho || π) is exactly the Fokker-Planck equation, and its implicit-Euler discretization is the JKO scheme. This is the geometry underlying diffusion models: the forward process descends the free energy, and each denoising step realizes one JKO step, which recovers DDPM, DDIM, NCSN/SMLD, and Energy Matching; this is one scheme, not separate theories. The same manifold supports a second variational principle. Its geodesics - the minimum-action curves of the Benamou-Brenier formula - are precisely the optimal-transport paths that Flow Matching learns. Fixing both endpoints and following the geodesic, generation becomes a deterministic ODE along a straight line, hence far fewer sampling steps. Placing both families of models on one manifold makes their relationship exact: diffusion follows a free-energy gradient flow, an initial-value problem; optimal-transport Flow Matching follows a Wasserstein geodesic, a boundary-value problem. The two reach the same endpoints along different paths.
We propose a neural algorithm for sampling from distributions specified by unnormalized Boltzmann densities. Our approach is based on the Jordan--Kinderlehrer--Otto scheme for the Kullback--Leibler divergence in the Wasserstein--Fisher--Rao geometry (WFR JKO scheme). Our contributions are twofold. First, we prove that, for any fixed step size, the exact WFR JKO iterates converge exponentially fast to the target as the number of iterations tends to infinity. Notably, this result requires no structural assumptions on the target, such as log-concavity or a logarithmic Sobolev inequality. Second, we develop a neural implementation of the WFR JKO scheme that parametrizes its transport and reaction components using reweighted normalizing flows. Numerical experiments on challenging multimodal targets demonstrate the promising performance of the proposed method.
Chenguang Duan, Johannes Hertrich, Gabriele Steidl
Institut für Geometrie und Praktische Mathematik, RWTH Aachen University · Institute of Computer Science, University of Göttingen · Institute of Mathematics, TU Berlin
We study inverse problems where an unknown potential is observed only through samples from the measure it induces by a convex variational principle. Such problems arise in learning costs, energies, and dynamics from distributional data, but the associated forward solution map is typically nonlinear and implicit. We show that its optimality gap nevertheless yields convex empirical objectives for finite-dimensional potential classes, and we introduce sharpened Fenchel--Young losses that add a data-dependent discrepancy inside the forward problem. This keeps the estimator calibrated while improving the local geometry of the loss. Our main stability theorem separates the inverse error analysis into measurement error, forward perturbation, and empirical curvature. We instantiate this principle for inverse entropic unbalanced optimal transport and for inverse Jordan--Kinderlehrer--Otto (JKO) learning from independent snapshot samples, obtaining high-probability parameter recovery bounds. JKO schemes discretize Wasserstein gradient flows through a sequence of variational problems over measures, making them a natural language for population dynamics observed through snapshots. In this JKO case, the sharpened objective reduces to an unbalanced transport problem, which also clarifies the connection between variational gap losses and quadratic iJKO⋆ surrogates. Numerical experiments illustrate the conditioning effect of sharpening and its benefits for sparse inverse-gradient-flow recovery.
Francisco Andrade, Gabriel Peyré, Clarice Poon
INRIA (PreMeDICaL & HeKA) · CNRS and ENS, PSL Universit´e · Mathematics Institute, University of Warwick