如何将TensorFlow官方教程中的灰度DCGAN适配为RGB彩色图像生成版本?
Adapting DCGAN for RGB Image Generation
Great question! Converting the TensorFlow DCGAN tutorial from grayscale to RGB just requires adjusting channel dimensions in your data pipeline and model layers. Let's walk through each change you need to make, plus the full adapted code:
1. Update Data Loading & Preprocessing
First, tweak how you load and prepare your RGB image data:
- Set
color_mode="rgb"inimage_dataset_from_directoryto load 3-channel color images - Adjust the reshaping step to use 3 channels instead of 1
- The pixel normalization logic stays identical—we still want values scaled to the [-1, 1] range for stable GAN training
2. Modify the Generator Model
The only critical change here is the final transposed convolution layer:
- Change the output channel count from 1 to 3 (to match RGB's 3 color channels)
- Update the output shape assertion to reflect the new channel dimension
3. Adjust the Discriminator Model
Update the input shape to accept 3-channel images:
- Change the
input_shapeparameter in the firstConv2Dlayer from[112, 112, 1]to[112, 112, 3]
4. Fix Visualization Code
Since we're now generating RGB images:
- Remove the
cmap='gray'argument fromplt.imshow() - Convert generated images to
uint8when saving/visualizing (RGB values need to be 0-255 integers for proper display)
Full Adapted RGB DCGAN Code
from google.colab import drive drive.mount('/content/drive') import tensorflow as tf import glob import matplotlib.pyplot as plt import numpy as np import os import PIL from tensorflow.keras import layers import time from IPython import display # Load RGB dataset train_dataset = tf.keras.preprocessing.image_dataset_from_directory( "/content/drive/MyDrive/birds", seed=123, validation_split=0, image_size=(112, 112), color_mode="rgb", # Switched to RGB mode shuffle=True, batch_size=1) train_images_array = [] for images, _ in train_dataset: for i in range(len(images)): train_images_array.append(images[i]) train_images = np.array(train_images_array) # Reshape to match RGB channel count (3) train_images = train_images.reshape(train_images.shape[0],112,112,3).astype('float32') train_images = (train_images - 127.5) / 127.5 # Normalize to [-1, 1] BUFFER_SIZE = 60000 BATCH_SIZE = 8 # Batch and shuffle the data dataset_ = tf.data.Dataset.from_tensor_slices(train_images).shuffle(BUFFER_SIZE).batch(BATCH_SIZE) def make_generator_model(): model = tf.keras.Sequential() model.add(layers.Dense(7*7*256, use_bias=False, input_shape=(100,))) model.add(layers.BatchNormalization()) model.add(layers.LeakyReLU()) model.add(layers.Reshape((7, 7, 256))) assert model.output_shape == (None, 7, 7, 256) model.add(layers.Conv2DTranspose(128, (5, 5), strides=(1, 1), padding='same', use_bias=False)) assert model.output_shape == (None, 7, 7, 128) model.add(layers.BatchNormalization()) model.add(layers.LeakyReLU()) model.add(layers.Conv2DTranspose(64, (5, 5), strides=(2, 2), padding='same', use_bias=False)) assert model.output_shape == (None, 14, 14, 64) model.add(layers.BatchNormalization()) model.add(layers.LeakyReLU()) # Final layer: output 3 channels for RGB model.add(layers.Conv2DTranspose(3, (20, 20), strides=(8, 8), padding='same', use_bias=False, activation='tanh')) assert model.output_shape == (None, 112, 112, 3) # Updated channel count return model generator = make_generator_model() noise = tf.random.normal([1, 100]) generated_image = generator(noise, training=False) # Visualize RGB image without grayscale filter plt.imshow(generated_image[0, :, :, :] * 127.5 + 127.5) def make_discriminator_model(): model = tf.keras.Sequential() # Input shape updated to accept 3-channel RGB images model.add(layers.Conv2D(64, (10, 10), strides=(2, 2), padding='same', input_shape=[112, 112, 3])) model.add(layers.LeakyReLU()) model.add(layers.Dropout(0.3)) model.add(layers.Conv2D(128, (5, 5), strides=(2, 2), padding='same')) model.add(layers.LeakyReLU()) model.add(layers.Dropout(0.3)) model.add(layers.Flatten()) model.add(layers.Dense(1)) return model discriminator = make_discriminator_model() decision = discriminator(generated_image) print (decision) # Loss functions and optimizers remain unchanged cross_entropy = tf.keras.losses.BinaryCrossentropy(from_logits=True) def discriminator_loss(real_output, fake_output): real_loss = cross_entropy(tf.ones_like(real_output), real_output) fake_loss = cross_entropy(tf.zeros_like(fake_output), fake_output) total_loss = real_loss + fake_loss return total_loss def generator_loss(fake_output): return cross_entropy(tf.ones_like(fake_output), fake_output) generator_optimizer = tf.keras.optimizers.Adam(1e-4) discriminator_optimizer = tf.keras.optimizers.Adam(1e-4) checkpoint_dir = '/content/drive/MyDrive/training_checkpoints11' checkpoint_prefix = os.path.join(checkpoint_dir, "ckpt") checkpoint = tf.train.Checkpoint(generator_optimizer=generator_optimizer, discriminator_optimizer=discriminator_optimizer, generator=generator, discriminator=discriminator) EPOCHS = 50 noise_dim = 100 num_examples_to_generate = 16 seed = tf.random.normal([num_examples_to_generate, noise_dim]) def generate_and_save_images(model, epoch, test_input): predictions = model(test_input, training=False) fig = plt.figure(figsize=(4, 4)) for i in range(predictions.shape[0]): plt.subplot(4, 4, i+1) # Convert to uint8 for proper RGB display img = (predictions[i, :, :, :] * 127.5 + 127.5).numpy().astype(np.uint8) plt.imshow(img) plt.axis('off') plt.savefig('image_at_epoch_{:04d}.png'.format(epoch)) plt.show() @tf.function def train_step(images): noise = tf.random.normal([BATCH_SIZE, noise_dim]) with tf.GradientTape() as gen_tape, tf.GradientTape() as disc_tape: generated_images = generator(noise, training=True) real_output = discriminator(images, training=True) fake_output = discriminator(generated_images, training=True) gen_loss = generator_loss(fake_output) disc_loss = discriminator_loss(real_output, fake_output) gradients_of_generator = gen_tape.gradient(gen_loss, generator.trainable_variables) gradients_of_discriminator = disc_tape.gradient(disc_loss, discriminator.trainable_variables) generator_optimizer.apply_gradients(zip(gradients_of_generator, generator.trainable_variables)) discriminator_optimizer.apply_gradients(zip(gradients_of_discriminator, discriminator.trainable_variables)) def train(dataset, epochs): for epoch in range(epochs): start = time.time() for image_batch in dataset: train_step(image_batch) display.clear_output(wait=True) generate_and_save_images(generator, epoch + 1, seed) if (epoch + 1) % 8 == 0: checkpoint.save(file_prefix = checkpoint_prefix) print ('Time for epoch {} is {} sec'.format(epoch + 1, time.time()-start)) display.clear_output(wait=True) generate_and_save_images(generator, epochs, seed) train(dataset_, 128) # Post-training visualization noise = tf.random.normal([1, 100]) generated_image = generator(noise, training=False) print(generated_image.shape) img = (generated_image[0, :, :, :] * 127.5 + 127.5).numpy().astype(np.uint8) plt.imshow(img) # Restore checkpoint if needed checkpoint.restore(tf.train.latest_checkpoint(checkpoint_dir))
内容的提问来源于stack exchange,提问作者Brian Danowski
相关产品推荐
相关产品推荐

