使用Numpy argmin()查找最接近元组时结果异常的问题
解决元组数组中查找最接近查询元组的索引问题
问题原因
你当前的代码仅计算了单个元素的绝对差值,np.argmin默认会把二维差值数组扁平化,返回整个数组中最小元素的位置,而非「每行总差值最小」的行索引。比如索引8的元组(0.0, 21.0, 8.0)与查询元组的第三个元素差值为0.4,这个值是所有单个元素差值里最小的,因此np.argmin返回了它对应的扁平化索引,但这并非你需要的「整体最接近」的元组索引。
正确解决方案
需要先计算每个元组与查询元组的元素级绝对差总和(或其他距离指标,比如欧氏距离),再找总和最小的行索引:
方法1:使用曼哈顿距离(绝对差总和)
import numpy as np array_of_tuples = np.array([(0.0, 6.5, 1), (0.0, 6.5, 4.5), (0.0, 6.5, 8.0), (0.0, 13.5, 1), (0.0, 13.5, 4.5), (0.0, 13.5, 8.0), (0.0, 21.0, 1), (0.0, 21.0, 4.5), (0.0, 21.0, 8.0), (7.0, 6.5, 1), (7.0, 6.5, 4.5), (13.5, 21.0, 8.0)]) query = np.array((13.1, 20.3, 8.4)) # 计算每行的绝对差总和 total_diffs = np.sum(np.abs(array_of_tuples - query), axis=1) # 获取总差值最小的索引 closest_index = np.argmin(total_diffs) print(closest_index) # 输出11,对应最后一个元组
方法2:使用欧氏距离(平方差总和)
如果需要更侧重较大差值的影响,可以用欧氏距离(无需开根号,因为平方根不影响最小值的位置):
# 计算每行的平方差总和 squared_diffs = np.sum((array_of_tuples - query)**2, axis=1) closest_index = np.argmin(squared_diffs) print(closest_index) # 同样输出11
验证结果
索引11的元组(13.5, 21.0, 8.0)与查询元组的绝对差为(0.4, 0.7, 0.4),总和为1.5;而索引8的元组绝对差总和为13.1 + 0.7 + 0.4 = 14.2,显然前者整体更接近。
内容的提问来源于stack exchange,提问作者Josh
相关产品推荐
相关产品推荐

