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

Keras实现类UNet架构时Concatenate拼接层维度不兼容问题求助

Keras实现类UNet生成器时Concatenate层维度不匹配的解决方法

问题描述

刚开始基于Keras开展神经网络开发工作,尝试复现Sim、Oh、Kim、Jung在论文《Optimal Transport driven CycleGAN for Unsupervised Learning in Inverse Problems》图10中给出的类UNet架构生成器模型,对应实现代码如下:

def def_generator(image_shape=(256,256,3)):
  init= RandomNormal(stddev=0.02)
  #Start 1st Block
  in_image = Input(shape=image_shape)
  g1=Conv2D(64,(3,3))(in_image)
  g1=InstanceNormalization(axis=-1)(g1)
  g1=LeakyReLU(alpha=0.2)(g1)
  g1=Conv2D(64,(3,3))(g1)
  g1=InstanceNormalization(axis=-1)(g1)
  g1=LeakyReLU(alpha=0.2)(g1)
  #End of 1st Block
  #Start of 2nd Block
  g2=MaxPool2D()(g1)
  g2=Conv2D(128,(3,3))(g2)
  g2=InstanceNormalization(axis=-1)(g2)
  g2=LeakyReLU(alpha=0.2)(g2)
  g2=Conv2D(128,(3,3))(g2)
  g2=InstanceNormalization(axis=-1)(g2)
  g2=LeakyReLU(alpha=0.2)(g2)
  #End of 2nd Block
  #Start of 3rd Block
  g3=MaxPool2D()(g2)
  g3=Conv2D(256,(3,3))(g3)
  g3=InstanceNormalization(axis=-1)(g3)
  g3=LeakyReLU(alpha=0.2)(g3)
  g3=Conv2D(256,(3,3))(g3)
  g3=InstanceNormalization(axis=-1)(g3)
  g3=LeakyReLU(alpha=0.2)(g3)
  #End of 3rd Block
  #Start of 4th block
  g4=MaxPool2D()(g3)
  g4=Conv2D(512,(3,3))(g4)
  g4=InstanceNormalization(axis=-1)(g4)
  g4=LeakyReLU(alpha=0.2)(g4)
  g4=Conv2D(512,(3,3))(g4)
  g4=InstanceNormalization(axis=-1)(g4)
  g4=LeakyReLU(alpha=0.2)(g4)
  g4=Conv2D(256,(3,3))(g4)
  g4=InstanceNormalization(axis=-1)(g4)
  g4=LeakyReLU(alpha=0.2)(g4)
  g4=Conv2DTranspose(256,(2,2),strides=(4,4),output_padding=1)(g4)
  #End of 4th Block
  #Start of 5th Block
  g5input=Concatenate()([g4,g3])
  g5=Conv2D(256,(3,3))(g5input)
  g5=InstanceNormalization(axis=-1)(g5)
  g5=LeakyReLU(alpha=0.2)(g5)
  g5=Conv2D(256,(3,3))(g5)
  g5=InstanceNormalization(axis=-1)(g5)
  g5=LeakyReLU(alpha=0.2)(g5)
  g5=Conv2DTranspose(128,(2,2),strides=(3,3), padding='same', output_padding=0)(g5)
  #End of 5th Block
  #Start of 6th block
  g6input=Concatenate()([g5,g2])
  g6=Conv2D(128,(2,2))(g6input)
  g6=InstanceNormalization(axis=-1)(g6)
  g6=LeakyReLU(alpha=0.2)(g6)
  g6=Conv2D(128,(2,2))(g6)
  g6=InstanceNormalization(axis=-1)(g6)
  g6=LeakyReLU(alpha=0.2)(g6)
  g6=Conv2DTranspose(64,(2,2),strides=(2,2), padding='valid', output_padding=1)(g6)
  #End of 6th Block
  #Start of 7th block
  g7input=Concatenate()([g6,g1])
  g7=Conv2D(64,(2,2))(g7input)
  g7=InstanceNormalization(axis=-1)(g7)
  g7=LeakyReLU(alpha=0.2)(g7)
  g7=Conv2D(64,(2,2))(g7)
  g7=InstanceNormalization(axis=-1)(g7)
  g7=LeakyReLU(alpha=0.2)(g7)
  g7=Conv2DTranspose(1,(1,1))(g7)
  
  model=Model(in_image, g5)
  model.compile(loss='mse', optimizer=Adam(lr=2e-4,beta_1=0.5), loss_weights=[0.5], metrics=['accuracy'])
  return model

g=def_generator((120,120,1))
print(g.summary())

代码运行时始终报错,提示需要执行Concatenate拼接操作的对应层维度不兼容。已知该问题由前期MaxPool2D与Conv2D操作带来的特征图尺寸变化导致,需要通用的技巧或实现策略,规避或减少这类维度不匹配问题。

通用规避策略

  • 所有卷积、池化、转置卷积层统一加padding='same'参数。默认padding='valid'会在卷积时丢弃边缘像素,每做一次3*3 valid卷积,特征图宽高就会减2,几次下采样后尺寸偏差会被放大,直接导致跳连拼接时尺寸对不上。用same padding可以保证卷积操作不改变特征图宽高,只有步长不为1的层(池化、转置卷积)会调整尺寸,尺寸计算逻辑会简单很多。
  • 输入尺寸固定为2的n次幂相关数值。UNet类架构有多次2倍下采样、2倍上采样,输入尺寸选256256、128128、512*512这类能被2^下采样次数整除的数值,从根源上避免上采样后出现奇数尺寸、和下采样特征图差1-2个像素的问题。如果必须用非标准尺寸,优先在输入层后加ZeroPadding2D把尺寸补到符合整除要求,输出前再用Cropping2D裁回目标尺寸。
  • 不要手动硬转置卷积的步长、output_padding参数凑尺寸。上述代码里g4层用strides=(4,4)、g5层用strides=(3,3)是非常容易出问题的写法,标准UNet的上采样统一用2倍上采样,对应转置卷积步长固定为2,每上采样一次尺寸刚好放大2倍,和对应层级下采样的特征图尺寸完全匹配,不需要反复调output_padding。
  • 拼接前主动做尺寸对齐。如果已经出现1-2个像素的尺寸差,不需要回去改所有层参数,在Concatenate前加一个Cropping2D层裁掉上采样输出多出来的边缘像素,或者加ZeroPadding2D给尺寸小的特征图补边,把两个要拼接的特征图宽高调成完全一致就行。
  • 逐层打印特征图尺寸调试。不要写完整个模型再运行,每写完一个下采样/上采样块就打印当前输出的shape,比如print(g1.shape, g2.shape, g3.shape),第一时间发现尺寸偏差在哪一层产生,不要等拼接报错了再回头逐层排查。
  • 修正模型输出定义。上述代码里model=Model(in_image, g5)是明显的笔误,最终输出应该是最后一层g7,不然模型构建阶段就会出现逻辑错误。

针对贴出的这段代码,最核心的问题有两个:一是所有Conv2D都没加padding='same',卷积过程持续缩小特征图尺寸;二是上采样步长乱设,没有和下采样的2倍缩放逻辑对齐,把这两点改完90%的维度报错都会消失。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.01 20:45:53