TensorFlow生成器适配112×112输入时Concatenate层报错问题
报错原因
你使用的U-Net结构默认适配输入尺寸为2的整数次幂(如256=2^8):8次步长为2的下采样可将256x256输入降至1x1,搭配后续7次上采样和最后1层转置卷积,刚好能还原为256x256输出。
将输入调整为112x112后,由于112不是2的整数次幂,多次下采样后的特征尺寸和上采样阶段的特征尺寸无法匹配,导致Concatenate层拼接时出现长宽维度的冲突,你收到的报错就是第一次上采样输出2x2特征时,对应的跳连特征尺寸为1x1,无法直接拼接。
解决方案
提供两种可直接落地的修改方案:
方案1:调整网络层数量适配112x112输入
你使用的downsample(步长2的卷积)和upsample(步长2的转置卷积)默认每调用一次,特征长宽就会缩放一倍,112x112输入最多支持6次有效下采样,无需原有8层下采样结构。
修改点如下:
- 移除
down_stack最后1层下采样,保留前7层 - 移除
up_stack最后1层上采样,保留前6层 - 调整最后一层转置卷积的配置,确保输出尺寸符合你的需求
修改后的核心代码参考:
def Generator(): inputs = tf.keras.layers.Input(shape=[112, 112, 3]) down_stack = [ downsample(64, 4, apply_batchnorm=False), # (batch_size, 56, 56, 64) downsample(128, 4), # (batch_size, 28, 28, 128) downsample(256, 4), # (batch_size, 14, 14, 256) downsample(512, 4), # (batch_size, 7, 7, 512) downsample(512, 4), # (batch_size, 4, 4, 512) downsample(512, 4), # (batch_size, 2, 2, 512) downsample(512, 4), # (batch_size, 1, 1, 512) ] up_stack = [ upsample(512, 4, apply_dropout=True), # (batch_size, 2, 2, 1024) upsample(512, 4, apply_dropout=True), # (batch_size, 4, 4, 1024) upsample(512, 4, apply_dropout=True), # (batch_size, 8, 8, 1024) upsample(512, 4), # (batch_size, 16, 16, 1024) upsample(256, 4), # (batch_size, 32, 32, 512) upsample(128, 4), # (batch_size, 64, 64, 256) ] initializer = tf.random_normal_initializer(0., 0.02) last = tf.keras.layers.Conv2DTranspose(OUTPUT_CHANNELS, 4, strides=2, padding='same', kernel_initializer=initializer, activation='tanh') # (batch_size, 112, 112, 3) x = inputs # Downsampling through the model skips = [] for down in down_stack: x = down(x) skips.append(x) skips = reversed(skips[:-1]) # Upsampling and establishing the skip connections for up, skip in zip(up_stack, skips): x = up(x) x = tf.keras.layers.Concatenate()([x, skip]) x = last(x) return tf.keras.Model(inputs=inputs, outputs=x)
方案2:拼接前做尺寸对齐,无需调整网络层数
如果不想改动原有网络结构,直接在跳连拼接前增加尺寸对齐逻辑即可,改动量最小,适合快速验证效果:
仅需要修改跳连部分的代码:
# Upsampling and establishing the skip connections for up, skip in zip(up_stack, skips): x = up(x) # 新增尺寸对齐逻辑 if x.shape[1] != skip.shape[1] or x.shape[2] != skip.shape[2]: skip = tf.keras.layers.Resizing(height=x.shape[1], width=x.shape[2])(skip) x = tf.keras.layers.Concatenate()([x, skip])
内容的提问来源于stack exchange,提问作者Hazem
相关产品推荐
相关产品推荐

