You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

寻求快速计算Wasserstein距离的工具:Gudhi过慢,Geomloss结果偏差大

解决二维点集(持久性图)Wasserstein距离计算的结果差异与提速方案

问题核心原因

Gudhi计算的是持久性图专用的Wasserstein距离,默认会把点集里多余的点匹配到对角线(即(t,t)形式的虚拟点,对应同源类消失的情况),且默认用欧氏距离(L2)计算点间距离;而你当前的Geomloss代码直接计算普通点集的Wasserstein距离,既没有处理持久性图的对角线匹配规则,也用了L1点距离,这是结果差异巨大的根本原因。


方案一:调整Geomloss适配持久性图计算逻辑

关键调整点

  1. 补全对角线虚拟点,让两个点集大小一致(模拟Gudhi的对角线匹配逻辑)
  2. 统一用欧氏距离(L2)计算点间距离
  3. 调整Sinkhorn参数,让结果更接近精确值

调整后代码

import torch
from geomloss import SamplesLoss
import numpy as np

dgm1 = np.array([[2.7, 3.7],[9.6, 14.],[34.2, 34.974]])
dgm2 = np.array([[2.8, 4.45],[9.5, 14.1]])

# 补全对角线点,使两个点集大小一致
diff = len(dgm1) - len(dgm2)
if diff > 0:
    # 给dgm2添加diff个对角线点(取(0,0)不影响距离计算)
    dgm2 = np.vstack([dgm2, np.zeros((diff, 2))])
elif diff < 0:
    dgm1 = np.vstack([dgm1, np.zeros((-diff, 2))])

# 转换为torch张量
I1 = torch.Tensor(dgm1)
I2 = torch.Tensor(dgm2)

# 配置SamplesLoss:匹配Gudhi的计算规则
loss = SamplesLoss(
    loss='sinkhorn',
    debias=True,  # 去偏,让Sinkhorn结果更接近精确Wasserstein距离
    p=2,  # 点间距离用欧氏距离(和Gudhi默认一致)
    blur=1e-5,  # 模糊度越小,结果越精确(速度略降)
    scaling=0.9999,  # 缩放因子接近1,提升结果精度
    backend='auto'
)

# 因为Geomloss默认用均匀权重(1/N),乘以点集大小得到和Gudhi一致的结果
result = loss(I1, I2) * len(I1)
print(result)

方案二:更快且结果与Gudhi一致的替代工具

1. Persim库(专门针对持久性图)

Persim是持久性图处理的专用库,计算Wasserstein距离的速度远快于Gudhi,结果完全对齐。

代码示例

import numpy as np
from persim import wasserstein

dgm1 = np.array([[2.7, 3.7],[9.6, 14.],[34.2, 34.974]])
dgm2 = np.array([[2.8, 4.45],[9.5, 14.1]])

# 计算1阶Wasserstein距离,自动处理对角线匹配
dist = wasserstein(dgm1, dgm2, order=1)
print(dist)

2. POT库(通用最优传输库)

POT是专业的最优传输计算库,支持GPU加速,可灵活配置持久性图的计算规则。

代码示例

import numpy as np
import ot

dgm1 = np.array([[2.7, 3.7],[9.6, 14.],[34.2, 34.974]])
dgm2 = np.array([[2.8, 4.45],[9.5, 14.1]])

# 计算点间欧氏距离矩阵
M = ot.dist(dgm1, dgm2, metric='euclidean')

# 处理点集大小差异:添加对角线虚拟点的距离
if len(dgm1) > len(dgm2):
    # 计算dgm1多余点到对角线的最小欧氏距离
    diag_dist = np.array([np.linalg.norm([b, d] - [(b+d)/2, (b+d)/2]) for b, d in dgm1[len(dgm2):]])
    M = np.hstack([M, diag_dist.reshape(-1, 1)])
    # 定义权重(模拟Gudhi的等权重逻辑)
    a = np.ones(len(dgm1))
    b = np.hstack([np.ones(len(dgm2)), np.ones(len(dgm1)-len(dgm2))])
else:
    diag_dist = np.array([np.linalg.norm([b, d] - [(b+d)/2, (b+d)/2]) for b, d in dgm2[len(dgm1):]])
    M = np.vstack([M, diag_dist.reshape(1, -1)])
    a = np.hstack([np.ones(len(dgm1)), np.ones(len(dgm2)-len(dgm1))])
    b = np.ones(len(dgm2))

# 计算1阶Wasserstein距离
wasserstein_dist = ot.emd2(a / a.sum(), b / b.sum(), M) * max(len(dgm1), len(dgm2))
print(wasserstein_dist)

内容的提问来源于stack exchange,提问作者guest

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.24 05:32:03