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

如何用矢量化操作扩展多维数组元素索引查找代码?

多维数组中按指定值列表匹配对应索引的矢量化实现

问题背景

我有一段处理一维数组的代码,能找到input列表中与values列表元素匹配的索引,输出顺序和values一致:

import numpy as np

input = [1, 2, 8, 7, 3, 4, 6, 5, 9]
values = [4, 8, 3]

match_index_lst, match_index_values = np.where(np.array(input) == np.array(values)[:,None])
output_indice_lst = match_index_values[np.argsort(match_index_lst)]
# 输出: [5, 2, 4]

现在需要将这段代码矢量化扩展,适配维度为[a,b,c]的多维数组(原代码中input是一维的c)。比如当输入是(2,2,8)维度的数组时,要输出对应结构的索引数组,示例如下:

示例输入

import numpy as np

input = [[[[ 0.31, 1.56, 1.58, 0.16, 0.22, 0.54, 0.98, 0.35 ]],
          [[ 0.77, 2.62, 0.44, 0.08, 0.76, 0.87, 0.88, 0.51 ]]],

         [[[ 1.14, 0.48, 1.09, 0.93, 0.47, 0.13, 0.75, 0.19 ]],
          [[ 1.15, 0.17, 2.33, 0.46, 0.30, 2.60, 0.79, 1.07 ]]]]

values = [[[[ 0.54, 1.58 ]],
           [[ 0.77, 0.88 ]]],

          [[[ 0.48, 1.09 ]],
           [[ 2.60, 2.33 ]]]]

期望输出

[[[[ 5, 2 ]],
  [[ 0, 6 ]]],

 [[[ 1, 2 ]],
  [[ 5, 2 ]]]]

我试过扁平化数组后操作,但无法恢复正确顺序和结构;用循环实现虽然简单,但对速度要求高,希望用矢量化操作解决。


矢量化解决方案

核心思路是利用numpy的广播机制,在保持多维结构的前提下,对每个[a,b]位置的子数组,分别匹配对应values中同位置的元素,最后整理索引结构。

import numpy as np

# 转换为numpy数组,明确维度
input_arr = np.array(input)
values_arr = np.array(values)

# 获取维度信息:input_arr维度为(a,b,1,c),values_arr为(a,b,1,k)(k是每个子组的匹配值数量)
a, b, _, c = input_arr.shape
_, _, _, k = values_arr.shape

# 扩展维度实现广播匹配:将input_arr扩展为(a,b,1,1,c),values_arr扩展为(a,b,1,k,1)
matches = input_arr[:, :, :, np.newaxis, :] == values_arr[:, :, :, :, np.newaxis]

# 扁平化匹配矩阵,方便提取索引
matches_reshaped = matches.reshape(-1, c)

# 提取匹配的行列索引,按行排序保证每个values元素的索引对应正确
row_indices, col_indices = np.where(matches_reshaped)
sorted_indices = col_indices[np.argsort(row_indices)]

# 重塑回原多维结构
output = sorted_indices.reshape(a, b, 1, k)

print(output)

代码说明

  1. 维度对齐:通过np.newaxis扩展维度,让input_arr和values_arr在广播时,能逐个位置对应子数组进行元素匹配。
  2. 匹配矩阵生成:广播后的比较会生成(a,b,1,k,c)的布尔矩阵,标记每个values元素在对应input子数组中的位置。
  3. 索引提取与排序:将匹配矩阵扁平化后,用np.where提取匹配的行列索引,排序后确保索引顺序和values的元素顺序一致。
  4. 结构恢复:将提取的索引重塑回与values一致的多维结构,得到目标输出。

注意事项

  • 如果input子数组中存在重复的values元素,代码会返回第一个匹配的索引;若需获取所有匹配索引,可调整np.where后的处理逻辑。
  • 需确保input和values的维度结构完全对应(如示例中的(a,b,1,c)和(a,b,1,k)),否则需先调整维度对齐。

内容的提问来源于stack exchange,提问作者Tim

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 17:22:31