循环中沿第5维度拼接numpy矩阵失败,求解决方案
解决Numpy数组沿指定维度拼接的问题
嘿,我明白你的问题了——你尝试用np.stack来拼接多个形状为(1, 2, 128, 30, 3)的临时数组,想要最终得到(1, 2, 128, 30, 600)的结果,但几次迭代后就失败了,问题出在函数的选择上!
为什么np.stack会失败?
np.stack的作用是在一个新的维度上堆叠数组,而不是在已有的维度上扩展。比如你第一次把Total初始化为temp(形状(1,2,128,30,3)),然后执行Total = np.stack((Total, temp), axis=5),这会新增一个第6维(索引5),得到的形状会变成(1,2,128,30,3,2)——这完全不是你想要的结构。后续迭代会继续新增维度,很快形状就会彻底混乱,自然会报错。
正确的解决方案:使用np.concatenate
你需要的是在**已有的第5个维度(对应Numpy的索引4,因为维度从0开始计数)**上拼接数组,这时候应该用np.concatenate:
初始化Total:第一次迭代时直接把
Total赋值为第一个temp:import numpy as np # 假设第一次生成的temp temp = np.random.rand(1, 2, 128, 30, 3) Total = temp # 初始形状 (1,2,128,30,3)循环拼接:后续每次迭代生成
temp后,用np.concatenate在axis=4的位置拼接:for _ in range(199): # 因为已经初始化了1次,还需要199次 temp = np.random.rand(1, 2, 128, 30, 3) Total = np.concatenate([Total, temp], axis=4)完成后
Total的形状就是(1,2,128,30,600),正好符合你的需求。
更高效的优化:预分配内存
如果迭代次数很多(200次),每次concatenate都会重新分配内存,效率较低。你可以提前预分配好最终形状的数组,然后把每个temp放到对应的位置:
# 预分配数组,指定和temp一致的数据类型 Total = np.zeros((1, 2, 128, 30, 600), dtype=temp.dtype) for i in range(200): temp = np.random.rand(1, 2, 128, 30, 3) # 将temp放到Total的第i个3段位置 Total[..., i*3 : (i+1)*3] = temp
这种方式避免了多次内存拷贝,运行速度会更快。
内容的提问来源于stack exchange,提问作者Shinobii
相关产品推荐
相关产品推荐

