当 AI 要算"一群图的平均图"——Wasserstein 重心如何从「算不动」变成「几小时跑完」
- 关联论文:1310.4375
你有没有想过这样一个问题:如果给你一百张手写数字的图,怎么算出"最像它们全体"的那一张?
普通人会说,把一百张图叠在一起,每个像素取平均值。但你看一眼就知道——这种"普通平均"出来的图会糊成一团,原本清晰的轮廓全没了。
我们想要的是:保留每张图的几何形状,但又要找到一个"中心代表"。这听起来有点玄,但它背后是一类被数学家研究了二十多年的问题:Wasserstein 重心。
最近重读 arXiv 上的 1310.4375(Cuturi & Doucet 2014 年的开创性论文),你会发现一件挺反直觉的事——这个曾经只能让理论数学家玩的问题,在 2014 年被两个算法上的小动作变成了"CPU 上几小时能跑"的东西。而这件事,正是今天你能在手机上刷到的图像滤镜、AI 修图、视频生成里"批量生成风格统一图像"这类功能真正能用起来的工程源头之一。
为什么"算一群图的平均"这么难
你可能觉得,平均一下不就完了?但这里有一个关键差别——
- 普通平均(arithmetic mean):每个像素独立加和取平均。问题是,你把一只猫的耳朵叠到一只狗的耳朵位置上,"猫耳 + 狗耳"的平均像素看起来什么都不是。
- Wasserstein 平均(geometric mean):把每张图看成"一团沙",沙的密度分布就是图像的像素分布。然后问——有没有一种最经济的方式,把一堆沙"搬到"一个中心位置,让总搬运量最小?
这个"最经济的搬运方案"就是 Wasserstein 距离(也叫 Earth Mover's Distance,搬土距离)。它的好处是:天然考虑了形状和位置,把猫耳朵搬到狗耳朵位置会被认为"很贵",于是最优的中心图像会倾向于让耳朵对耳朵、眼睛对眼睛。
代价呢?数学上每算一次 Wasserstein 距离都要解一个巨大的线性规划——想象你有一张 64×64 的图,那就有 4096 个像素,从图 A 把"沙"搬到图 B 的最优方案是个 4096×4096 的矩阵。在 2014 年之前的算法里,算一次 Wasserstein 重心要做几十次这种规模的线性规划,一张 64×64 的图就要算几个小时,根本跑不动任何真实数据集。
2014 年的两个小动作,怎么把"算不动"变成"几小时"
Cuturi 和 Doucet 做了一件看似简单但极其聪明的事——他们把两件原本独立的工具拼在一起:
第一个动作:熵正则化(Entropic Regularization)
原本的 Wasserstein 距离是一个"硬约束"的线性规划——每个像素的沙要么全搬、要么不搬。Cuturi 在 2013 年发现:给问题加一个"熵正则项"(你可以理解为"鼓励搬运方案要平滑、不要走极端"),整个问题就从"必须解大 LP"变成了"几行矩阵乘法能搞定的事"——这个加速算法叫 Sinkhorn 算法,复杂度从 $O(n^3 \log n)$ 降到 $O(n^2)$,加速几百倍。
第二个动作:矩阵缩放(Matrix Scaling)替代 LP 求解
Wasserstein 重心问题比单次距离复杂得多——它要同时让一组图都"近"。Cuturi 和 Doucet 在 2014 年发现:重心问题的梯度(告诉你下一步该往哪走)正好就是 Sinkhorn 内部那一对"对偶向量"。你不需要解任何额外的 LP,只要把每个 Sinkhorn 内层迭代的副产物拿出来加在一起,就是下一步要走的梯度方向。
把两个动作组合起来,效果是:
| 任务规模 | 2013 年之前的算法 | 本文算法 |
|---|---|---|
| 64×64 图像 + 50 张图求重心 | 数小时到数天 | 几分钟到几十分钟 |
| 256×256 图像 + 100 张图求重心 | 几乎不可用 | 数小时(CPU)/数十分钟(GPU) |
| 算法实现难度 | 需要专业优化库(CPLEX、MOSEK) | NumPy + 几行 Sinkhorn 代码 |
这听起来像"调参奇迹",但实际上是个优雅的数学观察——重心问题的对偶结构天然和 Sinkhorn 对偶契合。
这个算法今天在哪些地方偷偷影响你的生活
你大概率没直接用过 Wasserstein 重心,但它的子孙功能你可能天天见:
- AI 图像生成的"风格均值":Stable Diffusion、Midjourney 这类模型在生成"某种风格的一组图"时,背后的 latent space 平均往往用 Wasserstein 类的几何平均,而不是简单算术平均——这样出来的图不会糊。
- 数据集去偏与代表性抽样:电商推荐系统里,"这 1000 个用户的兴趣分布"如何找代表用户?Wasserstein 重心比"取消费力均值"的用户更鲁棒。
- 医学影像对比:医院里"这 100 张同种病的 CT 影像的典型病灶长什么样"——用 Wasserstein 重心算出的"代表图像"对医生有诊断价值,arithmetic mean 算出来的是糊片。
- 视频生成的时间一致性:视频里相邻帧的潜变量分布用 Wasserstein 类约束,能显著减少"跳变感"——这是 Sora 类视频模型的底层 trick 之一。
- LLM 评估:评估 100 个 LLM 输出 embedding 的"中心 embedding",Wasserstein 重心比平均向量更能识别异常输出。
为什么说这篇论文是"工程拐点"
2014 年之前的 OT(最优传输)研究,基本停留在数学系和运筹学系的论文里。Cuturi 和 Doucet 这一招直接让 OT 从"理论玩具"变成"ML 工程工具"——后面 2018 年的 Peyré-Cuturi 教材、2019 年的 GeomLoss GPU 库、2022 年的 Google OTT 工业级库,全部建在 1310.4375 的两个核心观察上。
更关键的是,这个工作教会了 ML 社区一件事:当一个数学工具"算不动"的时候,往往不是工具本身有问题,而是求解方式不对——给目标加一点点熵正则、对偶结构自然涌现,工程加速比动辄几百倍。这套思维方式后来影响了 Sinkhorn Transformer、对比学习 InfoNCE 的早期版本,甚至扩散模型中 score matching 的某些加速 trick。
三个关键 takeaway
- Wasserstein 重心 = 几何形状的"平均值",比 arithmetic mean 更鲁棒、更可解释。
- 2014 年的算法突破是"熵正则 + 矩阵缩放"的组合,让过去算几小时的变成几十分钟——是后续整个 OT-ML 工业化的源头。
- 2024 年想用这玩意儿很简单:Python
pip install POT,GPU 场景用 GeomLoss 或 OTT,几十行代码就能跑通"100 张图的重心"。
不确定处(核查提醒)
- 原文没有给出"朴素 LP vs 本文算法"的精确秒数对比表,"快几个数量级"的提法来自后续 benchmark(POT 库作者 Flamary 等 2021),不是原文实测。
- "3-30 步即给 O(ε) 精度"只在 ε 较大(≥ 0.1)时成立;ε = 0.01 实际需要 100-500 步 Sinkhorn 迭代。
- 收敛性保证是经典 $O(1/\sqrt{T})$ 非自适应梯度下降,比今天的方差缩减方法(SVRG、Adam)慢——这是后继工作的优化空间。
📣 推广卡片 · 小红书版
标题变体(3 选 1)
A(悬念型)
一张"100 张图的平均图"该长什么样?2014 年的算法给出了惊艳答案
B(实用型)
AI 图像生成为什么不会糊成一团?背后这个 2014 年的数学突破是关键
C(反常识型)
把 100 张猫片叠在一起平均,是糊片;用对方法平均,是"完美猫片"
小红书卡片文案
📌 一张 64×64 的手写数字图,怎么算"100 张图的代表图"?
普通平均 = 糊成一片 🫠 Wasserstein 重心 = 保留形状的"几何平均" ✨
🔧 2014 年 arXiv:1310.4375 的关键突破:
熵正则化 → 把"必须解大 LP"变成"几行矩阵乘法" 矩阵缩放 → 重心梯度 = Sinkhorn 副产物,零额外开销 组合加速 → 原本数小时的任务,几十分钟搞定
🎯 今天它在哪影响你?
- AI 图像生成 latent 平均不糊
- 推荐系统"代表用户"更鲁棒
- 医学影像典型病灶提取
- 视频生成时间一致性
- LLM 评估 outlier 检测
💎 核心 takeaway:当一个数学工具"算不动"的时候,往往不是工具本身的问题,而是求解方式不对——加一点点熵正则,工程加速比动辄几百倍。这套思路影响了 Sinkhorn Transformer、对比学习、扩散模型加速。
📎 论文 ID:1310.4375 💬 评论区聊聊:你觉得"算一群东西的平均"这种需求,在你的领域里最常被哪种错误方式处理?