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
若未检测到GPU,检查显卡驱动、CUDA、CUDNN版本是否与当前TensorFlow版本匹配import tensorflow as tf print(tf.config.list_physical_devices('GPU')) - 释放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
相关产品推荐
相关产品推荐

