You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch interpolate函数报错:输入输出空间维度不匹配问题求助

问题分析与解决方法

错误原因拆解

  1. 插值参数传错导致维度不匹配
    F.interpolate针对4D输入(N,C,H,W)时,size参数只需要指定空间维度的尺寸(即H,W)。你传入的target.size()[1:]是(1,180,320),包含了目标的通道维度,导致插值函数误以为要把输入的2个空间维度(180,320)改成3个,直接触发空间维度数量不匹配的报错。

  2. 错误压缩类别通道导致损失计算失败
    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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.24 08:16:06