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

U型扩散网络上下采样阶段张量维度不匹配导致torch.cat运行时错误

U型扩散网络上下采样阶段张量维度不匹配导致torch.cat运行时错误

看起来你碰到了扩散模型U-Net结构里非常常见的维度对齐坑——上下采样后的张量空间维度(这里是第3个深度维度,25 vs 24)对不上,直接触发了torch.cat的运行时错误。咱们一步步拆解问题、给出修复方案:

问题根源定位

从你打印的张量shape可以直接锁定问题:

  • 下采样阶段保存的残差张量residual_x:torch.Size([4, 64, 25, 64, 64])
  • 上采样阶段当前的特征张量x:torch.Size([4, 64, 24, 64, 64])

torch.cat(dim=1)要求除了拼接的通道维度外,所有其他空间维度的尺寸必须完全一致,但这里深度维度的25和24不匹配,直接触发了RuntimeError。

出现这个差异的核心原因大概率是两个:

  1. 你的下采样模块self.downs和上采样模块self.ups的维度设计没对应上(比如下采样用了步长卷积,上采样的转置卷积参数没匹配,导致奇数维度的计算结果错位);
  2. 输入数据的深度尺寸是25(非2的整数次幂),多轮下采样/上采样后放大了维度偏差。

快速修复方案(代码层面)

在执行torch.cat之前,先把两个张量的空间维度对齐,有两种常用方式,你可以根据需求选择:

方案1:裁剪大张量匹配小张量

把残差张量residual_x的深度维度裁剪到和x一致的24:

for up in self.ups:
    residual_x = residual_inputs.pop()
    # 对齐深度维度:裁剪residual_x到x的尺寸
    residual_x = residual_x[:, :, :x.shape[2], :, :]
    # 确认维度一致后再拼接
    x = torch.cat((x, residual_x), dim=1)
    if feature:
        self.features.append(x.detach().cpu().numpy())
    x = up(x, t)

方案2:上采样小张量匹配大张量

用三线性插值(适合3D数据)把x的深度维度上采样到25,和残差张量对齐:

for up in self.ups:
    residual_x = residual_inputs.pop()
    # 对齐空间维度:上采样x到residual_x的尺寸
    x = torch.nn.functional.interpolate(
        x, 
        size=residual_x.shape[2:],  # 匹配深度、高度、宽度三个维度
        mode='trilinear', 
        align_corners=False
    )
    x = torch.cat((x, residual_x), dim=1)
    if feature:
        self.features.append(x.detach().cpu().numpy())
    x = up(x, t)

根源解决建议(模型设计层面)

如果想从根本上避免这类问题,建议做两个调整:

  1. 调整输入数据尺寸:把输入的深度维度从25改成24或者32(2的整数次幂),这样多轮下采样/上采样后维度会完全对应,不会出现奇数/偶数错位;
  2. 对齐上下采样模块的参数:检查self.downs里的下采样卷积(比如带步长的3D卷积)和self.ups里的上采样转置卷积,确保步长、卷积核大小、padding参数满足转置卷积的输出尺寸公式:output_size = (input_size -1)*stride - 2*padding + kernel_size + output_padding,保证每一组上下采样的维度完全匹配。

额外代码检查点

最后帮你扫一眼forward流程的细节:

  • 确认self.downs和self.ups的模块数量完全一致,否则residual_inputs.pop()会提前耗尽或者有剩余,引发连锁的维度问题;
  • 检查下采样模块是否有池化/步长卷积操作,对应的上采样模块是否有对应的上采样逻辑,避免单方向的维度变化。

备注:内容来源于stack exchange,提问作者user27940251

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.14 09:08:07