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

如何将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" in image_dataset_from_directory to 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_shape parameter in the first Conv2D layer from [112, 112, 1] to [112, 112, 3]

4. Fix Visualization Code

Since we're now generating RGB images:

  • Remove the cmap='gray' argument from plt.imshow()
  • Convert generated images to uint8 when 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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.30 13:42:43