熵正则化最优传输 [entropic_regularized_optimal_transport]

1. 离散最优传输定义

最优传输涉及到具体的计算问题,在计算机中实际计算时是离散化的,具体表示如下:

a \in \mathbb{R}^nb \in \mathbb{R}^m 分别表示源分布与目标分布,满足 a_i \ge 0b_j \ge 0\sum_i a_i = \sum_j b_j = 1

定义:

  • 运输矩阵 P \in \mathbb{R}^{n\times m},其中 P_{ij} 表示从第 i 个点运到第 j 个点的质量。
  • 代价矩阵 C \in \mathbb{R}^{n\times m},其中 C_{ij}=c(x_i,y_j)

离散最优传输问题就是求如下的式子

\min_{P}\sum_{i,j} C_{ij}P_{ij}

满足 P\mathbf{1}_m = a(行约束:源分布质量守恒),P^T\mathbf{1}_n = b(列约束:目标分布质量守恒),P_{ij}\ge0

2. 为什么需要熵正则?

普通 OT 存在问题:变量数很多(n\times m),约束条件数目是 (n + m),线性规划计算昂贵,解可能稀疏、不稳定。

因此加入熵正则项 H(P)=-\sum_{i,j}P_{ij}\log P_{ij},将目标函数变为:

\min_P \sum_{i,j} C_{ij}P_{ij} - \varepsilon H(P) = \min_P \sum_{i,j} C_{ij}P_{ij} + \varepsilon\sum_{i,j}P_{ij}\log P_{ij}

其中 \varepsilon>0 为正则化强度,熵越大,运输计划越平滑。

这里简单解释一下这个正则化的意思 当\varepsilon\rightarrow\infty时,正则项会驱使解向满足约束的熵最小方向靠近,根据一些信息论知识,这个分布就是边缘分布的向量积; 当\varepsilon\rightarrow 0时,就逐渐接近普通的最优传输。

3. 熵正则化 OT 的拉格朗日形式

根据简单的优化知识,用拉格朗日乘数法可以求解P的解析表达式,引入拉格朗日乘子 f_i(对应行约束)、g_j(对应列约束),可以构造拉格朗日函数如下:

\mathcal{L}(P,f,g) = \sum_{i,j}\left(C_{ij}P_{ij}+\varepsilon P_{ij}\log P_{ij}\right) - \sum_i f_i\left(a_i-\sum_j P_{ij}\right) - \sum_j g_j\left(b_j-\sum_i P_{ij}\right)

对每个 P_{ij} 求偏导并令其为 0可得:

\frac{\partial \mathcal{L}}{\partial P_{ij}} = C_{ij}+\varepsilon(\log P_{ij}+1)-f_i-g_j = 0

整理得 \log P_{ij} = \dfrac{f_i+g_j-C_{ij}}{\varepsilon}-1,于是:

P_{ij} = \exp\!\left(\frac{f_i}{\varepsilon}\right)\exp\!\left(-\frac{C_{ij}}{\varepsilon}\right)\exp\!\left(\frac{g_j}{\varepsilon}\right)e^{-1}

把常数吸收进变量,令 u_i = \exp\!\left(\dfrac{f_i}{\varepsilon}\right)v_j = \exp\!\left(\dfrac{g_j}{\varepsilon}\right)K_{ij}=\exp\!\left(-\dfrac{C_{ij}}{\varepsilon}\right),则:

P_{ij}=u_i K_{ij} v_j

4. 矩阵形式

在上一小节求出的表达式可以用矩阵表示,便于放进GPU中处理。 定义 u \in \mathbb{R}^nv\in\mathbb{R}^m,以及 K=\exp(-C/\varepsilon),则:

P = \mathrm{diag}(u)\,K\,\mathrm{diag}(v)

上面式子只是解析式,理论上来说要代入约束条件具体求值,但是现在有n\times m个未知数,约束方程只有n + m个,在大规模情况下,解肯定不止一个,所以还需要加点条件,或者引入什么启发性的算法来求解。下面介绍最原始的Sinkhorn算法

5. 由边缘约束得到 Sinkhorn 迭代

P\mathbf{1}_m=a,代入得 \mathrm{diag}(u)K\mathrm{diag}(v)\mathbf{1}_m=a,即 u \odot (Kv)=a,所以u应该是:

u = \frac{a}{Kv}

同理,由 P^T\mathbf{1}_n=b,得:

v=\frac{b}{K^Tu}

