Optimal Transport:从 Wasserstein 距离到 Sinkhorn 算法

4 minute read

Published:

参考:https://coderlemon17.github.io/posts/2022/07-16-ot/

问题

给定两个概率分布 P 和 Q,如何衡量它们“有多不一样”?

常见做法有:

• KL 散度

• JS 散度

假设你有 2 个工厂(A, B)和 2 家门店(X, Y)

  • 供应量(源分布r): 工厂A有3吨货,工厂B有7吨货

  • 需求量(目标分布c): 门店X需要4吨,门店Y需要6吨

  • 运输成本(Cost Martrix C):

    A->X 代价2, A->Y 代价5

    B->X 代价4 B->Y 代价1

Wasserstein 距离

现在我们就需要找出一个运输方案,在满足供需条件的前提下,让总运输最低。这在数学上被称为建立一个线性规划(Linear Programming)模型。

(T_{AX}), (T_{BX}), (T_{AY}), (T_{BY})

我们的终极目标是“总运费最低”:

\[\text{总运费} = 2 \cdot T_{AX} + 5 \cdot T_{AY} + 4 \cdot T_{BX} + 1 \cdot T_{BY}\]

设立条件

传统的线性规划不允许任何“模糊”,必须严丝合缝地满足供需:

\[T_{AX} + T_{AY} = 3\] \[T_{BX} + T_{BY} = 7\] \[T_{AX} + T_{BX} = 4\] \[T_{AY} + T_{BY} = 6\] \[T \ge 0\]

求解联立方程组

如果是简单的这种我们能一眼答案,但是在大部分的复杂的情况下满足这些条件的运输方案不是唯一的,而是有无数种。计算机的任务不是“找出一个解”,而是要在这无数个合法解中,像海底捞针一样,找出让“总运费最低”的那一个。这也就是最优化问题。

通用杀手锏:单纯形法 (Simplex Method)

构建多维空间: 计算机把每一个未知数(比如 (T_{AX}))看作空间中的一个维度。所有的限制条件(等式和 (T \ge 0))像是一把把刀,在这个多维空间里切出一个边缘清晰的几何体(叫做多面体)。

顶点即是方案: 这个几何体内部和表面的每一个点,都是满足供需的运输方案。数学家证明了一个伟大的定理:最优解(最低成本)必定出现在这个几何体的“顶点”上。

计算机的寻优逻辑(爬山算法的逆向):

  1. 随机空降: 计算机先随便找一个顶点作为初始方案(比如不管成本,先硬凑出一个满足供需的解)。

  2. 巡视邻居: 计算机计算这个顶点相邻的几个顶点,看看顺着哪条边走,总成本下降得最快。

  3. 不断移动: 顺着成本下降的边,走到下一个顶点。

  4. 到达谷底: 重复上述过程,直到发现周围所有相邻顶点的成本都比现在高。此时,计算机确信已经找到了全局最低成本,直接输出结果。

from scipy.optimize import linprog
import numpy as np

# 1. 目标函数:总运费最低
# 我们把 2x2 的矩阵展平成一维列表: [T_AX, T_AY, T_BX, T_BY]
c_obj = [2, 5, 4, 1]

# 2. 限制条件 (等式约束 A_eq * T = b_eq)
# 每一行代表一个方程的系数,对应未知数顺序 [T_AX, T_AY, T_BX, T_BY]
A_eq = [
    [1, 1, 0, 0],  # 方程 1: 工厂 A 供 3 吨 (T_AX + T_AY = 3)
    [0, 0, 1, 1],  # 方程 2: 工厂 B 供 7 吨 (T_BX + T_BY = 7)
    [1, 0, 1, 0],  # 方程 3: 门店 X 需 4 吨 (T_AX + T_BX = 4)
    [0, 1, 0, 1]   # 方程 4: 门店 Y 需 6 吨 (T_AY + T_BY = 6)
]
b_eq = [3, 7, 4, 6] # 方程右边的常数

# 3. 边界条件: 所有运输量必须 >= 0
bounds = [(0, None), (0, None), (0, None), (0, None)]

