使用numpy.sqrt计算向量间L2距离时触发RuntimeWarning的问题咨询
使用numpy.sqrt计算向量间L2距离时触发RuntimeWarning的问题咨询
大家好,我最近在实现向量间L2距离计算的时候遇到了一个RuntimeWarning,实在想不通原因,来求助各位大佬。
我收到的警告内容如下:
/tmp/ipykernel_4554/4230604056.py:37: RuntimeWarning: invalid value encountered in sqrt return np.sqrt((X2).sum(axis=1)[:, np.newaxis] - 2 * X.dot(self.X_train.T) + (self.X_train2).sum(axis=1))
我的距离计算方法是这样实现的:
def compute_distances(self, X): return np.sqrt((X**2).sum(axis=1)[:, np.newaxis] - 2 * X.dot(self.X_train.T) + (self.X_train**2).sum(axis=1))
按道理来说,这个返回的矩阵里每个(i,j)位置的元素就是X[i]和self.X_train[j]之间的L2距离,数学上这个表达式的结果肯定是非负的——毕竟L2距离的平方展开后就是这个式子,不可能出现负数。可为什么会触发np.sqrt的无效值警告呢?实在搞不懂问题出在哪里,有没有大佬能帮忙分析一下?
备注:内容来源于stack exchange,提问作者R J
相关产品推荐
相关产品推荐