uv是相互依赖的关系,如果明确获取其中的一个,那么就能得到另一个,问题是现在谁也不知道,所以就引入了Sinkhorn-Knopp 迭代算法,总要迈出第一步,然后逐步修正,到最后收敛,算法是这样的: 初始化 v^{(0)}=\mathbf{1},迭代:

u^{(t+1)} = \frac{a}{Kv^{(t)}}, \qquad v^{(t+1)} = \frac{b}{K^Tu^{(t+1)}}

最终 P^\star = \mathrm{diag}(u)K\mathrm{diag}(v)。 这个算法很简单,但是为什么会收敛到特定值上呢,聪明努力的数学家已经证明了收敛性,但是作为学习者,我们还需要一些直觉上的理解,下面给出两个直觉上的解释。

6. Sinkhorn 为什么会收敛:矩阵缩放角度

熵正则最优传输的最优解具有如下形式:

P^\star = \operatorname{diag}(u) K \operatorname{diag}(v)

其中

K_{ij} = e^{-C_{ij}/\varepsilon}

因为 \varepsilon > 0,并且通常假设 C_{ij} 有限,所以

K_{ij} > 0.

也就是说,熵正则把原来的最优传输问题转化成了一个正矩阵缩放问题

给定一个正矩阵 K,寻找两个正向量 u, v,使得矩阵 K 的行和等于 a,列和等于 b。 行边缘约束要求:\operatorname{diag}(u)Kv = a. 逐元素写就是:u_i (Kv)_i = a_i.

因此,在固定 v 的情况下,唯一能让行和满足 a 的更新是:

u_i = \frac{a_i}{(Kv)_i}.

同理,列边缘约束要求:\operatorname{diag}(v)K^\top u = b. 逐元素写就是:v_j (K^\top u)_j = b_j.

因此,在固定 u 的情况下,唯一能让列和满足 b 的更新是:

v_j = \frac{b_j}{(K^\top u)_j}.

所以 Sinkhorn 迭代就是:

u^{(t+1)} = \frac{a}{K v^{(t)}},
v^{(t+1)} = \frac{b}{K^\top u^{(t+1)}}.

直观地说,Sinkhorn 迭代就是不断做两件事:

  1. 固定列缩放 v,修正行缩放 u,让行和等于 a
  2. 固定行缩放 u,修正列缩放 v,让列和等于 b

也就是:

\text{修正行} \rightarrow \text{修正列} \rightarrow \text{修正行} \rightarrow \text{修正列} \rightarrow \cdots

为什么这种反复的行列缩放不会一直震荡,而是会收敛?一个简单的理解是:Sinkhorn 迭代对应的映射是压缩映射

把 Sinkhorn 更新只写成关于 v 的形式:

v^{(t+1)} = F(v^{(t)}) = \frac{b}{K^\top\left(\frac{a}{Kv^{(t)}}\right)}.

也就是说,每一轮迭代都是在对当前的 v 作用同一个映射 F

由于 K 是严格正矩阵,乘以 KK^\top 会把向量中极端的比例差异“平均掉”。因此,两次不同初始值经过一次 Sinkhorn 更新后,它们之间的差异会变小。

换句话说,存在某种合适的距离 d 和常数 0<\gamma<1,使得

d(F(v),F(w)) \leq \gamma d(v,w).

这就是压缩映射。

压缩映射有一个重要性质:反复迭代不会产生真正的周期震荡,而是会收敛到唯一的不动点。

所以 Sinkhorn 迭代最终会收敛到某个 v^\star,满足

F(v^\star)=v^\star.

再由

u^\star=\frac{a}{Kv^\star}

得到对应的 u^\star

于是

P^\star = \operatorname{diag}(u^\star)K\operatorname{diag}(v^\star)

同时满足行边缘和列边缘约束,这个式子又是从拉格朗日乘数法中求得,因此就是熵正则最优传输的最优传输计划。需要注意的是,这里的“唯一”指的是最终传输矩阵 P^\star 唯一;缩放因子 u^\star,v^\star 本身可以相差一个整体倍数。

7. Sinkhorn 是交替 KL 投影

还可以从集合投影的角度理解 Sinkhorn 迭代。 定义两个约束集合:\mathcal C_a = \{\pi : \pi \mathbf 1 = a\}, \mathcal C_b = \{\pi : \pi^\top \mathbf 1 = b\}.

其中:

  • \mathcal C_a 表示所有行边缘等于 a 的传输矩阵;
  • \mathcal C_b 表示所有列边缘等于 b 的传输矩阵。

