如何为NumPy数组添加维度?将(50,32)数组扩展为(50,32,16,16)
给NumPy数组添加指定大小的第三、第四维度
你需要的不是单纯新增维度(单纯新增会得到(50,32,1,1)),而是要把新增的维度扩展到16的大小,以下是几种实用实现方式:
方法1:用np.newaxis + np.broadcast_to(推荐,内存友好)
np.newaxis用于快速新增维度,np.broadcast_to通过广播机制扩展维度(不复制数据,仅创建视图):
import numpy as np # 初始化原数组 arr = np.random.rand(50, 32) # 新增两个大小为1的维度,得到形状(50,32,1,1) arr_with_new_dims = arr[..., np.newaxis, np.newaxis] # 广播到目标形状(50,32,16,16) result = np.broadcast_to(arr_with_new_dims, (50, 32, 16, 16))
...代表保留原数组的所有现有维度,后续的np.newaxis依次添加第三、第四维度。- 广播后的数组与原数组共享内存,适合不需要修改重复数据的场景。
方法2:用np.expand_dims + np.tile(需要复制数据时使用)
如果需要实际复制数据(比如后续要修改每个位置的元素),可以用np.expand_dims指定维度位置,再用np.tile复制填充:
import numpy as np arr = np.random.rand(50, 32) # 依次在第3、第4维度位置新增大小为1的维度 arr_expanded = np.expand_dims(np.expand_dims(arr, axis=2), axis=3) # 按维度复制:前两个维度不复制,后两个维度各复制16次 result = np.tile(arr_expanded, (1, 1, 16, 16))
np.expand_dims(arr, axis=2)表示在索引为2的位置(第三维)新增维度,同理axis=3对应第四维。np.tile的参数是每个维度的复制次数,(1,1,16,16)保证原数组的前两个维度不变,后两个维度扩展到16。
关于np.newaxis和np.expand_dims的说明
两者本质是等价的:
arr[..., np.newaxis]等价于np.expand_dims(arr, axis=-1)(-1表示最后一个位置)np.newaxis写法更简洁,适合手动指定维度;np.expand_dims适合动态传入axis参数的编程场景。
内容的提问来源于stack exchange,提问作者Prometheus
相关产品推荐
相关产品推荐

