求两个数组元素匹配所需位移量:现有实现结果存在偏差
计算NumPy数组元素匹配时的最小位移量
问题描述
给定两个等长的NumPy数组,需要为数组a的每个元素,找到数组b中同值且距离最近的元素的位置,计算两者位置差的绝对值作为位移量。
示例
import numpy as np a = np.array([1, 0, 2, 0, 0, 1, 0, 2, 0, 0, 1, 0, 2, 0, 0 ] ) b = np.array([2, 0, 0, 1, 0, 2, 0, 0, 1, 0, 2, 0, 0, 1, 0] ) # 预期结果 c = np.array ([3, 0, 2, 1, 0, 2, 0, 2, 1, 0, 2, 0, 2, 1, 0 ])
预期结果说明:
a[0] = 1,b中最近的1在索引3,位移量为|0-3|=3a[1] = 0,b[1] = 0,位移量为0a[2] = 2,b中最近的2在索引0,位移量为|2-0|=2,以此类推。
现有尝试及问题
尝试的函数实现如下:
def find_shifts_required(x, y): assert len(x) == len(y) result = np.zeros(len(x)) labels = np.unique(x) for lbl in labels: x_pos = np.where(x == lbl) y_pos = np.where(y == lbl) result[x_pos] = np.subtract(x_pos , y_pos ) _shifts = np.absolute(result).astype(int) return _shifts d = find_shifts_required(a,b) print("expected: " , c) print("actual: " , d)
运行结果对比:
expected: [3 0 2 1 0 2 0 2 1 0 2 0 2 1 0] actual: [3 0 2 1 0 3 0 2 1 0 3 0 2 1 0]
问题原因:该函数直接将x和y中同值元素的位置按索引顺序相减,并未找到最近的匹配项,导致部分位置的位移量计算错误。
正确实现方式
方法一:基于NumPy的高效向量化实现
利用np.searchsorted快速定位最近的匹配位置,避免低效循环:
import numpy as np def find_shifts_required(x, y): assert len(x) == len(y) result = np.zeros(len(x), dtype=int) # 预存每个值在y中的位置列表(已按顺序排列) pos_dict = {} unique_vals = np.unique(x) for val in unique_vals: pos_dict[val] = np.where(y == val)[0] for idx, val in enumerate(x): y_pos = pos_dict[val] # 找到当前索引在y_pos中的插入点 insert_idx = np.searchsorted(y_pos, idx) # 收集候选的最近位置:插入点和插入点前一个位置 candidates = [] if insert_idx < len(y_pos): candidates.append(y_pos[insert_idx]) if insert_idx > 0: candidates.append(y_pos[insert_idx - 1]) # 计算最小位移 result[idx] = min(abs(idx - p) for p in candidates) return result # 测试 a = np.array([1, 0, 2, 0, 0, 1, 0, 2, 0, 0, 1, 0, 2, 0, 0 ] ) b = np.array([2, 0, 0, 1, 0, 2, 0, 0, 1, 0, 2, 0, 0, 1, 0] ) c = np.array ([3, 0, 2, 1, 0, 2, 0, 2, 1, 0, 2, 0, 2, 1, 0 ]) d = find_shifts_required(a, b) print("expected: ", c) print("actual: ", d)
方法二:利用Scipy的距离计算(适合小数组)
如果数组规模较小,可以用scipy.spatial.distance.cdist计算所有位置对的距离,再取最小值:
import numpy as np from scipy.spatial.distance import cdist def find_shifts_required(x, y): assert len(x) == len(y) result = np.zeros(len(x), dtype=int) pos_dict = {} for val in np.unique(x): pos_dict[val] = np.where(y == val)[0].reshape(-1, 1) for idx, val in enumerate(x): y_pos = pos_dict[val] # 计算当前索引与所有同值位置的距离 distances = cdist([[idx]], y_pos, metric='cityblock').flatten() result[idx] = distances.min() return result
说明
- 方法一的效率更高,尤其适合大规模数组,因为通过分组和二分查找减少了不必要的计算。
- 方法二更简洁,但
cdist会计算所有位置对的距离,在数组较大时性能会下降。
内容的提问来源于stack exchange,提问作者Mansour
相关产品推荐
相关产品推荐

