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

如何避免Google Colab中GAN训练会话崩溃?

GAN训练循环会话崩溃的解决方法

问题背景

提供的GAN训练循环在执行时会话持续崩溃,已尝试切换TPU运行时、每个epoch后调用tf.keras.backend.clear_session(),但未解决问题。训练代码如下:

for epoch in range(EPOCHS):
tf.keras.backend.clear_session()
for batch_images in dataset:
# Generate noise for the generator
noise = tf.random.normal([BATCH_SIZE, latent_dim])

    with tf.GradientTape() as gen_tape, tf.GradientTape() as disc_tape:

        generated_frames = generator(noise, training=True)

        # Get discriminator predictions for real and fake frames
        real_output = discriminator(batch_images, training=True)
        fake_output = discriminator(generated_frames, training=True)

        # Calculate the generator and discriminator losses
        gen_loss = generator_loss(fake_output)
        disc_loss = discriminator_loss(real_output, fake_output)

    # Compute the gradients of the generator and discriminator with respect to their loss
    gradients_of_generator = gen_tape.gradient(gen_loss, generator.trainable_variables)
    gradients_of_discriminator = disc_tape.gradient(disc_loss, discriminator.trainable_variables)

    # Apply the gradients to update the generator and discriminator variables
    generator_optimizer.apply_gradients(zip(gradients_of_generator, generator.trainable_variables))
    discriminator_optimizer.apply_gradients(zip(gradients_of_discriminator, discriminator.trainable_variables))

    # Train the generator to predict the future frames
    for i in range(num_frames):
        with tf.GradientTape() as gen_tape:
            # Generate noise for each frame
            noise = tf.random.normal([BATCH_SIZE, latent_dim])

            # Generate a single frame
            generated_frame = generator(noise, training=True)

            # Calculate the generator loss for predicting the next frame
            next_frame_loss = next_frame_loss_function(batch_images[:, i+1], generated_frame)

        gradients_of_generator = gen_tape.gradient(next_frame_loss, generator.trainable_variables)

        generator_optimizer.apply_gradients(zip(gradients_of_generator, generator.trainable_variables))

可行解决方法

1. 修复梯度磁带作用域与内存泄漏

  • tf.keras.backend.clear_session()的位置错误,应放在epoch循环初始化前执行;同时梯度磁带使用后需手动清理引用,避免内存堆积。
  • 优化后的核心循环结构:
tf.keras.backend.clear_session()  # 初始化前全局清理一次
for epoch in range(EPOCHS):
    for batch_images in dataset:
        noise = tf.random.normal([BATCH_SIZE, latent_dim])
        # 限制磁带作用域在batch内,执行后手动释放资源
        with tf.GradientTape() as gen_tape, tf.GradientTape() as disc_tape:
            generated_frames = generator(noise, training=True)
            real_output = discriminator(batch_images, training=True)
            fake_output = discriminator(generated_frames, 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))
        # 手动删除临时变量,帮助垃圾回收
        del gen_tape, disc_tape, gradients_of_generator, gradients_of_discriminator

        # 帧预测训练部分同样清理磁带
        for i in range(num_frames):
            with tf.GradientTape() as gen_tape:
                noise = tf.random.normal([BATCH_SIZE, latent_dim])
                generated_frame = generator(noise, training=True)
                next_frame_loss = next_frame_loss_function(batch_images[:, i+1], generated_frame)
            
            gradients_of_generator = gen_tape.gradient(next_frame_loss, generator.trainable_variables)
            generator_optimizer.apply_gradients(zip(gradients_of_generator, generator.trainable_variables))
            del gen_tape, gradients_of_generator
    # epoch结束后可选再次清理
    tf.keras.backend.clear_session()

2. 限制显存使用与批量大小

  • 检查BATCH_SIZE是否过大,尝试减半或逐步调整;同时开启TensorFlow显存按需分配:
gpus = tf.config.list_physical_devices('GPU')
if gpus:
    try:
        for gpu in gpus:
            tf.config.experimental.set_memory_growth(gpu, True)
    except RuntimeError as e:
        print(e)

3. 合并生成器训练步骤,避免梯度冲突

  • 原代码中生成器被两次独立更新梯度,易导致计算图混乱。建议合并对抗损失与帧预测损失,单次更新梯度:
# 替代原有的两次生成器训练逻辑
with tf.GradientTape() as gen_tape:
    noise = tf.random.normal([BATCH_SIZE, latent_dim])
    generated_frames = generator(noise, training=True)
    fake_output = discriminator(generated_frames, training=True)
    gen_adv_loss = generator_loss(fake_output)
    
    # 计算帧预测平均损失
    frame_pred_loss = 0.0
    for i in range(num_frames):
        frame_noise = tf.random.normal([BATCH_SIZE, latent_dim])
        generated_frame = generator(frame_noise, training=True)
        frame_pred_loss += next_frame_loss_function(batch_images[:, i+1], generated_frame)
    frame_pred_loss /= num_frames
    
    # 总损失可调整权重平衡两个任务
    total_gen_loss = gen_adv_loss + 0.5 * frame_pred_loss

gradients_of_generator = gen_tape.gradient(total_gen_loss, generator.trainable_variables)
generator_optimizer.apply_gradients(zip(gradients_of_generator, generator.trainable_variables))

4. 优化数据集加载

  • 确保数据集迭代时无内存泄漏,使用prefetch优化加载流程:
dataset = dataset.batch(BATCH_SIZE).prefetch(tf.data.AUTOTUNE)
  • 检查batch_images的维度是否与模型输入匹配,维度不匹配可能引发计算图异常。

5. 切换到图模式训练(可选)

  • 若使用TensorFlow 2.x,尝试禁用eager execution,减少内存开销:
tf.compat.v1.disable_eager_execution()

注意:切换后需调整代码适配图模式,例如用tf.function装饰训练步骤。

内容的提问来源于stack exchange,提问作者PURVI VERMA

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.21 11:05:31