使用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)
问题根源分析
- 生成器输出与目标尺寸不匹配:当前生成器从100维噪声出发,经过3次转置卷积后输出尺寸为
(28,28,3),但目标灰度图尺寸是(224,224),且缺少通道维度。 - 生成器输入设计错误:图像上色是条件生成任务,生成器输入应为灰度图而非随机噪声,随机噪声适用于无约束的图像生成任务(如DCGAN)。
- 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
相关产品推荐
相关产品推荐

