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

如何为多维NumPy数组的各行应用不同k值执行partition操作以获取对应第k小元素

如何为多维NumPy数组的各行应用不同k值执行partition操作以获取对应第k小元素

你遇到的问题核心是NumPy的np.partition本身不支持为每行指定不同的kth值,不过我们可以结合argsort和take_along来实现向量化的高效解决方案,同时完美处理NaN的情况,还能适配你提到的通用形状要求。


步骤1:重塑数组和k值到便于批量处理的形状

因为你的输入数组只要满足axis=-2是1,kths最后两个维度是1,我们可以先把它们都重塑成二维数组(样本数 × 特征数)和一维k值数组,这样不管原数组是几维,都能统一处理:

# 把原数组重塑为(n_samples, n_features),n_samples是所有前面维度的乘积,n_features是最后一维的长度
arr_2d = array_4d.reshape(-1, array_4d.shape[-1])
# 把kths重塑为一维数组,每个元素对应一行的k值
k_vals = kths.reshape(-1)

步骤2:对每行进行排序索引并提取对应kth元素

利用np.argsort对每行排序(NaN会自动被排在末尾),然后用np.take_along根据每行对应的k值提取元素:

# 获取每行的排序索引,NaN会被排到索引的最后位置
sorted_indices = np.argsort(arr_2d, axis=1)
# 把k值转换成列向量,方便和排序索引进行广播匹配
k_col = k_vals.reshape(-1, 1)
# 提取每行对应k值位置的索引
selected_indices = np.take_along(sorted_indices, k_col, axis=1)
# 根据索引提取对应的第k小元素
result_2d = np.take_along(arr_2d, selected_indices, axis=1)

步骤3:把结果重塑回原数组的形状

最后把二维结果数组还原成和原数组前n-1维一致的形状(因为最后一维变成了1):

# 重塑回原数组的前n-1维形状,再加上最后一维的1
result = result_2d.reshape(array_4d.shape[:-1] + (1,))

完整测试代码

把上面的步骤整合到你的示例中,运行后就能得到你想要的结果:

import numpy as np

# 你的原始输入数组
array_4d = np.array(
    [
        [
            [
                [4, 1, np.nan, 20, 11, 12],
            ],
            [
                [np.nan, np.nan, np.nan, np.nan, np.nan, np.nan],
            ],
            [
                [33, 4, 55, 26, 17, 18],
            ],
        ],
        [
            [
                [7, 8, 9, np.nan, 11, 12],
            ],
            [
                [np.nan, np.nan, np.nan, np.nan, np.nan, np.nan],
            ],
            [
                [13, 14, 15, 16, 17, 18],
            ],
        ],
    ]
)

kths = np.array(
    [
        [
            [[1]],
            [[2]],
            [[0]],
        ],
        [
            [[0]],
            [[2]],
            [[1]],
        ],
    ]
)

# 步骤1:重塑
arr_2d = array_4d.reshape(-1, array_4d.shape[-1])
k_vals = kths.reshape(-1)

# 步骤2:排序索引+提取
sorted_indices = np.argsort(arr_2d, axis=1)
k_col = k_vals.reshape(-1, 1)
selected_indices = np.take_along(sorted_indices, k_col, axis=1)
result_2d = np.take_along(arr_2d, selected_indices, axis=1)

# 步骤3:重塑回原形状
result = result_2d.reshape(array_4d.shape[:-1] + (1,))

print(result)

运行这段代码,输出就是你预期的结果:

array([[[[ 4.]],

        [[nan]],

        [[ 4.]]],


       [[[ 7.]],

        [[nan]],

        [[14.]]]])

关于通用性的说明

这个方案完全符合你提到的通用要求:

  • 不管原数组是几维,只要axis=-2是1,reshape(-1, array_4d.shape[-1])都能正确把所有行展平成二维
  • kths只要最后两个维度是1,reshape(-1)就能正确提取每行对应的k值
  • 自动处理全NaN的行,返回NaN;非全NaN的行自动忽略NaN,提取有效元素中的第k小值

备注:内容来源于stack exchange,提问作者Andi

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.14 10:49:36