不同形状矩阵拼接补零实现:如何将a、b、c拼接为(2,4,3)规格的矩阵
实现方案
我们可以用numpy的维度扩展和填充功能完成需求,具体步骤如下:
步骤1:导入依赖并定义输入矩阵
import numpy as np a = np.array([[1, 2, 3, 4]]) b = np.array([[1, 2, 3, 4], [1, 2, 3, 4]]) c = np.array([[[1, 2, 3], [4, 5, 6], [7, 8, 9]]])
步骤2:统一矩阵维度为3维
三个矩阵当前维度分别是2、2、3,我们先把低维矩阵在末尾扩展维度,统一为3维:
def to_3d(arr): while arr.ndim < 3: arr = np.expand_dims(arr, axis=-1) return arr a_3d = to_3d(a) # 转换后shape为 (1, 4, 1) b_3d = to_3d(b) # 转换后shape为 (2, 4, 1) c_3d = to_3d(c) # 本身就是3维,shape为 (1, 3, 3)
步骤3:计算目标填充维度
按需求取各维度的最大值作为目标shape:
target_shape = ( max(a_3d.shape[0], b_3d.shape[0], c_3d.shape[0]), max(a_3d.shape[1], b_3d.shape[1], c_3d.shape[1]), max(a_3d.shape[2], b_3d.shape[2], c_3d.shape[2]), ) # 计算得到目标shape为 (2, 4, 3),和预期一致
步骤4:对所有矩阵补0对齐到目标shape
我们默认在每个维度的末尾填充0,如果你有前后填充的需求可以自行调整pad_width的参数:
def pad_arr(arr, target_shape): pad_width = [] for dim_idx in range(3): diff = target_shape[dim_idx] - arr.shape[dim_idx] # 元组第一个元素是维度前填充的长度,第二个是维度后填充的长度 pad_width.append((0, diff)) return np.pad(arr, pad_width=pad_width, mode='constant', constant_values=0) a_padded = pad_arr(a_3d, target_shape) b_padded = pad_arr(b_3d, target_shape) c_padded = pad_arr(c_3d, target_shape)
步骤5:合并矩阵(按需选择)
如果你需要把三个矩阵拼接为单个矩阵,可以根据业务场景选择对应的合并方式:
- 如果需要新增维度堆叠三个矩阵,输出shape为
(3, 2, 4, 3):stacked = np.stack([a_padded, b_padded, c_padded], axis=0) - 如果需要沿最后一个维度拼接,输出shape为
(2, 4, 9):concatenated = np.concatenate([a_padded, b_padded, c_padded], axis=-1) - 如果需要按位置合并为单个
(2,4,3)的矩阵,比如直接相加:merged = a_padded + b_padded + c_padded
内容的提问来源于stack exchange,提问作者Eduardo Morales
相关产品推荐
相关产品推荐