熵正则最优传输可以理解为:在满足边缘约束的所有传输计划中,寻找一个离 Gibbs 核 K 最近的矩阵。这里的“距离”不是欧氏距离,而是 KL 散度。

具体来说,熵正则最优传输目标可以写成

\min_{\pi \in \mathcal C_a \cap \mathcal C_b} \sum_{i,j} C_{ij}\pi_{ij} + \varepsilon \sum_{i,j}\pi_{ij}\log \pi_{ij}.

K_{ij}=e^{-C_{ij}/\varepsilon},

则有

C_{ij}=-\varepsilon \log K_{ij}.

代入目标函数:

\sum_{i,j} C_{ij}\pi_{ij} + \varepsilon \sum_{i,j}\pi_{ij}\log \pi_{ij} = \sum_{i,j}(-\varepsilon \log K_{ij})\pi_{ij} + \varepsilon \sum_{i,j}\pi_{ij}\log \pi_{ij}.

整理可得

= \varepsilon \sum_{i,j} \pi_{ij} \left( \log \pi_{ij} - \log K_{ij} \right).

也就是

= \varepsilon \sum_{i,j} \pi_{ij} \log \frac{\pi_{ij}}{K_{ij}}.

因此,在约束集合

\mathcal C_a \cap \mathcal C_b

上,原问题等价于

\min_{\pi \in \mathcal C_a \cap \mathcal C_b} \sum_{i,j} \pi_{ij} \log \frac{\pi_{ij}}{K_{ij}}.

这个量可以理解为 \pi 到未归一化核 K 的 KL 型距离。

所以从集合投影的角度看,熵正则 OT 就是在所有满足边缘约束的传输矩阵中,寻找一个最接近 Gibbs 核 K 的矩阵:

\pi^\star = \arg\min_{\pi \in \mathcal C_a \cap \mathcal C_b} \sum_{i,j} \pi_{ij} \log \frac{\pi_{ij}}{K_{ij}}.

Sinkhorn 迭代做的事情是交替把当前矩阵投影到两个集合上:

\mathcal C_a \rightarrow \mathcal C_b \rightarrow \mathcal C_a \rightarrow \mathcal C_b \rightarrow \cdots

也就是说:

  1. 先把当前矩阵投影到满足行边缘约束的集合 \mathcal C_a
  2. 再把它投影到满足列边缘约束的集合 \mathcal C_b
  3. 不断重复这个过程。

在 KL 几何下,对 \mathcal C_a 的投影正好对应“按行缩放”:

u_i = \frac{a_i}{(Kv)_i}.

\mathcal C_b 的投影正好对应“按列缩放”:

v_j = \frac{b_j}{(K^\top u)_j}.

因此,Sinkhorn 迭代也可以看成:

在两个凸集合 \mathcal C_a\mathcal C_b 之间做交替 KL 投影。

这和欧氏空间里的交替投影类似:如果有两个凸集,我们反复投影到第一个集合,再投影到第二个集合,通常会逐渐靠近两个集合的交集。

在这里,两个集合的交集就是:

\mathcal C_a \cap \mathcal C_b = \{\pi : \pi \mathbf 1 = a,\ \pi^\top \mathbf 1 = b\}.

也就是所有满足两个边缘分布约束的 coupling。

由于 \mathcal C_a\mathcal C_b 都是凸集,KL 散度对应的 Bregman 投影具有良好的收敛性质,所以交替 KL 投影会收敛到交集中的最优点。

因此,从集合角度看,Sinkhorn 收敛的原因是:

它不是随意地更新矩阵,而是在 KL 几何下不断把矩阵投影回”正确行边缘”和”正确列边缘”这两个凸集合;凸性和 KL 投影结构保证了这个过程会收敛。

小结

这篇笔记的主线是:普通 OT 计算昂贵 \to 加入熵正则 \to 拉格朗日求导得到解析结构 P_{ij}=u_iK_{ij}v_j \to 写成矩阵形式 P=\operatorname{diag}(u)K\operatorname{diag}(v) \to 代入边缘约束导出 Sinkhorn 迭代。

Sinkhorn 迭代的收敛性可以从两个角度理解:从矩阵缩放角度看,迭代映射 F 是压缩映射,不动点唯一;从 KL 投影角度看,每一步都是在两个凸集合 \mathcal{C}_a\mathcal{C}_b 之间做交替 Bregman 投影,凸性保证收敛。

两个视角是互补的:前者说明”为什么不震荡”,后者说明”收敛到的东西正是我们想要的最优解”。