3D ResUNet分割CT图像张量尺寸不匹配报错如何解决
问题根因
错误来源于3D UNet跳跃连接拼接阶段,编码器输出特征和对应解码器上采样特征的空间维度不匹配。你采用的步长为2的下采样卷积会对奇数维度向下取整,而步长2的转置卷积上采样会直接将尺寸翻倍,二者累计就会出现1像素的尺寸差,也就是你遇到的55和54的维度差异。
解决方案
提供两个可直接落地的方案,优先推荐方案1,无需调整网络结构和输入尺寸:
方案1:拼接前动态对齐特征尺寸(最省事)
在forward函数中每次拼接跳跃连接之前,统一对齐两个特征的空间维度,可选择对尺寸小的特征做边缘填充,或对尺寸大的特征做中心裁剪。
你只需要新增一个通用尺寸对齐函数,再修改3处拼接逻辑即可:
首先定义对齐函数(可放在ResUNet类内部或外部,注意导入torch.nn.functional as F):
def align_tensor_size(source, target): # source为待对齐张量,target为目标尺寸张量,仅对齐后三个空间维度 diff_d = target.size(2) - source.size(2) diff_h = target.size(3) - source.size(3) diff_w = target.size(4) - source.size(4) # PyTorch pad顺序为[w前补, w后补, h前补, h后补, d前补, d后补] source = F.pad(source, [ diff_w//2, diff_w - diff_w//2, diff_h//2, diff_h - diff_h//2, diff_d//2, diff_d - diff_d//2 ]) return source
然后修改forward中的三处拼接逻辑:
- 报错行修改为:
short_range6_aligned = align_tensor_size(short_range6, long_range3) outputs = self.decoder_stage2(torch.cat([short_range6_aligned, long_range3], dim=1)) + short_range6_aligned
- 第二处拼接修改为:
short_range7_aligned = align_tensor_size(short_range7, long_range2) outputs = self.decoder_stage3(torch.cat([short_range7_aligned, long_range2], dim=1)) + short_range7_aligned
- 第三处拼接修改为:
short_range8_aligned = align_tensor_size(short_range8, long_range1) outputs = self.decoder_stage4(torch.cat([short_range8_aligned, long_range1], dim=1)) + short_range8_aligned
如果你使用Dice损失等对边缘敏感的损失函数,可将对齐逻辑改为裁剪大尺寸张量到小尺寸,避免填充0值影响损失计算。
方案2:调整输入尺寸为2的n次幂倍数
当前网络共3次步长2的下采样,所以输入的三个空间维度(D、H、W)都需要是2^3=8的整数倍,你可以将输入的三个维度填充或裁剪到最近的8的倍数,下采样和上采样后尺寸就会完全匹配。
额外优化提示
代码中的dropout参数名拼写错误,你写的是dorp_rate,不影响运行但建议修正,避免后续维护混淆。
内容的提问来源于stack exchange,提问作者dzhang
相关产品推荐
相关产品推荐

