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

使用Keras Sequential实现GAN图像上色时的维度不匹配报错

问题:GAN图像上色任务中的维度匹配错误

需求背景

基于少量ImageNet狗图像,用Keras Sequential构建GAN实现黑白图像上色。

处理流程

  • 将./ImageNet/dogs/中的狗图像转为224×224分辨率,保存至./ImageNet/dogs_lowres/(注:原描述为244,推测是笔误,统一为224)
  • 将低分辨率图像转为灰度图,保存至./ImageNet/dogs_bnw/
  • 输入低分辨率黑白图至GAN生成彩色图像

问题卡点

训练阶段出现维度匹配错误,报错信息:

ValueError: `logits` and `labels` must have the same shape, received ((32, 28, 28, 3) vs (32, 224, 224)).

报错发生在g_loss = generator.train_on_batch(noise, real_images)行。

相关代码

生成器与判别器代码

# GAN model for recoloring black and white images
generator = Sequential()
generator.add(Dense(7 * 7 * 128, input_dim=100))
generator.add(Reshape((7, 7, 128)))
generator.add(Conv2DTranspose(64, kernel_size=5, strides=2, padding='same'))
generator.add(Conv2DTranspose(32, kernel_size=5, strides=2, padding='same'))
generator.add(Conv2DTranspose(3, kernel_size=5, activation='sigmoid', padding='same'))

# Discriminator model
discriminator = Sequential()
discriminator.add(Flatten(input_shape=(224, 224, 3)))
discriminator.add(Dense(1, activation='sigmoid'))

# Compile the generator model
optimizer = Adam(learning_rate=0.0002, beta_1=0.5)
generator.compile(loss='binary_crossentropy', optimizer=optimizer)

# Train the GAN to recolor images
epochs = 10000
batch_size = 32

训练循环代码

for epoch in range(epochs):
    idx = np.random.randint(0, bw_images.shape[0], batch_size)
    real_images = bw_images[idx]

    noise = np.random.normal(0, 1, (batch_size, 100))
    generated_images = generator.predict(noise)

    # noise_rs = noise.reshape(-1, 1)
    g_loss = generator.train_on_batch(noise, real_images)

    if epoch % 100 == 0:
        print(f"Epoch: {epoch}, Generator Loss: {g_loss}")

变量形状

  • real_images.shape:(32, 224, 224)(灰度图,无通道维度)
  • noise.shape:(32, 100)

问题根源分析

  1. 生成器输出与目标尺寸不匹配:当前生成器从100维噪声出发,经过3次转置卷积后输出尺寸为(28,28,3),但目标灰度图尺寸是(224,224),且缺少通道维度。
  2. 生成器输入设计错误:图像上色是条件生成任务,生成器输入应为灰度图而非随机噪声,随机噪声适用于无约束的图像生成任务(如DCGAN)。
  3. GAN训练逻辑错误:直接用生成器拟合真实图像,未遵循GAN的对抗训练逻辑——生成器需通过判别器的反馈优化,而非直接以真实图像为标签。

解决方案

1. 修正生成器的输入输出结构

将生成器输入改为灰度图的形状(224,224,1),调整网络结构使输出为(224,224,3)的彩色图:

from tensorflow.keras.layers import LeakyReLU, Input
from tensorflow.keras.models import Model

generator = Sequential()
# 输入:224×224灰度图,通道数1
generator.add(Conv2D(64, kernel_size=5, strides=2, padding='same', input_shape=(224,224,1)))
generator.add(LeakyReLU(0.2))
generator.add(Conv2D(128, kernel_size=5, strides=2, padding='same'))
generator.add(LeakyReLU(0.2))
generator.add(Conv2D(256, kernel_size=5, strides=2, padding='same'))
generator.add(LeakyReLU(0.2))
# 上采样回到224×224
generator.add(Conv2DTranspose(128, kernel_size=5, strides=2, padding='same'))
generator.add(LeakyReLU(0.2))
generator.add(Conv2DTranspose(64, kernel_size=5, strides=2, padding='same'))
generator.add(LeakyReLU(0.2))
generator.add(Conv2DTranspose(3, kernel_size=5, strides=2, padding='same', activation='sigmoid'))

2. 调整图像数据维度

将灰度图扩展通道维度,同时准备对应的彩色真实图像作为判别器的真实样本:

# 假设color_images是原始彩色图,形状为(样本数,224,224,3),已归一化到[0,1]
# 扩展灰度图的通道维度
bw_images = np.expand_dims(bw_images, axis=-1)  # 形状变为(样本数,224,224,1)

3. 修正GAN训练逻辑

遵循标准GAN的对抗训练流程,交替训练判别器和生成器:

编译完整GAN模型

# 先编译判别器
discriminator.compile(loss='binary_crossentropy', optimizer=optimizer, metrics=['accuracy'])

# 构建完整GAN(冻结判别器,仅训练生成器)
discriminator.trainable = False
gan_input = Input(shape=(224,224,1))
generated_color = generator(gan_input)
gan_output = discriminator(generated_color)
gan = Model(gan_input, gan_output)
gan.compile(loss='binary_crossentropy', optimizer=optimizer)

修正训练循环

for epoch in range(epochs):
    # ---------------------
    # 训练判别器
    # ---------------------
    # 1. 用真实彩色图像训练
    idx = np.random.randint(0, color_images.shape[0], batch_size)
    real_color = color_images[idx]
    real_labels = np.ones((batch_size, 1))
    d_loss_real = discriminator.train_on_batch(real_color, real_labels)
    
    # 2. 用生成的彩色图像训练
    bw_batch = bw_images[idx]
    fake_color = generator.predict(bw_batch, verbose=0)
    fake_labels = np.zeros((batch_size, 1))
    d_loss_fake = discriminator.train_on_batch(fake_color, fake_labels)
    
    # 计算判别器总损失
    d_loss = 0.5 * np.add(d_loss_real, d_loss_fake)
    
    # ---------------------
    # 训练生成器
    # ---------------------
    # 让判别器误以为生成的图像是真实的
    g_loss = gan.train_on_batch(bw_batch, real_labels)
    
    # 打印日志
    if epoch % 100 == 0:
        print(f"Epoch: {epoch}, D Loss: {d_loss[0]:.4f}, G Loss: {g_loss:.4f}")

4. 额外注意事项

  • 确保所有图像尺寸统一为224×224(原描述中的244应为笔误)
  • 将所有图像数据归一化到[0,1]区间,匹配生成器sigmoid输出的范围
  • 少量数据易导致过拟合,可加入数据增强(如随机水平翻转、随机裁剪)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.03 11:19:51