如何在Numpy中创建可兼容规则/不规则结构的嵌套整数数组?
问题原因
你遇到的报错本质是numpy的自动维度推断逻辑导致的:
- 当你传入长度不一致的嵌套序列调用
np.array(..., dtype=object)时,numpy无法创建高维数组,会直接将每个子序列作为单个object元素存入一维数组,因此取出的test[0]是原生Python列表,可直接作为索引使用。 - 当你传入所有子序列长度一致的嵌套序列时,numpy会自动生成二维object数组,每个元素是单独的Python整型值,取出的
test[0]是一维object dtype的numpy数组,numpy索引不支持object类型的整型数组,因此触发报错。
通用解决方案
不需要提前判断子序列长度是否一致,你可以先创建指定长度的空object一维数组,再直接赋值嵌套序列,强制numpy将每个子序列作为单个元素存储:
import numpy as np def create_index_array(nested_seqs): # 先初始化一维object数组,长度等于子序列的数量 res = np.empty(len(nested_seqs), dtype=object) # 直接赋值所有子序列,numpy不会自动堆叠高维 res[:] = nested_seqs return res
验证效果
用你给出的示例测试:
other_array = np.array([0,1,2,3]) # 不等长子序列场景 test1 = create_index_array([[1,2],[1,2,3]]) print(other_array[test1[0]]) # 输出 [1 2],运行正常 # 等长子序列场景 test2 = create_index_array([[1,2],[1,3]]) print(other_array[test2[0]]) # 输出 [1 2],无报错
可选临时方案
如果不想修改数组创建逻辑,也可以在使用索引时做类型转换:
print(other_array[test2[0].astype(int)])
内容的提问来源于stack exchange,提问作者Combinatorialist
相关产品推荐
相关产品推荐

