Scipy中Wasserstein距离的支撑集定义及计算疑问
关于scipy.stats.wasserstein_distance的使用误区与支撑集的作用
你遇到的问题核心是对wasserstein_distance的参数含义理解有误——这个函数的前两个输入参数不是“按索引对应支撑点的权重数组”,而是两个分布的样本点集合(或支撑点列表)。
为什么你的调用结果是0?
你传入的[0,1,0]和[1,0,0]会被函数当作两个分布的样本数据。函数会自动统计样本的出现频率,生成经验分布:
- 第一个数组的样本是
0、1、0,对应经验分布:点0的权重是2/3,点1的权重是1/3 - 第二个数组的样本是
1、0、0,对应经验分布:点0的权重是2/3,点1的权重是1/3
两个分布完全一致,所以Wasserstein距离自然是0。
支撑集的作用
Wasserstein距离的计算完全依赖于两个分布的支撑集(即概率质量不为0的点的集合)和每个支撑点上的权重。它衡量的是将一个分布的所有质量移动到另一个分布所需的最小总代价(质量×移动距离),只有支撑集上的点才会参与代价计算。
正确实现你预期结果的方式
如果你想表达“第一个分布在点1处有全部质量,第二个分布在点0处有全部质量”,有两种正确调用方式:
- 直接传入单个支撑点(因为每个分布只有一个非零质量的点):
from scipy.stats import wasserstein_distance result = wasserstein_distance([1], [0]) # 结果为1,符合预期
- 指定支撑点和对应权重(适合多支撑点的场景):
from scipy.stats import wasserstein_distance # 支撑点列表是[0,1],u_weights对应第一个分布各点的权重,v_weights对应第二个 result = wasserstein_distance([0, 1], [0, 1], u_weights=[0, 1], v_weights=[1, 0]) # 结果同样为1
内容的提问来源于stack exchange,提问作者mzzx
相关产品推荐
相关产品推荐

