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

UNet模型Concatenate层维度不匹配报错求助

Keras构建UNet时Concatenate层维度不匹配错误

基于keras.Model构建UNet模型时,调用Concatenate层触发InvalidArgumentError。明明打印出拼接的两个输入形状一致,但报错显示二者维度不匹配。

相关代码

class UNet(keras.Model):
  def __init__(self, shape=(572, 572, 1), **kwargs):
    self.concat = keras.layers.Concatenate(axis=-1) # 沿深度维度拼接
    ...

  class CONV2_BLOCK(keras.layers.Layer):
    ...      

  class CONV_T(keras.layers.Layer):
    def __init__(self, filters, **kwargs):
      super().__init__(**kwargs)
      self.conv_t = keras.layers.Conv2DTranspose(filters=filters, kernel_size=2, strides=2)

    def call(self, inputs):
      outputs = self.conv_t(inputs)
      return outputs

  class CROP(keras.layers.Layer):
    def __init__(self, cropping, **kwargs):
      super().__init__(**kwargs)
      self.cropping = cropping
      self.crop = keras.layers.Cropping2D(cropping=self.cropping)

    def call(self, inputs):
      outputs = self.crop(inputs)
      return outputs

  def call(self, inputs):
    # self.conv_arr = [64, 128, 256, 512, 1024]
    # self.crop_arr = [4, 17, 40, 88] 从下到上的裁剪量

    x1 = self.CONV2_BLOCK(filters=64)(inputs)
    print(x1.shape)
    x = self.maxpool(x1)
    print(x.shape)
    ...

    x = self.CONV2_BLOCK(filters=1024)(x)
    print(x.shape)

    print(f"convt shape{self.CONV_T(filters=512)(x).shape}")
    print(f"crop shape{self.CROP(cropping=4)(x4).shape}")
    x = self.concat([self.CONV_T(filters=512)(x), self.CROP(cropping=4)(x4)])
    x = self.CONV2_BLOCK(filters=512)(x)

    ...

    x = self.concat([self.CONV_T(filters=64)(x), self.CROP(cropping=88)(x1)])
    x = self.CONV2_BLOCK(filters=64)(x)

    outputs = self.conv_sz1(x)

    return outputs

打印输出

打印输出:

conv_t shape(2, 56, 56, 512)
crop shape(2, 56, 56, 512)

报错信息

报错信息:

-->83 x = self.concat([self.CONV_T(filters=216)(x), self.CROP(cropping=17)(x3)])
    84 x = self.CONV2_BLOCK(filters=256)(x)

Dimension 1 in both shapes must be equal: shape[0] = [2,104,104,216] vs. shape[1] = [2,102,102,256] [Op:ConcatV2] name: concat


解决方案

核心问题

你打印的是临时创建的层实例的输出形状,但实际拼接时又重新创建了新的层实例。如果CONV2_BLOCK的卷积参数(如padding)、裁剪数值计算有误,新实例的输出形状会和打印值不一致。

修复步骤

  1. 提前初始化所有层实例
    在__init__方法中创建好所有需要的CONV_T和CROP层,避免在call中重复创建:

    class UNet(keras.Model):
      def __init__(self, shape=(572, 572, 1), **kwargs):
        super().__init__(**kwargs)
        self.concat = keras.layers.Concatenate(axis=-1)
        # 提前创建所有转置卷积和裁剪层
        self.conv_t512 = self.CONV_T(filters=512)
        self.crop4 = self.CROP(cropping=4)
        self.conv_t256 = self.CONV_T(filters=256)  # 注意与conv_arr对应,报错中误写为216
        self.crop17 = self.CROP(cropping=17)
        # 其他层同理初始化
        ...
    
  2. 复用已创建的层实例
    在call方法中使用提前初始化的实例,确保打印和实际拼接用的是同一层:

    def call(self, inputs):
      ...
      # 打印同一实例的形状
      print(f"convt shape{self.conv_t512(x).shape}")
      print(f"crop shape{self.crop4(x4).shape}")
      x = self.concat([self.conv_t512(x), self.crop4(x4)])
      ...
      # 报错行改为复用实例
      x = self.concat([self.conv_t256(x), self.crop17(x3)])
      ...
    
  3. 重新计算裁剪数值
    UNet的裁剪尺寸需要匹配下采样过程的特征图收缩量:每次maxpool使尺寸减半,两次无padding的3x3卷积会让每个维度减少4(输入尺寸-2*2)。重新核对crop_arr的数值,确保裁剪后特征图与转置卷积输出尺寸完全一致。

  4. 检查CONV2_BLOCK的padding设置
    确保CONV2_BLOCK中的卷积使用padding='valid'(默认值),如果用padding='same'会导致特征图尺寸不收缩,与UNet原始结构冲突,引发尺寸不匹配。


内容的提问来源于stack exchange,提问作者user19676560

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 22:17:44