Numpy如何沿指定维度折叠数组,按特定索引取值比较保留对应元素
解决方案
numpy没有直接实现该逻辑的单一内置函数,但可以组合np.argmax和高级索引实现,全程向量化操作,适配任意维度、任意分组数的场景:
实现逻辑
- 先提取你用来做比较规则的基准值数组:沿指定折叠维度,取每组特定索引位置的所有值
- 对基准值数组沿折叠维度求最大值对应的索引
max_idx - 通过高级索引用
max_idx从原数组中提取对应分组的全部元素,调整维度后得到最终结果
示例1:数组A的实现
import numpy as np A = np.array([[[1, 2, 3, 4], [0, 1, 2, 1]], [[5, 6, 7, 8], [1, 0, 3, 1]]]) # 提取基准值:沿axis=0折叠,每组第1行(dim=1索引为1)的所有值 compare_vals = A[:, 1, :] # 求axis=0维度上最大值对应的索引 max_idx = np.argmax(compare_vals, axis=0) # 高级索引取值并调整维度 res = A[max_idx, :, np.arange(A.shape[2])].T print(res)
输出:
[[5 2 7 4] [1 1 3 1]]
示例2:数组B的实现
B = np.array([[[ 1, 2, 3, 4], [ 5, 6, 7, 8], [ 0, 1, 2, 1]], [[ 9, 10, 11, 12], [13, 14, 15, 16], [ 1, 0, 3, 1]], [[17, 18, 19, 20], [21, 22, 23, 24], [ 0, 0, 0, 2]]]) # 提取基准值:沿axis=0折叠,每组第2行(dim=1索引为2)的所有值 compare_vals = B[:, 2, :] max_idx = np.argmax(compare_vals, axis=0) res = B[max_idx, :, np.arange(B.shape[2])].T print(res)
输出:
[[ 9 2 11 20] [13 6 15 24] [ 1 1 3 2]]
通用封装函数
可以封装为通用函数适配任意维度场景:
def fold_arr_by_max_pos(arr, fold_axis=0, compare_pos=None): """ 沿指定维度折叠数组,保留基准位置值最大的分组的全部元素 :param arr: 输入numpy数组 :param fold_axis: 要折叠的维度,默认0 :param compare_pos: 基准位置的索引元组,长度为arr.ndim-1(排除折叠维度) :return: 折叠后的数组 """ if compare_pos is None: raise ValueError("请指定基准位置compare_pos") # 构造基准值的索引 idx = [] pos_ptr = 0 for dim in range(arr.ndim): if dim == fold_axis: idx.append(slice(None)) else: idx.append(compare_pos[pos_ptr]) pos_ptr += 1 compare_vals = arr[tuple(idx)] max_idx = np.argmax(compare_vals, axis=fold_axis) # 构造高级索引网格 grid = np.ogrid[tuple(slice(0, s) for s in compare_vals.shape)] grid.insert(fold_axis, max_idx) res = arr[tuple(grid)] # 调整维度顺序 return np.moveaxis(res, fold_axis, -1)
使用示例:
# 数组A调用:折叠维度0,基准位置为dim=1的索引1 res_a = fold_arr_by_max_pos(A, fold_axis=0, compare_pos=(1,)) # 数组B调用:折叠维度0,基准位置为dim=1的索引2 res_b = fold_arr_by_max_pos(B, fold_axis=0, compare_pos=(2,))
内容的提问来源于stack exchange,提问作者baard
相关产品推荐
相关产品推荐

