Wasserstein距离简介

Wasserstein 距离是一种衡量两个概率分布之间差异的度量,在最优传输理论中有着核心地位。相比于 KL 散度等传统度量,Wasserstein 距离能更好地反映分布之间的几何结构差异,因此在机器学习、计算机视觉等领域得到了广泛应用。

直观理解:推土机距离

Wasserstein 距离最经典的直观解释是推土机距离(Earth Mover’s Distance, EMD):

把一堆土从当前位置搬运到目标位置,最小化总搬运工作量。

假设有两堆形状不同的土堆 P\mathbb{P}Q\mathbb{Q}

  • P\mathbb{P}:当前位置的土堆(源分布)
  • Q\mathbb{Q}:目标形状的土堆(目标分布)
  • c(x,y)c(x, y):从位置 xx 搬运一单位土到位置 yy代价,通常取欧氏距离的 pp 次方 xyp\|x - y\|^p
  • 在所有可能的搬运方案(Transport Plans)中,找到能将 P\mathbb{P} 完全重塑为 Q\mathbb{Q} 且总做功最小的方案

Wasserstein 距离 = 将所有土从 P\mathbb{P} 的形状搬运成 Q\mathbb{Q} 的形状所需的最小总代价

推土机距离示例
图1:推土机距离示意图——将左侧分布 P\mathbb{P} 的土搬运成右侧分布 Q\mathbb{Q} 的形状

Kantorovich 问题:从直觉到数学

从直觉到公式

搬运土堆比喻很直观——把 P\mathbb{P} 形状的土堆搬成 Q\mathbb{Q} 的形状,最小化总搬运代价。现在需要把这个直觉写成数学。

要精确描述一个运输方案,只需要回答三个问题:

  1. 谁搬到谁——xx 处的土有多少运到了 yy
  2. 搬多少——从 xx 运出的总量 = P\mathbb{P}xx 的质量;运入 yy 的总量 = Q\mathbb{Q}yy 的质量
  3. 花多少钱——每单位土从 xxyy 的代价是 c(x,y)c(x,y),总代价 = 每段路程的运量 × 单价,求和

Kantorovich(1942)用联合分布 γ(x,y)\gamma(x, y) 统一回答了这三个问题:

  • γ(x,y)\gamma(x, y):从 xx 运到 yy 的土量(回答"谁搬到谁")
  • γ\gamma 作为运输计划的概率密度(或质量函数),称为 运输计划(transport plan / coupling)

Kantorovich 问题

Wc(P,Q)=infγΓ(P,Q)X×Yc(x,y)dγ(x,y)\mathcal{W}_c(P, Q) = \inf_{\gamma \in \Gamma(\mathbb{P}, \mathbb{Q})} \int_{\mathcal{X} \times \mathcal{Y}} c(x, y) \, d\gamma(x, y)

其中:

Γ(P,Q)={γP(X×Y):Proj1γ=PProj2γ=Q}\Gamma(\mathbb{P}, \mathbb{Q}) = \left\{ \gamma \in \mathcal{P}(\mathcal{X} \times \mathcal{Y}) : \begin{array}{l} \text{Proj}_{1\sharp}\gamma = \mathbb{P} \\ \text{Proj}_{2\sharp}\gamma = \mathbb{Q} \end{array} \right\}

Γ(P,Q)\Gamma(\mathbb{P}, \mathbb{Q}) 是所有边际分别为 P\mathbb{P}Q\mathbb{Q} 的联合分布的集合。

为什么这个形式是好的?

  1. 凸优化问题:目标函数关于 γ\gamma线性的,约束是凸的 → 一个线性规划(在无穷维空间上)
  2. 解总是存在:不要求一个 xx 只能映射到一个 yy,允许分拆运输(mass splitting),总能构造出满足边际约束的 γ\gamma
  3. 对偶性:凸结构使得强对偶成立,因此可以推导出对偶形式

离散情形下的直观理解

在离散情形下,P\mathbb{P}Q\mathbb{Q} 是两个离散分布:

P=i=1mpiδxi,Q=j=1nqjδyj\mathbb{P} = \sum_{i=1}^{m} p_i \delta_{x_i}, \quad \mathbb{Q} = \sum_{j=1}^{n} q_j \delta_{y_j}

运输计划 γ\gamma 变成一个 m×nm \times n 的矩阵:

γ=(γ11γ12γ1nγ21γ22γ2nγm1γm2γmn)\gamma = \begin{pmatrix} \gamma_{11} & \gamma_{12} & \cdots & \gamma_{1n} \\ \gamma_{21} & \gamma_{22} & \cdots & \gamma_{2n} \\ \vdots & \vdots & \ddots & \vdots \\ \gamma_{m1} & \gamma_{m2} & \cdots & \gamma_{mn} \end{pmatrix}

约束条件:

jγij=pi(行和=P的质量),iγij=qj(列和=Q的质量)\sum_{j} \gamma_{ij} = p_i \quad (\text{行和} = \mathbb{P} \text{的质量}), \qquad \sum_{i} \gamma_{ij} = q_j \quad (\text{列和} = \mathbb{Q} \text{的质量})

这正是一个标准的运输线性规划(transportation LP)。

Kantorovich 问题变为:

minγ0i=1mj=1nc(xi,yj)γij,s.t. γ1=p,  γ1=q\min_{\gamma \ge 0} \sum_{i=1}^{m} \sum_{j=1}^{n} c(x_i, y_j) \cdot \gamma_{ij}, \quad \text{s.t. } \gamma\mathbf{1} = p,\; \gamma^\top\mathbf{1} = q