PyTorch interpolate函数报错:输入输出空间维度不匹配问题求助
问题分析与解决方法
错误原因拆解
插值参数传错导致维度不匹配
F.interpolate针对4D输入(N,C,H,W)时,size参数只需要指定空间维度的尺寸(即H,W)。你传入的target.size()[1:]是(1,180,320),包含了目标的通道维度,导致插值函数误以为要把输入的2个空间维度(180,320)改成3个,直接触发空间维度数量不匹配的报错。错误压缩类别通道导致损失计算失败
PyTorch的CrossEntropyLoss要求模型输出(logits)必须保留类别通道:形状为(N,C,H,W)(C是类别数,这里是36)。你用squeeze(dim=1)把模型输出的通道维度去掉后,张量变成(N,H,W),损失函数无法识别每个位置的类别概率分布,自然报错。
修正后的代码
情况1:模型输出空间尺寸已和目标一致(无需插值)
如果你的模型输出out的空间尺寸180×320已经和目标匹配,直接跳过插值步骤:
# 保留模型输出的类别通道,形状(5,36,180,320) outs = out # 压缩目标的通道维度,得到(5,180,320),也可以直接传原target(损失会自动忽略通道) target_squeezed = target.squeeze(dim=1) crit_loss = crit(outs, target_squeezed) loss += (loss_coeff * crit_loss)
情况2:模型输出空间尺寸和目标不一致(需要插值)
如果确实需要下采样,正确指定插值的空间维度:
# 取目标的空间维度尺寸(H,W),即target.size()[2:] outs = F.interpolate(out, size=target.size()[2:], mode='bilinear', align_corners=False) # 此时outs形状仍为(5,36,180,320),保留类别通道 target_squeezed = target.squeeze(dim=1) crit_loss = crit(outs, target_squeezed) loss += (loss_coeff * crit_loss)
关键注意点
CrossEntropyLoss不需要模型输出和目标的通道数一致:模型输出是N×C×H×W(每个位置对应C个类别的概率logit),目标是N×H×W(每个位置对应类别索引0-35)。- 目标如果是
N×1×H×W形状,不需要手动squeeze,损失函数会自动忽略单通道维度。
内容的提问来源于stack exchange,提问作者Varun
相关产品推荐
相关产品推荐

