1. 初识POT:Python中的最优传输瑞士军刀
如果你在机器学习和数据分析领域摸爬滚打过一段时间,大概率听说过“最优传输”这个听起来有点玄乎的词。我第一次接触这个概念是在处理图像颜色迁移项目时,当时需要把一张图片的色调风格转移到另一张图片上,试了各种直方图匹配方法效果都不理想,直到发现了POT这个库。简单来说,最优传输要解决的是这样一个问题:如何用最小的“成本”把一个分布(比如一堆沙子)搬运成另一个分布(比如一座沙堡)。这里的“成本”可以是距离、时间、能量等等。而POT库,就是Python中实现这套数学理论的工具箱,它让原本需要大量数学推导才能上手的算法,变成了几行代码就能调用的函数。
POT的全称是Python Optimal Transport,由Rémi Flamary和Nicolas Courty等学者维护,现在已经成为了这个领域的标准工具之一。我刚开始用的时候还是0.7版本,现在都到0.9.6了,功能越来越丰富。这个库最吸引我的地方在于它的“实用性”——它不是为了炫技而存在的数学玩具,而是真正能解决实际问题的工具。无论是计算两个分布之间的Wasserstein距离(这个我们后面会详细讲),还是实现复杂的颜色迁移、领域自适应,POT都提供了简洁的接口。
安装POT非常简单,直接用pip就行。我建议你创建一个新的虚拟环境来安装,避免依赖冲突。打开终端,输入:
pip install pot
如果你用的是Anaconda,也可以用conda安装:
conda install -c conda-forge pot
安装完成后,在Python里导入试试:
import ot
print(ot.__version__)
如果能看到版本号(比如0.9.6),说明安装成功了。这里有个小细节要注意:导入的模块名是ot而不是pot,这是开发者为了简化导入而设计的。我第一次用的时候还疑惑了半天,以为安装错了。
2. 核心概念:从“搬沙子”到Wasserstein距离
在深入代码之前,我们得先搞懂几个核心概念。最优传输理论起源于18世纪,法国数学家蒙日(Monge)研究怎么用最少的成本把一堆沙子运到指定位置。后来康托罗维奇(Kantorovich)把它推广成了更一般的线性规划问题。听起来很学术?我举个更生活的例子:假设你有三个仓库的货物要分配到五个商店,每个仓库的货物量不同,每个商店的需求量也不同,运输成本也各不相同。怎么安排运输计划,让总运费最低?这就是最优传输要解决的问题。
在机器学习中,我们经常把数据看作“分布”。比如,一组图片的像素颜色分布、一组用户的行为特征分布。比较两个分布是否相似,传统方法有KL散度、JS散度等,但这些方法有个致命缺点:它们要求两个分布有相同的“支撑集”(简单理解就是定义域要完全一样)。而Wasserstein距离(也叫Earth Mover‘s Distance,推土机距离)就没这个限制,它直接衡量的是把一个分布“搬”成另一个分布需要的最小成本。
我第一次用Wasserstein距离是在比较两个不同数据集的特征分布时。传统方法因为分布范围不同,结果总是很奇怪,而Wasserstein距离给出了更合理的相似度度量。POT库计算Wasserstein距离的核心函数是ot.emd2()和ot.sinkhorn2(),前者是精确解(计算量大),后者是加了熵正则化的近似解(计算快)。我们来看个最简单的例子:
import numpy as np
import ot
# 创建两个简单的1D分布(直方图)
a = np.array([0.4, 0.6]) # 第一个分布:40%在位置1,60%在位置2
b = np.array([0.3, 0.7]) # 第二个分布:30%在位置1,70%在位置2
# 创建成本矩阵:这里用简单的欧氏距离
M = np.array([[0, 1], # 从位置1到位置1成本0,到位置2成本1
[1, 0]]) # 从位置2到位置1成本1,到位置2成本0
# 计算精确的Wasserstein距离
wd_exact = ot.emd2(a, b, M)
print(f"精确Wasserstein距离: {wd_exact}")
# 计算带熵正则化的Wasserstein距离(计算更快)
reg = 0.1 # 正则化参数
wd_approx = ot.sinkhorn2(a, b, M, reg)
print(f"近似Wasserstein距离: {wd_approx}")
运行这段代码,你会看到两个相似但不完全相同的值。emd2给出的是精确解,但计算复杂度是O(n^3),当分布维度高时就吃不消了。sinkhorn2通过引入熵正则化,把问题变成了可以用迭代快速求解的形式,复杂度降到O(n^2),适合大规模数据。我在实际项目中,除非分布非常小(比如小于100个点),否则基本都用Sinkhorn算法。
3. 实战入门:计算Wasserstein距离的三种姿势
知道了基本概念,我们来点实际的。计算Wasserstein距离有三种常见场景,我挨个给你演示。第一种场景是最简单的:两个一维直方图的比较。比如比较两幅图像的亮度直方图。POT提供了专门的一维函数,计算速度特别快:
import numpy as np
import ot
import matplotlib.pyplot as plt
# 生成两个一维分布(模拟两个图像的亮度直方图)
np.random.seed(42)
n_bins = 100
# 第一个分布:偏向低亮度
hist1 = np.random.beta(a=2, b=5, size=n_bins)
hist1 = hist1 / hist1.sum() # 归一化成概率分布
# 第二个分布:偏向高亮度
hist2 = np.random.beta(a=5, b=2, size=n_bins)
hist2 = hist2 / hist2.sum()
# 方法1:使用通用函数(适用于任意维度)
positions = np.arange(n_bins).reshape(-1, 1) # 位置信息
M = ot.dist(positions, positions, metric='euclidean') # 成本矩阵
wd1 = ot.emd2(hist1, hist2, M)
# 方法2:使用一维专用函数(更快!)
wd2 = ot.wasserstein_1d(positions.flatten(), positions.flatten(), hist1, hist2)
print(f"通用函数结果: {wd1:.4f}")
print(f"一维专用函数结果: {wd2:.4f}")
# 可视化两个分布
plt.figure(figsize=(10, 4))
plt.subplot(1, 2, 1)
plt.bar(range(n_bins), hist1, alpha=0.7, label='分布1')
plt.bar(range(n_bins), hist2, alpha=0.7, label='分布2')
plt.title(f"两个分布,Wasserstein距离={wd2:.4f}")
plt.legend()
# 再生成一个更相似的分布对比
hist3 = np.random.beta(a=2.2, b=4.8, size=n_bins)
hist3 = hist3 / hist3.sum()
wd3 = ot.wasserstein_1d(positions.flatten(), positions.flatten(), hist1, hist3)
plt.subplot(1, 2, 2)
plt.bar(range(n_bins), hist1, alpha=0.7, label='分布1')
plt.bar(range(n_bins), hist3, alpha=0.7, label='分布3')
plt.title(f"两个相似分布,Wasserstein距离={wd3:.4f}")
plt.legend()
plt.tight_layout()
plt.show()
运行这段代码,你会看到第一个图中两个分布差异较大,Wasserstein距离值也较大;第二个图中两个分布形状相似,距离值就小很多。这就是Wasserstein距离的直观意义——分布越相似,“搬运成本”越低。
第二种场景是二维或更高维的点集比较。比如在机器学习中,我们经常需要比较两个特征空间中的数据分布。这时候成本矩阵的计算就关键了:
# 生成两个二维点集(模拟两个数据集的特征分布)
n_points = 50
# 第一个点集:集中在左下角
np.random.seed(42)
X = np.random.randn(n_points, 2) * 0.5 + np.array([0, 0])
# 第二个点集:集中在右上角
Y = np.random.randn(n_points, 2) * 0.5 + np.array([3, 3])
# 为每个点分配均匀权重(每个点重要性相同)
a = ot.unif(n_points) # 均匀分布 [1/n, 1/n, ...]
b = ot.unif(n_points)
# 计算成本矩阵:点之间的欧氏距离
M = ot.dist(X, Y, metric='euclidean')
# 计算Wasserstein距离
wd = ot.emd2(a, b, M)
print(f"两个点集之间的Wasserstein距离: {wd:.4f}")
# 我们还可以得到最优传输矩阵T,看看具体怎么“搬运”的
T = ot.emd(a, b, M)
# 可视化
plt.figure(figsize=(12, 5))
plt.subplot(1, 2, 1)
plt.scatter(X[:, 0], X[:, 1], c='blue', alpha=0.6, label='分布X')
plt.scatter(Y[:, 0], Y[:, 1], c='red', alpha=0.6, label='分布Y')
plt.tit


338

被折叠的 条评论
为什么被折叠?



