调用ot.wasserstein_1d计算Wasserstein距离时出现TypeError错误求助
问题:POT库计算1-Wasserstein距离报错TypeError
用户代码:
import ot import numpy as np tab1 = np.random.normal(2,1,1000) tab2 = np.random.normal(0,1,1000) ot.wasserstein_1d(tab1,tab2)
报错信息:
--------------------------------------------------------------------------- TypeError Traceback (most recent call last) <ipython-input-31-5e5baa2dbe40> in <module> 5 tab2 = np.random.normal(0,1,1000) 6 ----> 7 ot.wasserstein_1d(tab1,tab2) C:\ProgramData\Miniconda3\envs\py37_v1\lib\site-packages\ot\lp\solver_1d.py in wasserstein_1d(u_values, v_values, u_weights, v_weights, p, require_sort) 125 u_quantiles = quantile_function(qs, u_cumweights, u_values) 126 v_quantiles = quantile_function(qs, v_cumweights, v_values) --> 127 qs = nx.zero_pad(qs, pad_width=[(1, 0)] + (qs.ndim - 1) * [(0, 0)]) 128 delta = qs[1:, ...] - qs[:-1, ...] 129 diff_quantiles = nx.abs(u_quantiles - v_quantiles) C:\ProgramData\Miniconda3\envs\py37_v1\lib\site-packages\ot\backend.py in zero_pad(self, a, pad_width) 1026 1027 def zero_pad(self, a, pad_width): --> 1028 return np.pad(a, pad_width) 1029 1030 def argmax(self, a, axis=None): TypeError: pad() missing 1 required positional argument: 'mode'
错误原因
这是numpy版本与POT库版本不兼容导致的问题:
- 旧版本numpy中,
np.pad()的mode参数是必填项,必须明确指定填充模式 - 当前使用的POT库版本在实现
zero_pad方法时,直接调用np.pad(a, pad_width)却未传入mode参数,和旧版numpy的要求冲突
解决办法
有两种可行方案:
- 升级numpy版本:将numpy升级到1.17及以上版本,这些版本中
np.pad()的mode参数默认值为'constant',可以省略参数调用 - 降级POT库版本:安装0.8.3及之前的POT版本,这些版本的代码适配了旧版numpy的
np.pad()调用要求
内容的提问来源于stack exchange,提问作者Artashes
相关产品推荐
相关产品推荐