# 4. 求解线性规划 (使用单纯形法改进的高效算法)
res = linprog(c_obj, A_eq=A_eq, b_eq=b_eq, bounds=bounds, method='highs')

print("\n--- 传统最优传输 (线性规划) 结果 ---")
if res.success:
    # 将一维结果重新变回 2x2 矩阵以便观察
    optimal_T = np.round(res.x.reshape(2, 2), 2)
    print("绝对完美的运输方案 T:")
    print(optimal_T)
    print("最低总成本:", res.fun)
else:
    print("求解失败")

Sinkhorn

供应量(行要求): (r = \begin{bmatrix} 3 \ 7 \end{bmatrix})(工厂 A 有 3 吨,B 有 7 吨)

需求量(列要求): (c = \begin{bmatrix} 4 \ 6 \end{bmatrix})(门店 X 需 4 吨,Y 需 6 吨)

成本矩阵: (C = \begin{bmatrix} 2 & 5 \ 4 & 1 \end{bmatrix})

第 0 步:将成本转化为“亲和力”矩阵 (K)

算法不直接用成本来算,而是将成本转化为一个“亲和力”矩阵 (K)。成本越低,亲和力越高(使用自然指数 (e) 转换,公式为 (K = \exp(-C))):

\[K = \exp\left(-\frac{C}{\epsilon}\right)\] \[K = \begin{bmatrix} e^{-2} & e^{-5} \\ e^{-4} & e^{-1} \end{bmatrix} \approx \begin{bmatrix} 0.135 & 0.007 \\ 0.018 & 0.368 \end{bmatrix}\]

观察这个矩阵: (K_{22})((0.368))最大,因为工厂 B 到门店 Y 的成本只要 1,亲和力最高。

同时,我们初始化列缩放系数 (v),通常设为全 1:

\[v^{(0)} = \begin{bmatrix} 1 \\ 1 \end{bmatrix}\]

第 1 轮迭代:Sinkhorn 交替缩放

\[P = \text{diag}(u) \cdot K \cdot \text{diag}(v)\]

通俗来说,就是给矩阵 (K) 的每一行乘上一个系数 (u_i),每一列乘上一个系数 (v_j)。

算法通过交替更新这两个系数,直到满足供需条件:

  1. 更新行系数 (u)(强行满足供应量)

