如何沿Numpy数组第2轴拆分,生成指定形状的数组列表?
解决方法:沿指定轴拆分Numpy数组
你遇到的问题是因为直接用list(a)时,Numpy默认会沿着轴0(第一个维度)拆分数组,所以得到的是6个形状为(150,29,29,29,1)的数组,和你想要的结果正好相反。
要实现沿轴1(第二个维度,长度150)拆分并得到150个(6,29,29,29,1)的数组,有两种简单的方法:
方法1:列表推导式直接切片(最直观)
直接遍历轴1的每个索引,取出对应切片即可,不需要额外处理维度:
import numpy as np # 示例数组 original_arr = np.random.rand(6, 150, 29, 29, 29, 1) # 沿轴1拆分 split_list = [original_arr[:, idx, :, :, :, :] for idx in range(150)] # 验证形状 print(split_list[0].shape) # 输出 (6, 29, 29, 29, 1)
方法2:使用np.split配合squeeze
np.split可以指定拆分的轴和份数,但拆分后每个子数组会保留长度为1的轴(比如(6,1,29,29,29,1)),需要用np.squeeze去掉这个多余的轴:
split_arrays = np.split(original_arr, 150, axis=1) split_list = [np.squeeze(arr, axis=1) for arr in split_arrays] # 验证形状 print(split_list[0].shape) # 输出 (6, 29, 29, 29, 1)
为什么这两种方法有效?
原数组的维度顺序是(轴0:6, 轴1:150, 轴2:29, 轴3:29, 轴4:29, 轴5:1),我们需要拆分的是轴1(第二个维度):
- 列表推导式中
[:, idx, :, :, :, :]表示保留轴0的所有元素,取轴1的第idx个元素,保留后面所有轴的元素,正好得到目标形状。 np.split(original_arr, 150, axis=1)会把轴1平均拆分成150份,每份长度为1,再用squeeze(axis=1)去掉这个长度为1的轴,就得到了想要的形状。
如果你的数组轴1长度不是刚好能被拆分份数整除,np.split会报错,这时候列表推导式的方法更稳妥,因为它不依赖维度长度的整除性。
内容的提问来源于stack exchange,提问作者ssm
相关产品推荐
相关产品推荐

