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

GAN训练速度过慢求助:本地运行远慢于Colab环境

GAN本地训练速度慢的优化方案

问题背景

我编写了一段GAN代码,在Colab环境训练时单步耗时仅3-4ms,但在新购置的本地电脑运行相同代码时,单步耗时高达72-82ms。已尝试切换PyTorch、VS Code开发环境,以及降级TensorFlow版本,问题仍未解决,需要有效的优化方案。

原始代码

import os
import numpy as np
import cv2
import matplotlib.pyplot as plt
from tensorflow import keras
from tensorflow.keras import models, layers
from tensorflow.keras.optimizers import Adam
from tensorflow.keras.models import Model

def build_generator(latent_dim):
    """Build the generator model."""
    model = models.Sequential([
        layers.Dense(256 * 8 * 8, input_dim=latent_dim),
        layers.LeakyReLU(alpha=0.2),
        layers.BatchNormalization(momentum=0.8),
        layers.Reshape((8, 8, 256)),
        layers.Conv2DTranspose(256, (4, 4), strides=(2, 2), padding='same'),
        layers.LeakyReLU(alpha=0.2),
        layers.BatchNormalization(momentum=0.8),
        layers.Conv2DTranspose(128, (4, 4), strides=(2, 2), padding='same'),
        layers.LeakyReLU(alpha=0.2),
        layers.BatchNormalization(momentum=0.8),
        layers.Conv2DTranspose(64, (4, 4), strides=(2, 2), padding='same'),
        layers.LeakyReLU(alpha=0.2),
        layers.BatchNormalization(momentum=0.8),
        layers.Conv2D(3, (3, 3), activation='tanh', padding='same')
    ])
    return model

def generate_and_visualize_images(generator, latent_dim, num_samples=10, save_path='./generated_images3/'):
    """Generate and visualize sample images and save them."""
    os.makedirs(save_path, exist_ok=True)

    noise = np.random.normal(0, 1, (num_samples, latent_dim))
    generated_images = generator.predict(noise)

    for i in range(num_samples):
        plt.imshow((generated_images[i] * 127.5 + 127.5).astype(np.uint8))
        plt.axis('off')
        plt.tight_layout()
        plt.savefig(f'{save_path}/generated_image_{i}.png')
        plt.close()

latent_dim = 50
epochs = 10000
batch_size = 64

# Assuming build_generator is a function that returns a compiled generator model
generator = build_generator(latent_dim)

discriminator = models.Sequential([
    layers.Conv2D(64, (3, 3), strides=(2, 2), padding='same', input_shape=(64, 64, 3)),
    layers.LeakyReLU(alpha=0.2),
    layers.Dropout(0.4),
    layers.Conv2D(128, (3, 3), strides=(2, 2), padding='same'),
    layers.LeakyReLU(alpha=0.2),
    layers.Dropout(0.4),
    layers.Conv2D(256, (3, 3), strides=(2, 2), padding='same'),
    layers.LeakyReLU(alpha=0.2),
    layers.Dropout(0.4),
    layers.Flatten(),
    layers.Dense(1, activation='sigmoid')
])
discriminator.compile(loss='binary_crossentropy', optimizer=Adam(learning_rate=0.0004, beta_1=0.5), metrics=['accuracy'])

discriminator.trainable = False
gan_input = layers.Input(shape=(latent_dim,))
gan_output = discriminator(generator(gan_input))
gan = Model(gan_input, gan_output)
gan.compile(loss='binary_crossentropy', optimizer=Adam(learning_rate=0.0001, beta_1=0.5))

# Load and preprocess dataset
folder_path = "/content/drive/MyDrive/alan/deep learning/resimdosyaları/IMAGES2/IMAGES2/"
images = []

if not os.path.exists(folder_path):
    raise FileNotFoundError(f"Error: The folder {folder_path} does not exist.")

for filename in os.listdir(folder_path):
    img_path = os.path.join(folder_path, filename)
    if os.path.isfile(img_path):
        try:
            img = cv2.imread(img_path)
            if img is not None:
                img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
                img = cv2.resize(img, (64, 64))
                images.append(img)
        except Exception as e:
            print(f"Error reading {img_path}: {e}")