计算公式:(u \leftarrow \frac{r}{K v})(注意这里的除法是逐个元素相除

\[K v^{(0)} = \begin{bmatrix} 0.135 & 0.007 \\ 0.018 & 0.368 \end{bmatrix} \begin{bmatrix} 1 \\ 1 \end{bmatrix} = \begin{bmatrix} 0.142 \\ 0.386 \end{bmatrix}\]

然后用供应量 (r) 去除以它:

\[u^{(1)} = \begin{bmatrix} 3 \\ 7 \end{bmatrix} \div \begin{bmatrix} 0.142 \\ 0.386 \end{bmatrix} \approx \begin{bmatrix} 21.13 \\ 18.13 \end{bmatrix}\]

更新列系数 (v):使得矩阵的列和等于门店的需求量 (c)

计算公式:(v \leftarrow \frac{c}{K^T u})

首先算分母 (K^T u^{(1)}):

\[K^T u^{(1)} = \begin{bmatrix} 0.135 & 0.018 \\ 0.007 & 0.368 \end{bmatrix} \begin{bmatrix} 21.13 \\ 18.13 \end{bmatrix} \approx \begin{bmatrix} 3.18 \\ 6.82 \end{bmatrix}\]

然后用需求量 (c) 去除以它:

\[v^{(1)} = \begin{bmatrix} 4 \\ 6 \end{bmatrix} \div \begin{bmatrix} 3.18 \\ 6.82 \end{bmatrix} \approx \begin{bmatrix} 1.26 \\ 0.88 \end{bmatrix}\]

公式为 (P = \text{diag}(u) \cdot K \cdot \text{diag}(v)):

  • (P_{11})(A 到 X)= (21.13 \times 0.135 \times 1.26 \approx \mathbf{3.59})

  • (P_{12})(A 到 Y)= (21.13 \times 0.007 \times 0.88 \approx \mathbf{0.13})

  • (P_{21})(B 到 X)= (18.13 \times 0.018 \times 1.26 \approx \mathbf{0.41})

  • (P_{22})(B 到 Y)= (18.13 \times 0.368 \times 0.88 \approx \mathbf{5.87})

也就是现在的运输方案矩阵 (P^{(1)}) 是:

\[P^{(1)} = \begin{bmatrix} 3.59 & 0.13 \\ 0.41 & 5.87 \end{bmatrix}\]

接下来的过程

由于现在的行和还不准,计算机会开启第 2 轮迭代。这样不断交替更新,通常只需要迭代几十次,“跷跷板”式的偏差就会越来越小。

import numpy as np

# 1. 准备已知数据
r = np.array([3.0, 7.0])  # 供应量 (工厂 A, B)
c = np.array([4.0, 6.0])  # 需求量 (门店 X, Y)
C = np.array([[2.0, 5.0], 
              [4.0, 1.0]]) # 成本矩阵

# 2. Sinkhorn 算法参数
epsilon = 1.0  # 控制模糊度的参数
n_iters = 100  # 交替迭代的次数

# 3. 将成本转化为“亲和力”矩阵 K
# 公式: K = exp(-C / epsilon)
K = np.exp(-C / epsilon)

# 初始化列缩放系数 v (行缩放系数 u 会在循环中计算)
v = np.ones(len(c))

# 4. 开始 Sinkhorn 交替缩放迭代
for i in range(n_iters):
    # 第一步:更新行系数 u,强行满足供应量 r
    u = r / np.dot(K, v)
    
    # 第二步:更新列系数 v,强行满足需求量 c
    v = c / np.dot(K.T, u)

# 5. 计算最终的运输方案 P
# 公式: P = diag(u) * K * diag(v)
P = np.diag(u) @ K @ np.diag(v)

print("--- Sinkhorn 算法结果 ---")
print("最终运输方案矩阵 P:")
print(np.round(P, 2))
print("Sinkhorn 总成本:", np.round(np.sum(P * C), 2))

个人思考

这个 Sinkhorn 通过更新两个系数来进行缩放,让我想到 Transformer 的 QKV 矩阵也是类似的思想:通过给 (x) 乘以 QKV 矩阵来进行变换。

最像的就是中间亲和力矩阵的构建:

在于 Transformer 的 Attention 机制(缩放点积注意力)Sinkhorn 核心逻辑的完美映射。

让我们对比一下它们的计算过程:

A. 第一步:计算“亲和力”矩阵

  • Transformer: 用 Query 和 Key 的点积来衡量 Token 之间的相关性,得到一个原始分数矩阵 (S = QK^T / \sqrt{d})。

  • Sinkhorn: 用工厂到门店的成本来计算亲和力,得到一个核矩阵 (K = \exp(-C/\epsilon))。

  • 共性: 两者都在构建一个全局的 (N \times N) 关系矩阵!

B. 第二步:施加缩放(魔法发生的地方)

  • Transformer 使用 Softmax(单向缩放):

  • Attention 会对原始矩阵 (S) 沿着每一行做 Softmax 操作。Softmax 的本质是什么?就是给每一行乘上一个缩放系数,强行让这一行的和等于 1

  • 这在数学上叫单随机矩阵(Singly Stochastic Matrix)。它保证了“每个词分配出去的注意力总和为 1”。

  • Sinkhorn 使用交替缩放(双向缩放):

  • Sinkhorn 不仅要求每一行(工厂供应)的和等于目标值,还要求每一列(门店需求)的和也等于目标值。所以它用 (u) 缩放行,用 (v) 缩放列。

  • 如果供需目标都是 1,这在数学上叫双随机矩阵(Doubly Stochastic Matrix)

注:欧式距离是量“两个点”之间的直线距离;而 Wasserstein 距离是量“两堆沙子(两个分布)”互相转换的搬运成本。

View the original note on GitHub