跳转至

Amortized Optimal Transport from Sliced Potentials

会议: NeurIPS2026
arXiv: 2604.15114
代码: https://github.com/tmp0810/Sliced-Amortized-OT
领域: 优化 / 最优传输
关键词: 摊销优化、切片势函数、熵正则化、传输计划、条件流匹配

一句话总结

以廉价的一维切片 Kantorovich 势函数作为特征,用共享线性系数预测原空间势函数,再恢复近似传输计划;RA-OT 与 OA-OT 减少训练成本并支持可变原子数,但并非所有场景下推理最快或生成质量最好。

研究背景与动机

最优传输(optimal transport,OT)不仅给出两个分布之间的距离,还给出质量如何从源位置送到目标位置的传输计划。颜色迁移需要这种对应关系,条件流匹配也可以用它为噪声和数据样本配对,使生成轨迹更直。然而,在每个新图像对或每个新 mini-batch 上重新求解 OT,计算成本会反复出现。熵正则化把问题变成可由 Sinkhorn 迭代求解的平滑问题,但仍然需要访问源与目标之间的代价矩阵,并逐对更新势函数。

摊销优化希望把以往问题中的计算经验压进一个可重复使用的预测器。Meta-OT 已经不是直接预测整个传输矩阵,而是预测一个 Kantorovich 势函数,再由熵正则化的关系恢复另一个势函数和计划。困难在于它仍从原始测度表示出发:固定支持点实验中,MLP 输入两组质量权重,输出长度等于源原子数的势向量;支持位置也变化时,还需要编码点云。这样的模型需要学习几何对应关系,参数较多,固定维度 MLP 也不能直接处理不同原子数。这里的限制针对本文讨论和采用的架构,并不意味着所有神经网络都不能接收变长集合。

一维 OT 对严格凸的差值代价具有按分位数匹配的结构,离散问题可通过排序和累计质量匹配高效计算。与其让模型从零理解原始分布,不如先计算多个投影方向上的一维势函数,把已包含传输结构的结果交给模型。需要注意的是,投影丢失信息,一维问题的精确解并不是原空间问题的精确解。核心 idea:学习跨测度对共享的切片势函数组合,将一维 OT 的结构性特征用于预测原空间势函数,再用原空间代价恢复近似计划。

方法详解

整体框架

输入是一对带权测度及其原空间代价函数,离散情况下分别包含源原子、目标原子及其质量。输出是熵正则化 OT 的近似计划,而不是只输出一个 Wasserstein 距离。整个流程依次经过“切片势特征”“共享系数学习”“原空间计划恢复”:第一步为当前输入求解多个一维投影问题,第二步用此前学好的系数组合这些势值,第三步回到原空间代价矩阵计算传输质量。

RA-OT 与 OA-OT 使用相同的特征和预测形式,区别是共享系数从何而来。RA-OT 在训练时需要原空间 OT 求解器提供真实势函数,做回归;OA-OT 不需要真实势标签,而是在原空间熵正则化对偶目标上学习系数。测试时两者都只复用系数,仍须为新测度对重新计算切片势特征,不能把训练阶段的切片结果直接当成所有新输入的解。

%%{init: {'flowchart': {'rankSpacing': 24, 'nodeSpacing': 28, 'padding': 6, 'wrappingWidth': 400}}}%%
flowchart TD
    A["带权源与目标测度<br/>原空间代价"] --> B["切片势特征"]
    B --> C["共享系数学习"]
    R["训练分支 RA-OT<br/>求解器真实势标签"] -.-> C
    O["训练分支 OA-OT<br/>原空间对偶目标"] -.-> C
    C -->|推理使用已学系数| D["原空间计划恢复"]
    A -->|原空间代价与质量约束| D
    D --> E["近似计划<br/>可选舍入或 Sinkhorn 精化"]

图中的虚线只表示训练监督,不属于推理的数据输入。新测度对的特征计算、线性组合和计划恢复都在推理时执行;舍入或迭代精化是可选后处理,不是把一维计划提升到高维的替代名称。

关键设计

1. 切片势特征:让输入先包含传输几何,而不是只包含原始质量

对每个固定投影方向,将源和目标原子映射到一维,同时保留原来的质量权重。求解投影后的 OT,得到该方向上的源势函数,再在每个源原子的投影位置上读取势值。于是,一个源原子不再只由坐标或质量描述,而由多个方向上的“传输势响应”描述。这些值已经依赖当前源与目标的联合关系,并不是脱离目标分布的通用图像特征。

