You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Numpy如何沿指定维度折叠数组,按特定索引取值比较保留对应元素

解决方案

numpy没有直接实现该逻辑的单一内置函数,但可以组合np.argmax和高级索引实现,全程向量化操作,适配任意维度、任意分组数的场景:

实现逻辑

  1. 先提取你用来做比较规则的基准值数组:沿指定折叠维度,取每组特定索引位置的所有值
  2. 对基准值数组沿折叠维度求最大值对应的索引max_idx
  3. 通过高级索引用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.25 13:15:04