PyTorch调用torch.cat拼接张量时报尺寸不匹配错误如何解决
报错原因
torch.cat的规则是:参与拼接的所有张量,除了指定的拼接维度外,其余所有维度的尺寸必须完全一致。
你现在执行的是torch.cat((x4,x3,x2,x1,x),dim=1),也就是在第1维(通道维度)拼接,要求所有张量的第0维(batch)、第2维(高度)、第3维(宽度)尺寸完全相等。
从你打印的张量尺寸可以看到:
- x1/x2/x3/x4的尺寸都是
[5, C, 32, 32],高宽为32*32 - aspp输出的x尺寸为
[5, 256, 16, 16],高宽为16*16
两者的高、宽维度尺寸都不匹配,因此触发报错。
解决方案
你可以根据你的网络设计需求,选择以下任意一种方案对齐张量尺寸:
方案1:上采样aspp输出的x到32*32
如果需要保留浅层特征的高分辨率,就把低分辨率的x上采样到和x1~x4一致:
x = self.aspp(x) # 双线性插值上采样到32*32 x = torch.nn.functional.interpolate(x, size=(32, 32), mode='bilinear', align_corners=False) x = torch.cat((x4,x3,x2,x1,x), dim=1)
方案2:下采样x1~x4到16*16
如果你的网络后续需要低分辨率特征降低计算量,就把高分辨率的x1~x4下采样到和x一致:
x1 = self.avg_pool(l1) x2 = self.avg_pool(l2) x3 = self.avg_pool(l3) x4 = self.avg_pool(l4) # 平均池化下采样到16*16 x1 = torch.nn.functional.avg_pool2d(x1, kernel_size=2, stride=2) x2 = torch.nn.functional.avg_pool2d(x2, kernel_size=2, stride=2) x3 = torch.nn.functional.avg_pool2d(x3, kernel_size=2, stride=2) x4 = torch.nn.functional.avg_pool2d(x4, kernel_size=2, stride=2) x = self.aspp(x) x = torch.cat((x4,x3,x2,x1,x), dim=1)
方案3:调整ASPP层参数
如果ASPP输出的1616不符合你的设计预期,你可以修改ASPP层的空洞卷积padding、步长参数,让它输出特征的高宽保持为3232,无需额外做采样操作。
内容的提问来源于stack exchange,提问作者Priyanka Mishra
相关产品推荐
相关产品推荐