x_train_s = np.array(images, dtype=np.float32)
x_train_s = (x_train_s - 127.5) / 127.5  # Normalize to [-1, 1]

# Lists to store loss values
discriminator_losses = []
generator_losses = []

for epoch in range(epochs):
    noise = np.random.normal(0, 1, (batch_size, latent_dim))
    fake_images = generator.predict(noise)
    real_images = x_train_s[np.random.randint(0, x_train_s.shape[0], batch_size)]

    # Apply label smoothing to real labels
    real_labels = np.ones((batch_size, 1)) * 0.9
    fake_labels = np.zeros((batch_size, 1)) + 0.1

    discriminator_loss_real = discriminator.train_on_batch(real_images, real_labels)
    discriminator_loss_fake = discriminator.train_on_batch(fake_images, fake_labels)
    discriminator_loss = 0.5 * np.add(discriminator_loss_real, discriminator_loss_fake)
    discriminator_losses.append(discriminator_loss[0])

    noise = np.random.normal(0, 1, (batch_size, latent_dim))
    generator_loss = gan.train_on_batch(noise, np.ones((batch_size, 1)) * 0.9)
    generator_losses.append(generator_loss)

    if epoch % 100 == 0:
        print(f"Epoch: {epoch}, Discriminator Loss: {discriminator_loss[0]}, Generator Loss: {generator_loss}")

# Plot the losses
plt.figure(figsize=(10, 5))
plt.plot(discriminator_losses, label='Discriminator Loss')
plt.plot(generator_losses, label='Generator Loss')
plt.xlabel('Epoch')
plt.ylabel('Loss')
plt.legend()
plt.title('Discriminator and Generator Losses')
plt.show()

# Generate final images
generate_and_visualize_images(generator, latent_dim)

优化方案

1. 硬件与环境验证

  • 确认GPU加速是否启用:运行以下代码检查TensorFlow是否识别到GPU
    import tensorflow as tf
    print(tf.config.list_physical_devices('GPU'))
    
    若未检测到GPU,检查显卡驱动、CUDA、CUDNN版本是否与当前TensorFlow版本匹配
  • 释放GPU资源:关闭本地占用GPU的程序(如浏览器、视频软件),避免资源抢占

2. 代码核心优化

  • 用tf.function替代predict:predict函数存在Python开销,改用tf.function装饰推理逻辑,加速假图像生成
    @tf.function
    def generate_fake(generator, noise):
        return generator(noise, training=False)
    
    # 训练循环中替换:
    noise = tf.random.normal((batch_size, latent_dim))
    fake_images = generate_fake(generator, noise)
    
  • 改用TensorFlow Dataset加载数据:手动循环读取图片效率低下,用tf.data实现并行加载、预取,减少数据准备耗时
    def load_img(img_path):
        img = tf.io.read_file(img_path)
        img = tf.image.decode_jpeg(img, channels=3)
        img = tf.image.resize(img, (64, 64))
        return (img - 127.5) / 127.5
    
    # 替换原有数据加载逻辑
    dataset = tf.data.Dataset.list_files(f"{folder_path}/*")  # 匹配所有图片文件
    dataset = dataset.map(load_img, num_parallel_calls=tf.data.AUTOTUNE)
    dataset = dataset.shuffle(buffer_size=len(x_train_s)).batch(batch_size).prefetch(tf.data.AUTOTUNE)
    
    # 训练循环中取真实图片:
    real_images = next(iter(dataset))
    
  • 减少numpy与Tensor转换:用tf.random.normal替代np.random.normal,避免跨框架数据转换开销
  • 移除不必要的计算指标:训练判别器时不需要计算accuracy,编译时去掉metrics=['accuracy'],减少额外计算

3. TensorFlow配置调优

  • 启用混合精度训练:利用半精度计算提升速度,同时降低显存占用
    from tensorflow.keras.mixed_precision import set_global_policy
    set_global_policy('mixed_float16')
    
  • 设置GPU内存增长模式:避免一次性占用全部显存,减少碎片化问题
    gpus = tf.config.list_physical_devices('GPU')
    if gpus:
        for gpu in gpus:
            tf.config.experimental.set_memory_growth(gpu, True)
    

内容的提问来源于stack exchange,提问作者Muhammed Gülsoy

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 09:12:32