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)、裁剪数值计算有误,新实例的输出形状会和打印值不一致。
修复步骤
提前初始化所有层实例
在__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) # 其他层同理初始化 ...复用已创建的层实例
在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)]) ...重新计算裁剪数值
UNet的裁剪尺寸需要匹配下采样过程的特征图收缩量:每次maxpool使尺寸减半,两次无padding的3x3卷积会让每个维度减少4(输入尺寸-2*2)。重新核对crop_arr的数值,确保裁剪后特征图与转置卷积输出尺寸完全一致。检查CONV2_BLOCK的padding设置
确保CONV2_BLOCK中的卷积使用padding='valid'(默认值),如果用padding='same'会导致特征图尺寸不收缩,与UNet原始结构冲突,引发尺寸不匹配。
内容的提问来源于stack exchange,提问作者user19676560