欧氏空间可使用线性投影;球面供需实验则使用适合球面几何的立体投影。投影族应由空间几何和代价选择,不能把任何数据都机械投影到 RGB 或欧氏直线后宣称保留了原始几何。离散一维势可利用排序匹配及互补松弛,或通过 OT 代价对质量权重求梯度获得;本文调用已有一维求解结构,而非学习另一个复杂的一维网络。

若有 100 个方向和 784 个源原子,就形成一个 784 行、100 列的势特征矩阵:行对应源原子,列对应投影方向。目标原子数变化会改变每个一维问题及其势值,却不会改变“100 个特征列”的含义。该组织方式是支持变长测度的关键,不能误说成把所有测度都压成一个与原子位置无关的 100 维全局向量。

2. 共享系数学习:用回归标签或对偶目标校准同一个线性预测器

每个源原子的预测势都是其切片势值的加权和。所有原子、所有训练测度对共享同一组系数;当投影数固定为 100 时,模型只学习 100 个系数,而不是为每个源原子单独存储一个参数。核心形式是原文式(13)后的线性模型:

\[ \hat{f}_{\boldsymbol{\omega}}[\mu,\nu,c](x)=\sum_{l=1}^{L}\omega_l f^{\star}_{\theta_l}[\mu,\nu,c](P^{c}_{\theta_l}(x)). \]

RA-OT 先在训练测度对上求解原空间 OT,得到源势标签,再最小化预测势与真实势之间的平方误差。离散情况下,把各测度对的特征矩阵与标签向量形成正规方程,累计特征的交叉乘积即可求共享系数,不要求每个问题的矩阵行数相同。正文给出逆矩阵形式,同时建议实践中用线性方程求解器;附录实际加入岭正则项,系数为 0.001。因此,“闭式训练”不等于没有训练成本:标签求解、切片计算和线性求解都需要计入。

OA-OT 不先计算真实势标签,而是把预测势带入原空间熵正则化对偶问题,由一次对应更新得到另一侧势,再用梯度优化系数。它与 Meta-OT 共享基于目标的训练思想,改变的是预测器输入及参数化:Meta-OT 接收原始测度表示,OA-OT 接收已经求好的切片势。它节省标签生成,但仍需评估原空间代价相关的对偶目标,并非训练也只做一维排序。

两种策略的监督对象也不同:RA-OT 要拟合求解器给出的势值,OA-OT 要让预测势在原空间对偶目标上表现好。势函数存在加性常数自由度,因此回归依赖求解器给出的势表示及其一致性;缓存未明确给出跨样本规范化细节,不能自行补写一种对齐规则作为作者实现。

原文用最小化符号写式(10)和式(18),却将其中的目标称为此前以最大化形式定义的对偶目标。这里保留“优化对偶目标”的机制表述,不猜补作者精确使用的负号。OA-OT 在本文中针对熵正则化 OT;RA-OT 在概念上也可对无正则 OT 的两侧势分别回归,但下面主实验比较的是熵正则化计划。

共享系数带来很强的结构约束:原空间势必须能被这些切片势有效近似。增加方向数可以丰富候选函数空间,却不能保证有限切片线性组合精确覆盖所有高维 OT 势,更不能保证分布外测度仍采用相同最优组合。附录不同数字组共用系数的实验提供局部经验支持,不是普遍表达能力定理。

3. 原空间计划恢复:预测的是原空间势,不是一维计划的直接平均

得到源势预测后,利用原空间代价和目标质量执行相应的熵正则化势更新,计算另一侧势,再恢复矩阵。按原文离散负熵约定,计划恢复式(7)为:

\[ \hat{P}_{ij}=\exp\left(\frac{\hat{f}_i+\hat{g}_j-C_{ij}}{\epsilon}\right). \]

这里的代价来自原空间的源目标原子对,并不是投影后的距离;质量约束进入另一侧势的更新。原文连续形式还采用了相对于乘积测度的表述,不能不加说明地把其质量因子与上面的离散势约定混用。缓存中的式(8)和式(9)存在质量向量与输出维度不一致的问题,本文笔记不把这些可疑更新式重排成声称精确的作者公式。

该指数矩阵保证非负,但预测势并不等于最优势。只更新另一侧势通常只能使一侧边际准确,另一侧仍可能偏离指定质量,不能因此宣称直接输出的矩阵已经是严格可行的 OT 计划。附录为此测试标准舍入,使矩阵满足两侧边际,也测试把预测势作为 Sinkhorn 热启动,继续迭代到更高精度。

这也解释了复杂度边界:预测势阶段可以主要由投影、排序和线性组合完成,但显式形成完整计划仍需访问源目标代价并输出一个源原子数乘目标原子数的矩阵。不能把正文只列出的切片和线性预测复杂度,当成包含完整矩阵恢复的端到端复杂度;参数个数不随原子数变化,也不意味着内存和运行时间不随原子数变化。

