寻求快速计算Wasserstein距离的工具:Gudhi过慢,Geomloss结果偏差大
解决二维点集(持久性图)Wasserstein距离计算的结果差异与提速方案
问题核心原因
Gudhi计算的是持久性图专用的Wasserstein距离,默认会把点集里多余的点匹配到对角线(即(t,t)形式的虚拟点,对应同源类消失的情况),且默认用欧氏距离(L2)计算点间距离;而你当前的Geomloss代码直接计算普通点集的Wasserstein距离,既没有处理持久性图的对角线匹配规则,也用了L1点距离,这是结果差异巨大的根本原因。
方案一:调整Geomloss适配持久性图计算逻辑
关键调整点
- 补全对角线虚拟点,让两个点集大小一致(模拟Gudhi的对角线匹配逻辑)
- 统一用欧氏距离(L2)计算点间距离
- 调整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
相关产品推荐
相关产品推荐

