如何将shape(20,3)的NumPy数组拆分为3个(20,1)子数组
数组形状理解纠正
shape为(20, 3)的NumPy数组是20行3列的二维数组,你理解为包含20个长度为3的子数组是正确的,不影响后续操作。
原有代码报错原因
你的双层循环写法存在两个核心问题:
- 内层循环每次执行都会重新运行
array_elements = np.zeros(3)初始化数组,之前存储的数值会被直接覆盖,最终只能拿到最后一次循环的赋值结果,这也是你之前用字典存储只能拿到单个值的原因 - 切片逻辑错误:
a[j:]是取数组从第j行到末尾的所有行,后续接[l]是取该切片的第l行,最终得到的是一整行长度为3的序列,把长度为3的序列赋值给单个数组元素位置,就会触发ValueError: setting an array element with a sequence报错
实现方法
不需要编写自定义循环函数,直接用NumPy内置的拆分功能即可,一行代码就能得到你要的包含3个(20,1)形状数组的可索引结构:
import numpy as np a = np.load("你的数组文件路径.npy") # 补全np.load的文件路径参数 # 沿列方向(axis=1)将数组拆分为3个列数组,每个数组shape自动为(20,1) split_result = np.hsplit(a, 3)
使用时直接通过索引取值即可:
split_result[0]:所有子数组第0位元素组成的(20,1)数组split_result[1]:所有子数组第1位元素组成的(20,1)数组split_result[2]:所有子数组第2位元素组成的(20,1)数组
如果想通过显式索引实现(方便理解逻辑),也可以直接按列切片,注意保留二维维度:
# 注意索引写法:列索引用方括号包裹,才会保留(20,1)的二维形状 col_0 = a[:, [0]] col_1 = a[:, [1]] col_2 = a[:, [2]] split_result = [col_0, col_1, col_2]
注意:如果写a[:,0]得到的是shape为(20,)的一维数组,不符合你需要的(20,1)形状要求。
内容的提问来源于stack exchange,提问作者ira_s16
相关产品推荐
相关产品推荐

