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

从NumPy ndarray中高效提取多组元素的实现方法

高效处理NumPy中不规则嵌套索引列表的解决方案

核心问题

你的索引列表b是长度不规则的嵌套结构,直接用普通NumPy索引会因维度不匹配报错,而列表推导的Python循环开销大,需要用向量级别的NumPy原生操作来提升效率。


最优解决方案:扁平化索引+分割

通过将嵌套索引扁平化处理,用一次向量索引获取所有元素,再按原结构分割,全程无Python循环,效率拉满:

import numpy as np

a = np.array([1,2,3])
b = [[0,2],[1]]

# 1. 扁平化索引并批量取值(向量操作,无循环)
flat_indices = np.concatenate(b)
flat_result = a[flat_indices]

# 2. 计算分割点,还原嵌套结构
split_lengths = [len(sub_idx) for sub_idx in b]
split_points = np.cumsum(split_lengths)[:-1]  # 去掉最后一个分割点(避免空数组)
c = np.split(flat_result, split_points)

# 输出结果:[array([1, 3]), array([2])]
# 如果需要纯Python列表,可追加:c = [arr.tolist() for arr in c]

为什么之前的方法失败?

  • a[b]报错:NumPy会将二维的b视为对二维数组的索引,但a是一维数组,维度不匹配导致索引过多错误。
  • tile后索引报错:np.tile(a, (2,1))生成的是(2,3)数组,b中的索引2对应数组的第0轴(长度为2),超出边界所以报错。

效率说明

  • 列表推导[a[b_] for b_ in b]:每次循环都会触发一次NumPy索引操作,Python循环的开销在数据量大时会被放大。
  • 上述方案:所有核心操作都是NumPy的向量级运算,完全规避Python循环,数据量越大,效率优势越显著。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.20 19:05:06