一个完整示例

以附录的跨分辨率 MNIST 为例,源图像有 784 个原子,目标有 196 个原子。每幅图像的像素强度归一化为质量,100 个投影方向分别产生带权一维 OT 问题;源侧势特征矩阵有 784 行、100 列。

同一组 100 个系数将这些特征合成为 784 个源势预测,随后结合 784 行、196 列的原空间代价计算目标势和计划。若下一对图像的源只有 400 个原子,改变的是矩阵行数及当前问题的势值,而不是重新训练一个 400 维输出头。

训练时,RA-OT 使用原空间求解器的势作为这类矩阵对应的标签;OA-OT 使用对应原空间问题的对偶目标。推理时都不再需要真实势标签。若下游要求严格边际可行性,还需明确加入舍入或继续 Sinkhorn,而不是在示例最后默认把近似预测称为精确解。

损失函数 / 训练策略

主实验每项任务构建 1,000 个测度对,按 70/30 划分训练池和测试集,再从训练池取 10、20、50 或 200 对,测试集为 300 对。主设置使用 100 个投影方向;这里的“训练对数”不是图像原子数,也不是每个一维 OT 的样本数。

RA-OT 使用平方误差回归及 0.001 的岭系数;OA-OT、Meta-OT 和可训练 Min-STP 使用 5,000 次梯度更新,OA-OT 学习率为 0.001。MNIST、球面运输和颜色迁移的熵参数分别为 0.1、0.5、0.005。多数实验在 T4 上进行,CIFAR-10 微调使用 40GB A100,不能把跨硬件时长混为同一速度基准。

原文参数独立于原子数的主张针对固定投影族及系数维度。新问题必须重新计算切片特征,RA-OT 的标签预算与 OA-OT 的对偶优化预算也不同;训练较快、参数较少、每对推理较快是三个不同结论。

实验关键数据

主实验

下表摘取原文表 1—3 的训练对数为 50、投影数为 100 的结果,均为 300 个测试对。计划 RMSE 是预测矩阵与收敛 Sinkhorn 参考矩阵逐元素差值的平方平均再开方;不同任务矩阵规模和熵参数不同,RMSE 数值不可跨任务直接比较。

任务 方法 计划 RMSE,均值 ± 标准差 训练时间(s) 每对推理(ms)
MNIST,RMSE 单位 10⁻⁶ Meta-OT 15.54 ± 4.74 37.11 2.39 ± 0.29
MNIST,RMSE 单位 10⁻⁶ RA-OT 7.77 ± 3.06 3.03 39.36 ± 3.61
MNIST,RMSE 单位 10⁻⁶ OA-OT 6.02 ± 2.52 15.78 38.92 ± 2.23
球面运输,RMSE 单位 10⁻⁷ Meta-OT 4.42 ± 1.55 52.07 16.09 ± 2.17
球面运输,RMSE 单位 10⁻⁷ RA-OT 7.82 ± 1.89 2.53 41.96 ± 6.29
球面运输,RMSE 单位 10⁻⁷ OA-OT 3.93 ± 1.92 19.37 41.03 ± 2.76
颜色迁移,RMSE 单位 10⁻⁶ Meta-OT 33.16 ± 10.45 32.91 15.16 ± 0.64
颜色迁移,RMSE 单位 10⁻⁶ RA-OT 9.99 ± 5.61 7.40 17.40 ± 0.83
颜色迁移,RMSE 单位 10⁻⁶ OA-OT 9.00 ± 5.02 18.13 17.76 ± 0.79

OA-OT 在这三项摘取结果中的 RMSE 最低,但 RA-OT 在球面任务上不如 Meta-OT;两种方法的逐对推理都比 Meta-OT 慢。因此,主结论是较低训练成本与更好的部分任务精度,不是全面支配所有基线。

消融实验

下表摘取附录表 9 的颜色迁移消融,RMSE 单位为 10⁻⁶。该附录没有在表题重述训练对数,因此不自行补充消融的训练对数;颜色迁移使用主实验的离散 RGB 任务设置。

