Amortized Optimal Transport from Sliced Potentials¶
Conference: NeurIPS2026
arXiv: 2604.15114
Code: https://github.com/tmp0810/Sliced-Amortized-OT
Area: Optimization / Optimal Transport
Keywords: amortized optimization, sliced potentials, entropic regularization, transport plans, conditional flow matching
TL;DR¶
The paper uses inexpensive one-dimensional Kantorovich potentials as features, predicts original-space potentials with shared linear coefficients, and reconstructs approximate transport plans; RA-OT and OA-OT reduce training costs and support variable numbers of atoms, but are not universally fastest at inference or best in generation quality.
Background & Motivation¶
Optimal transport (OT) provides not only a distance between distributions but also a plan specifying how mass moves from source locations to target locations. Color transfer needs these correspondences, and conditional flow matching can use them to pair noise with data samples and straighten generation trajectories. However, solving OT again for every image pair or mini-batch repeatedly incurs computational costs. Entropic regularization makes the problem smooth and amenable to Sinkhorn iterations, but still requires access to the source–target cost matrix and per-pair potential updates.
Amortized optimization compresses experience from previous problems into a reusable predictor. Meta-OT already avoids predicting the entire transport matrix directly: it predicts one Kantorovich potential and uses entropic relationships to recover the other potential and the plan. Its remaining difficulty is that it starts from raw measure representations. In fixed-support experiments, an MLP receives two mass-weight vectors and outputs a potential vector whose length equals the number of source atoms; changing support locations also requires a point-cloud encoder. Such models must learn geometric correspondences, have more parameters, and a fixed-dimensional MLP cannot directly handle changing atom counts. This limitation concerns the architectures discussed and used in this paper, not all neural networks accepting variable-length sets.
For strictly convex difference costs, one-dimensional OT has a quantile-matching structure, allowing efficient discrete computation through sorting and cumulative-mass matching. Rather than asking a model to understand raw distributions from scratch, the paper first computes one-dimensional potentials along several projection directions and feeds these transport-aware results to the model. Projection nevertheless discards information: an exact solution to a one-dimensional problem is not an exact solution to the original-space problem. Core Idea: learn combinations of sliced potentials shared across measure pairs, use the structural features of one-dimensional OT to predict original-space potentials, and reconstruct approximate plans using the original-space cost.
Method¶
Overall Architecture¶
The inputs are two weighted measures and their original-space cost function; discrete measures contain source atoms, target atoms, and their masses. The output is an approximate entropic OT plan, not merely a Wasserstein distance. The pipeline follows “Sliced Potential Features,” “Shared Coefficient Learning,” and “Original-Space Plan Recovery”: it solves multiple one-dimensional projection problems for the current input, combines their potential values with previously learned coefficients, and returns to the original-space cost matrix to compute transported mass.
RA-OT and OA-OT use the same features and prediction form, differing in how their shared coefficients are obtained. RA-OT requires an original-space OT solver to provide ground-truth potentials during training and performs regression; OA-OT learns coefficients through the original-space entropic dual objective without ground-truth potential labels. At test time, both reuse only the coefficients and must recompute sliced potential features for each new measure pair. Training-time slice solutions are not universal solutions for future inputs.
%%{init: {'flowchart': {'rankSpacing': 24, 'nodeSpacing': 28, 'padding': 6, 'wrappingWidth': 400}}}%%
flowchart TD
A["Weighted source and target measures<br/>Original-space cost"] --> B["Sliced Potential Features"]
B --> C["Shared Coefficient Learning"]
R["Training branch: RA-OT<br/>Solver-provided potential labels"] -.-> C
O["Training branch: OA-OT<br/>Original-space dual objective"] -.-> C
C -->|Inference uses learned coefficients| D["Original-Space Plan Recovery"]
A -->|Original-space cost and mass constraints| D
D --> E["Approximate plan<br/>Optional rounding or Sinkhorn refinement"]
Dashed edges represent training supervision only, not inference inputs. Feature computation for the new measure pair, linear combination, and plan recovery all occur at inference time. Rounding and iterative refinement are optional post-processing, not alternative names for lifting a one-dimensional plan into higher dimensions.
Key Designs¶
1. Sliced Potential Features: encode transport geometry before learning, rather than only raw mass
For each fixed projection direction, source and target atoms are mapped into one dimension while retaining their mass weights. Solving projected OT yields a source potential for that direction, which is evaluated at the projected location of each source atom. A source atom is therefore described not just by its coordinate or mass, but by its transport-potential responses across directions. These values already depend on the relationship between the current source and target; they are not generic image features independent of the target distribution.
Euclidean spaces can use linear projections, whereas the spherical supply–demand experiment uses stereographic projections appropriate to spherical geometry. The projection family should reflect the space and cost. Mechanically projecting arbitrary data into RGB or onto Euclidean lines does not preserve its original geometry. Discrete one-dimensional potentials can be obtained through sorting-based matching and complementary slackness, or from gradients of OT cost with respect to mass weights. The paper uses existing one-dimensional solver structure instead of learning another elaborate one-dimensional network.
With 100 directions and 784 source atoms, the resulting potential-feature matrix has 784 rows and 100 columns: rows correspond to source atoms and columns to projection directions. Changing the number of target atoms changes each one-dimensional problem and its potential values, but not the meaning of the 100 feature columns. This organization enables variable-length measures; it should not be described as compressing every measure into a single 100-dimensional global vector independent of atom locations.
2. Shared Coefficient Learning: calibrate one linear predictor with regression labels or a dual objective
The predicted potential at each source atom is a weighted sum of its sliced potential values. All atoms and training measure pairs share the same coefficient vector. With 100 fixed projections, the model learns only 100 coefficients rather than storing a separate parameter for each source atom. The central form is the linear model following Equation (13):
RA-OT first solves original-space OT on the training pairs to obtain source-potential labels, then minimizes the squared error between predicted and ground-truth potentials. In the discrete setting, feature matrices and label vectors define normal equations. Accumulating feature cross-products yields shared coefficients without requiring every problem to have the same number of matrix rows. The main text gives an inverse-matrix expression while recommending linear-system solvers in practice; the appendix uses ridge regularization with coefficient 0.001. Thus, closed-form training does not imply zero training cost: label generation, slice computation, and the linear solve all count.
OA-OT does not first compute ground-truth potential labels. Instead, it inserts the predicted potential into the original-space entropic dual problem, obtains the other potential through a corresponding update, and optimizes the coefficients with gradients. It shares objective-based training with Meta-OT but changes the predictor's input and parameterization: Meta-OT receives raw measure representations, whereas OA-OT receives already computed sliced potentials. Avoiding label generation does not eliminate evaluation of the original-space cost-dependent dual objective; training is not restricted to one-dimensional sorting.
The supervision targets also differ. RA-OT fits the solver's potential values, whereas OA-OT seeks potentials that perform well under the original-space dual objective. Potentials have additive-constant freedom, so regression depends on a consistent representation of solver-provided potentials. The cache does not clearly specify cross-sample gauge normalization, and no alignment rule is invented here as the authors' implementation.
Equations (10) and (18) use minimization symbols while identifying their objective with the previously maximized dual objective. This note therefore describes the mechanism as optimizing the dual objective without guessing the exact sign used by the authors. OA-OT addresses entropic OT in this paper. RA-OT can conceptually regress both potentials separately for unregularized OT, but the main experiments below compare entropically regularized plans.
Shared coefficients impose a strong structural restriction: original-space potentials must be well approximated by the available sliced potentials. More directions can enrich the candidate function space, but cannot guarantee that a finite linear combination exactly represents every high-dimensional OT potential or that the same optimal combination transfers out of distribution. The appendix's shared-coefficient experiment across digit groups provides local empirical support, not a general expressivity theorem.
3. Original-Space Plan Recovery: predict original-space potentials rather than directly average one-dimensional plans
After predicting the source potential, the method uses the original-space cost and target masses in the corresponding entropic potential update to obtain the other potential and reconstruct the matrix. Under the paper's discrete negative-entropy convention, Equation (7) gives:
The cost here is between source and target atoms in the original space, not their projected distance; mass constraints enter through the other-potential update. The continuous formulation additionally uses a product-measure convention, whose mass factors should not be mixed with the discrete potential convention without explanation. Cached Equations (8) and (9) contain inconsistencies between mass-vector and output dimensions. This note does not rearrange those questionable updates into purportedly exact author formulas.
The exponential matrix is nonnegative, but predicted potentials are not optimal potentials. Updating only the other potential typically makes one marginal accurate while leaving the other potentially inconsistent with prescribed masses. The directly reconstructed matrix therefore cannot automatically be called a strictly feasible OT plan. The appendix tests standard rounding to enforce both marginals, and also uses predicted potentials to warm-start Sinkhorn and continue to higher accuracy.
This distinction also establishes the complexity boundary. Potential prediction can primarily involve projection, sorting, and linear combination, but explicitly forming a complete plan still accesses source–target costs and outputs a matrix with source-count times target-count entries. The slice and linear-prediction complexities listed in the main text should not be treated as end-to-end complexities including full matrix recovery. Parameter counts independent of atom counts do not imply runtime or memory independent of atom counts.
A Worked Example¶
Consider the appendix's cross-resolution MNIST setting, with 784 source atoms and 196 target atoms. Pixel intensities are normalized into mass weights, and 100 projection directions produce weighted one-dimensional OT problems. The source-potential feature matrix has 784 rows and 100 columns.
The same 100 coefficients combine these features into 784 source-potential predictions. A 784-by-196 original-space cost matrix then enables target-potential and plan recovery. If the next pair has only 400 source atoms, the feature-matrix row count and current potential values change; there is no need to retrain a 400-dimensional output head.
During training, RA-OT uses original-space solver potentials as labels for these matrices; OA-OT uses each original-space problem's dual objective. Neither requires ground-truth potential labels at inference. If the downstream application requires strict marginal feasibility, rounding or further Sinkhorn iterations must be explicitly added rather than silently labeling the approximate prediction an exact solution.
Loss & Training¶
Each main task constructs 1,000 measure pairs and splits them 70/30 into a training pool and test set. The methods then use 10, 20, 50, or 200 pairs from the training pool, with 300 test pairs. The main setting uses 100 projection directions. Training-pair counts are neither image atom counts nor sample counts within each one-dimensional OT problem.
RA-OT uses squared-error regression and a ridge coefficient of 0.001. OA-OT, Meta-OT, and trainable Min-STP use 5,000 gradient updates; OA-OT's learning rate is 0.001. The entropy parameters for MNIST, spherical transport, and color transfer are 0.1, 0.5, and 0.005, respectively. Most experiments use a T4, while CIFAR-10 fine-tuning uses a 40GB A100; timings across these devices are not a single speed benchmark.
The claim that parameters are independent of atom counts concerns the dimensions of coefficients under a fixed projection family. New problems still require recomputing sliced features, and RA-OT's labeling budget differs from OA-OT's dual-optimization budget. Faster training, fewer parameters, and faster per-pair inference are three distinct claims.
Key Experimental Results¶
Main Results¶
The following selection from Tables 1–3 uses 50 training pairs, 100 projections, and 300 test pairs. Plan RMSE is the square root of the mean squared elementwise difference between the predicted matrix and a converged Sinkhorn reference. Matrix sizes and entropy parameters differ across tasks, so RMSE magnitudes should not be compared directly across tasks.
| Task | Method | Plan RMSE, mean ± standard deviation | Training time (s) | Per-pair inference (ms) |
|---|---|---|---|---|
| MNIST, RMSE unit 10⁻⁶ | Meta-OT | 15.54 ± 4.74 | 37.11 | 2.39 ± 0.29 |
| MNIST, RMSE unit 10⁻⁶ | RA-OT | 7.77 ± 3.06 | 3.03 | 39.36 ± 3.61 |
| MNIST, RMSE unit 10⁻⁶ | OA-OT | 6.02 ± 2.52 | 15.78 | 38.92 ± 2.23 |
| Spherical transport, RMSE unit 10⁻⁷ | Meta-OT | 4.42 ± 1.55 | 52.07 | 16.09 ± 2.17 |
| Spherical transport, RMSE unit 10⁻⁷ | RA-OT | 7.82 ± 1.89 | 2.53 | 41.96 ± 6.29 |
| Spherical transport, RMSE unit 10⁻⁷ | OA-OT | 3.93 ± 1.92 | 19.37 | 41.03 ± 2.76 |
| Color transfer, RMSE unit 10⁻⁶ | Meta-OT | 33.16 ± 10.45 | 32.91 | 15.16 ± 0.64 |
| Color transfer, RMSE unit 10⁻⁶ | RA-OT | 9.99 ± 5.61 | 7.40 | 17.40 ± 0.83 |
| Color transfer, RMSE unit 10⁻⁶ | OA-OT | 9.00 ± 5.02 | 18.13 | 17.76 ± 0.79 |
OA-OT has the lowest RMSE in these three selections, but RA-OT is worse than Meta-OT on spherical transport; both proposed methods are slower than Meta-OT at per-pair inference. The main conclusion is lower training costs and better accuracy on some tasks, not universal dominance over every baseline.
Ablation Study¶
The following color-transfer ablation is selected from Appendix Table 9, with RMSE in units of 10⁻⁶. Its caption does not restate the training-pair count, so no count is supplied here. Color transfer uses the main experiments' discrete RGB task setting.
| Projection count | Method | Plan RMSE, mean ± standard deviation | Training time (s) | Per-pair inference (ms) |
|---|---|---|---|---|
| 3 | RA-OT | 25.60 ± 10.42 | 6.68 | 16.48 ± 0.96 |
| 3 | OA-OT | 23.96 ± 6.79 | 16.86 | 17.29 ± 1.15 |
| 20 | RA-OT | 9.39 ± 5.50 | 7.03 | 17.64 ± 1.28 |
| 20 | OA-OT | 9.11 ± 5.11 | 17.31 | 17.48 ± 1.37 |
| 100 | RA-OT | 9.99 ± 5.61 | 7.40 | 17.40 ± 0.83 |
| 100 | OA-OT | 9.00 ± 5.02 | 18.13 | 17.76 ± 0.79 |
Increasing from 3 to 20 directions substantially improves color-plan prediction, but RA-OT's mean error increases slightly from 20 to 100. The evidence supports insufficient information at very low projection counts and diminishing returns, not strictly monotonic improvement. Limited timing variation within the measured range also does not imply zero overhead for arbitrary projection counts.
Key Findings¶
- Variable atom counts are directly evaluated: the appendix's multi-resolution MNIST mixes 784, 400, and 196 atoms in one model with only 100 coefficients. Overall RMSE is 5.18×10⁻⁵ ± 5.50×10⁻⁵ for RA-OT and 3.08×10⁻⁵ ± 3.00×10⁻⁵ for OA-OT. This is not the same evaluation distribution as the fixed-784-atom main experiment.
- Predicted plans require feasibility post-processing: Table 17 reduces OA-OT's source marginal error from 0.1441 to 1.340×10⁻¹⁶ and its plan RMSE from 6.02×10⁻⁶ to 4.91×10⁻⁶. The table specifies entropy parameter 0.01, whereas the main MNIST Table 1 uses 0.1; repeated RMSE values do not justify treating these as identical settings.
- Sequential processing exhibits cost crossovers: in a separate appendix MNIST timing setup with entropy parameter 0.01, RA-OT and OA-OT break even against independent cold-start Sinkhorn after 34 and 137 pairs. Their crossovers against Meta-OT are approximately 1,013.5 and 599.2 pairs, beyond which Meta-OT benefits from lower per-pair costs.
The CIFAR-10 experiment fine-tunes the same I-CFM checkpoint previously trained for 400,000 steps for 10 epochs; it does not train from scratch. The following results are from Appendix Table 10, with batch size 2,048 and 50,000 generated evaluation samples. NFE counts function evaluations by the adaptive solver; lower values indicate less sampling computation.
| Method | FID | NFE/sample | Training time (s) | Pretraining time (s) |
|---|---|---|---|---|
| ICFM | 3.638 | 146.61 | 481.9 | 0.0 |
| OT-CFM | 3.630 | 146.86 | 744.2 | 0.0 |
| OA-OT | 3.575 | 146.00 | 607.0 | 12.0 |
| RA-OT | 3.543 | 147.10 | 618.1 | 16.4 |
Both proposed methods fine-tune faster than OT-CFM but remain slower than ICFM. RA-OT has the lowest FID, whereas OA-OT has the lowest NFE; RA-OT does not improve NFE. In the two-dimensional scurve experiment, OT-CFM's endpoint distance is 0.0991, compared with 0.4675/0.4672 for RA-OT/OA-OT. Straighter paths do not substitute for endpoint distribution quality, and the main text's description of slight degradation should be interpreted alongside these values.
Highlights & Insights¶
- Learning from inexpensive solver potentials rather than raw distributions lets 100 coefficients perform cross-problem calibration instead of learning all transport geometry. This suggests a reusable strategy of solving simple structure first and learning a high-dimensional correction afterward.
- One feature representation supports both labeled regression and objective optimization without potential labels, allowing the strategy to reflect the offline solver budget. Saving label computation and saving gradient training are distinct benefits.
- Approximate potentials can directly recover plans or warm-start iterative solvers. For applications with strict feasibility requirements, the latter is safer than calling an approximate matrix an optimal plan.
Limitations & Future Work¶
- A global linear combination has an expressivity ceiling, and finite projections do not guarantee exact original-space OT. More complex measures and costs may require conditioned coefficients or nonlinear operators.
- Parameter dimensions independent of support sizes do not make inference costs support-independent; sliced features and complete cost matrices can still bottleneck throughput or memory.
- Distribution-shift experiments examine only MNIST rotations and do not establish generalization to arbitrary domain shifts or cost changes.
- Dual-optimization signs and discrete potential-update dimensions are questionable; Table 17 uses a different entropy parameter from the main MNIST experiment, and the spherical ablation uses a different RMSE unit. Reproduction should check the code rather than silently harmonize these differences.
- CIFAR-10 establishes only short-term fine-tuning gains from a particular checkpoint. The substantial endpoint-distance degradation on two-dimensional tasks warrants more systematic evaluation of fast pairing versus final generation quality.
Related Work & Insights¶
- vs Meta-OT: both predict potentials to indirectly recover entropic plans, but Meta-OT starts from raw measure encodings while this paper starts from one-dimensional potential features. The variable-atom advantage concerns a fixed-dimensional Meta-OT MLP; the color experiment actually uses a PointCloud Encoder for Meta-OT, so it should not uniformly be described as receiving only mass weights.
- vs Min-STP / min-SWGG: these methods produce fast plans through projections, whereas this paper learns the relationship between projected and original-space potentials before recovering a plan with original-space costs. It aims to approximate original-space reference plans, but finite features do not guarantee exact recovery.
- vs Sinkhorn / OT-CFM: the paper does not replace the mathematical definition of OT; it predicts potentials or accelerates repeated pairing with approximate plans. Sinkhorn can still follow prediction for stricter solving. Suitability depends jointly on total training costs, per-pair costs, marginal feasibility, and downstream quality.
Rating¶
- Novelty: 4/5 — Sliced potentials provide a clear route to an amortized predictor independent of support size.
- Experimental Thoroughness: 4/5 — Multiple geometries, projection ablations, variable atom counts, and high-dimensional fine-tuning are covered, but generation quality and setting differences require care.
- Writing Quality: 3/5 — The main argument is accessible, while formula conventions, objective signs, and some cross-table settings remain ambiguous.
- Value: 4/5 — Useful for repeated OT under limited training budgets and practical warm-starting, but not a universally fastest solver.