如何在PyTorch或NumPy中拼接不同形状张量并补零?
在NumPy和PyTorch中拼接不同形状张量并补零的方法
错误原因解释
你遇到的ValueError是因为np.concatenate要求除拼接轴外的所有维度必须完全匹配。你的两个数组n1(1,64,112,112)和n2(1,512,7,7),除第一个维度外,通道、高、宽维度都不一致,所以直接拼接会失败。解决思路是先将两个张量填充到相同的目标形状(每个维度取两者的最大值),再执行拼接。
NumPy 实现方案
先确定每个维度的目标大小,创建对应尺寸的零数组,将原张量复制到零数组的对应区域,最后拼接:
import numpy as np n1 = np.random.rand(1, 64, 112, 112) n2 = np.random.rand(1, 512, 7, 7) # 计算目标形状:每个维度取两个张量的最大值 target_shape = tuple(max(s1, s2) for s1, s2 in zip(n1.shape, n2.shape)) # 创建零填充数组 padded_n1 = np.zeros(target_shape, dtype=n1.dtype) padded_n2 = np.zeros(target_shape, dtype=n2.dtype) # 将原数组复制到填充数组的对应位置 padded_n1[:n1.shape[0], :n1.shape[1], :n1.shape[2], :n1.shape[3]] = n1 padded_n2[:n2.shape[0], :n2.shape[1], :n2.shape[2], :n2.shape[3]] = n2 # 在通道维度(第1轴)拼接,可根据需求修改axis参数 result = np.concatenate((padded_n1, padded_n2), axis=1) print(result.shape) # 输出 (1, 576, 112, 112)
PyTorch 实现方案
PyTorch可以用torch.nn.functional.pad实现填充,注意pad参数的顺序是从最后一个维度往前数(比如4D张量的pad顺序为:宽度左、宽度右、高度左、高度右、通道左、通道右、批次左、批次右)。
手动计算padding版
import torch t1 = torch.randn(1, 64, 112, 112) t2 = torch.randn(1, 512, 7, 7) # 确定目标形状 target_shape = tuple(max(s1, s2) for s1, s2 in zip(t1.shape, t2.shape)) # 计算t2各维度的填充量,按最后维度到第一维度的顺序组合 pad_w = (0, target_shape[3] - t2.shape[3], 0, target_shape[2] - t2.shape[2]) pad_c = (0, target_shape[1] - t2.shape[1]) pad_b = (0, target_shape[0] - t2.shape[0]) padding = pad_w + pad_c + pad_b # 填充t2 padded_t2 = torch.nn.functional.pad(t2, padding, mode='constant', value=0) # 在通道维度拼接 result = torch.cat((t1, padded_t2), dim=1) print(result.shape) # 输出 torch.Size([1, 576, 112, 112])
通用封装函数版
如果需要处理任意维度的张量,可以封装一个自动计算padding的函数:
def pad_to_match(tensor, target_shape): padding = [] # 从最后一个维度往前遍历,计算每个维度的左右填充量 for i in reversed(range(len(tensor.shape))): diff = target_shape[i] - tensor.shape[i] padding.extend([0, diff]) return torch.nn.functional.pad(tensor, padding, value=0) # 使用函数处理 target_shape = tuple(max(s1, s2) for s1, s2 in zip(t1.shape, t2.shape)) padded_t1 = pad_to_match(t1, target_shape) padded_t2 = pad_to_match(t2, target_shape) result = torch.cat((padded_t1, padded_t2), dim=1)
内容的提问来源于stack exchange,提问作者An Min Su
相关产品推荐
相关产品推荐