投影数 方法 计划 RMSE,均值 ± 标准差 训练时间(s) 每对推理(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

从 3 个方向增至 20 个方向能显著改善颜色计划预测,但 RA-OT 从 20 增至 100 时均值反而略差;结果支持低投影数不足与收益趋于饱和,而非严格单调改善。测量范围内时间变化有限,也不能推导任意投影数下都没有额外开销。

关键发现

  • 可变原子数得到直接验证:附录多分辨率 MNIST 混合 784、400、196 个原子,同一模型仅有 100 个系数;整体 RMSE 分别为 RA-OT 的 5.18×10⁻⁵ ± 5.50×10⁻⁵ 和 OA-OT 的 3.08×10⁻⁵ ± 3.00×10⁻⁵。这与固定 784 原子的主实验不是同一评估分布。
  • 预测计划需要可行性后处理:表 17 中 OA-OT 的源边际误差从 0.1441 降到 1.340×10⁻¹⁶,计划 RMSE 从 6.02×10⁻⁶ 降到 4.91×10⁻⁶。该表标注熵参数为 0.01,而主 MNIST 表 1 为 0.1,尽管部分 RMSE 重复,也不能擅自认作完全相同设置。
  • 顺序处理新问题时存在成本交叉:附录另一组熵参数为 0.01 的 MNIST 计时中,RA-OT、OA-OT 相对逐对从零运行 Sinkhorn 的盈亏平衡点为 34、137 对;相对 Meta-OT 的交叉点约为 1,013.5、599.2 对,之后 Meta-OT 的低逐对成本占优。

CIFAR-10 实验是从同一个已训练 400,000 步的 I-CFM 检查点微调 10 个 epoch,而非从零训练。下表为附录表 10,batch size 为 2,048,使用 50,000 个生成样本评估;NFE 是自适应求解器的函数评估次数,越少表示采样所需计算越少。

方法 FID NFE/sample 训练时间(s) 预训练时间(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

两种方法比 OT-CFM 的微调更快,但仍慢于 ICFM;RA-OT 的 FID 最低,OA-OT 的 NFE 最低,RA-OT 并未改善 NFE。二维 scurve 实验中,OT-CFM 的终点距离为 0.0991,而 RA-OT/OA-OT 为 0.4675/0.4672;轨迹较直不能替代终点分布质量,正文所说的轻微退化应结合这些具体数值理解。

亮点与洞察

  • 将廉价求解器的势而不是原始分布作为学习特征,使 100 个系数承担跨问题校准,而不是学习全部运输几何。这是“先解简单结构,再学习高维修正”的可复用思路。
  • 同一特征表示兼容有标签回归和无势标签的目标优化,可根据离线求解预算选择策略。节省标签计算与节省梯度训练不是同一种收益。
  • 近似势既能直接恢复计划,也能给迭代求解器热启动。对严格可行性要求高的应用,后者比将近似矩阵直接称为最优计划更稳妥。

局限与展望

  • 全局线性组合存在表达能力上限,有限投影不能保证精确原空间 OT;更复杂测度及代价可能需要条件化系数或非线性算子。
  • 参数维度独立于支持大小不等于推理成本独立于支持大小;切片特征和完整代价矩阵仍可能成为吞吐量或内存瓶颈。
  • 分布偏移实验仅考察 MNIST 旋转,不能支持任意领域迁移或任意代价变化下的泛化结论。
  • 原文对偶优化符号、离散势更新维度存在疑点;表 17 与主 MNIST 的熵参数不同,消融球面列也使用不同 RMSE 单位。复现实验应核对代码,而不是把这些差异自行修成一致。
  • CIFAR-10 只证明特定检查点上的短期微调收益;二维任务的终点距离明显退化,说明快速配对与最终生成质量之间仍需更系统的评估。

相关工作与启发

  • vs Meta-OT:两者都通过预测势间接恢复熵正则化计划;Meta-OT 从原始测度编码出发,本文从一维势特征出发。可变原子数优势针对固定维度 Meta-OT MLP,颜色实验的 Meta-OT 实际使用 PointCloud Encoder,不能统一描述为只输入质量权重。
  • vs Min-STP / min-SWGG:这些方法通过投影得到快速计划,本文学习投影势与原空间势之间的联系,再使用原空间代价恢复计划。其目标是逼近原空间参考计划,但有限特征不赋予精确恢复保证。
  • vs Sinkhorn / OT-CFM:本文不是替换 OT 的数学定义,而是预测势或用近似计划加速重复配对;严格求解仍可在预测后继续 Sinkhorn。适用性应由总训练成本、逐对成本、边际可行性和下游质量共同判断。

评分

  • 新颖性: 4/5 — 用切片势构造与支持大小无关的摊销预测器,切入点清晰。
  • 实验充分度: 4/5 — 包含多几何任务、投影消融、可变原子数和高维微调,但生成质量及设置差异需谨慎解读。
  • 写作质量: 3/5 — 主线易懂,公式约定、符号方向和部分跨表设置仍有歧义。
  • 价值: 4/5 — 适合重复 OT 的低训练预算场景,也提供实用热启动;不构成普遍最快求解器。