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

文本到图像GAN训练时Generator反向传播报错求助

问题描述

我正在训练一款接收文本嵌入、输出图像的GAN模型,因资源限制将Generator部署在GPU、Discriminator部署在CPU。训练循环中已删除无用变量并清理内存,但在Generator的反向传播阶段出现RuntimeError:Trying to backward through the graph a second time, but the saved intermediate results have already been freed. Specify retain_graph=True when calling backward the first time.。我尝试过替换反向传播方式、设置retain_graph=True等方法,但均未解决问题,且retain_graph=True会增加资源消耗,希望找到有效解决方案。

训练函数代码

import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader
from torchvision.utils import save_image
import torch.nn.functional as F
from tqdm import tqdm

def train_gan(generator, discriminator, dataset, batch_size, num_epochs, device):
    
    # Set up loss functions and optimizers
    adversarial_loss_generator = nn.BCELoss()
    adversarial_loss_discriminator=nn.BCELoss()
    generator_optimizer = optim.Adam(generator.parameters(), lr=0.0002, betas=(0.5, 0.999))
    discriminator_optimizer = optim.Adam(discriminator.parameters(), lr=0.0002, betas=(0.5, 0.999))
    
    # Set up data loader
    data_loader = DataLoader(dataset, batch_size=batch_size, shuffle=True)
    generator.to(device)
    discriminator.to("cpu")
    
    # Train the GAN
    for epoch in range(num_epochs):
        for i, data in enumerate(tqdm(data_loader)):
            # Load data and labels onto the device
            text_embeddings = data['text_embedding'].to(device)            
            # Generate fake images using the generator and the text embeddings
            noise = torch.randn(batch_size,generator.latent_dim).to(device)
            fake_images = generator(text_embeddings,noise)
            del noise
            clear_memory()
            fake_images = F.interpolate(fake_images, size=(512, 512), mode='bilinear', align_corners=False)
            # Train the discriminator
            discriminator_optimizer.zero_grad()
            real_images = data['image'].to("cpu")
            real_labels = torch.ones(real_images.size(0), 1).to("cpu")
            
            text_embeddings=text_embeddings.to("cpu")
            real_predictions = discriminator(real_images, text_embeddings)
            del real_images
            clear_memory()
            real_predictions=real_predictions.to("cpu")
            real_labels= real_labels.to("cpu")
            real_loss = adversarial_loss_discriminator(real_predictions, real_labels)
            real_loss=real_loss.to("cpu")
            fake_images=fake_images.to("cpu")
            text_embeddings=text_embeddings.to("cpu")
            fake_predictions = discriminator(fake_images, text_embeddings)
            fake_predictions=fake_predictions.to("cpu")
            fake_labels = torch.zeros(fake_images.size(0), 1).to("cpu")
            fake_loss = adversarial_loss_discriminator(fake_predictions, fake_labels)
            del fake_labels
            clear_memory()
            fake_loss=fake_loss.to ("cpu")
            discriminator_loss = real_loss + fake_loss
            discriminator_loss=discriminator_loss.to("cpu")
            discriminator_loss.backward()
            discriminator_optimizer.step()
            
            # Train the generator
            generator_optimizer.zero_grad()
            #fake_predictions = discriminator(fake_images, text_embeddings)
            del text_embeddings
            del fake_images
            clear_memory()
            fake_predictions=fake_predictions.to(device)
            real_labels=real_labels.to(device)
            
            generator_loss = adversarial_loss_generator(fake_predictions, real_labels)
            generator_loss=generator_loss.to(device)
            generator_loss.backward()
            generator_optimizer.step()
            del fake_predictions
            del real_labels
            clear_memory()

            # Save generated images and model checkpoints every 500 batches
            if i % 100 == 0:
                with torch.no_grad():
                    noise = torch.randn(batch_size,generator.latent_dim).to("cpu")
                    text_embeddings = data['text_embedding'].to("cpu")            
                    fake_images = generator(noise,text_embeddings)
                save_image(fake_images, f"images\generated_images_epoch_{epoch}_batch_{i}.png", normalize=True, nrow=4)
                torch.save(generator.state_dict(), f"models\generator_checkpoint_epoch_{epoch}_batch_{i}.pt")
                torch.save(discriminator.state_dict(), f"models\discriminator_checkpoint_epoch_{epoch}_batch_{i}.pt")
                
        # Print loss at the end of each epoch
        print(f"Epoch [{epoch+1}/{num_epochs}] Discriminator Loss: {discriminator_loss.item()}, Generator Loss: {generator_loss.item()}")

内存清理函数代码

import gc
def clear_memory():
    gc.collect()
    torch.cuda.empty_cache()
解决方案

问题根源

错误核心是:训练判别器时调用discriminator_loss.backward()会默认释放计算图,而后续训练生成器时,你复用了判别器输出的fake_predictions,这个张量的计算图已经被销毁,导致生成器反向传播时无法追溯到生成器的参数。同时代码中大量不必要的设备迁移操作,不仅增加开销,还可能破坏计算图连续性。

具体修复步骤

  1. 重新计算生成器训练所需的判别器输出
    训练生成器时不能复用判别器阶段的fake_predictions,必须重新生成假图像并让判别器再次预测,确保生成器的计算图独立。修改生成器训练部分代码:
# Train the generator
generator_optimizer.zero_grad()
# 重新生成假图像(使用GPU上的文本嵌入,减少设备迁移)
text_embeddings_gen = data['text_embedding'].to(device)
noise_gen = torch.randn(batch_size, generator.latent_dim).to(device)
fake_images_gen = generator(text_embeddings_gen, noise_gen)
fake_images_gen = F.interpolate(fake_images_gen, size=(512, 512), mode='bilinear', align_corners=False)
# 将假图像和文本嵌入移到CPU给判别器
fake_images_gen = fake_images_gen.to("cpu")
text_embeddings_gen = text_embeddings_gen.to("cpu")
# 获取判别器对新假图像的预测
fake_predictions_gen = discriminator(fake_images_gen, text_embeddings_gen)
# 计算生成器损失并反向传播
real_labels_gen = torch.ones(batch_size, 1).to("cpu")
generator_loss = adversarial_loss_generator(fake_predictions_gen, real_labels_gen).to(device)
generator_loss.backward()
generator_optimizer.step()
# 清理变量
del text_embeddings_gen, noise_gen, fake_images_gen, fake_predictions_gen, real_labels_gen
clear_memory()
  1. 减少不必要的设备迁移
    避免反复将同一个张量(比如text_embeddings)在GPU和CPU之间来回移动,尽量一次性完成设备分配,降低计算图被破坏的风险。

  2. 优化内存清理时机
    不要过于频繁调用clear_memory(),建议在每个训练步骤(判别器+生成器训练完成后)统一调用一次,减少GC开销。

修改后的核心训练循环片段

for epoch in range(num_epochs):
    for i, data in enumerate(tqdm(data_loader)):
        # ---------------------
        #  训练判别器
        # ---------------------
        discriminator_optimizer.zero_grad()
        
        # 处理真实样本
        real_images = data['image'].to("cpu")
        real_labels = torch.ones(real_images.size(0), 1).to("cpu")
        text_embeddings_disc = data['text_embedding'].to("cpu")
        real_predictions = discriminator(real_images, text_embeddings_disc)
        real_loss = adversarial_loss_discriminator(real_predictions, real_labels)
        
        # 处理生成样本
        text_embeddings_gen_disc = data['text_embedding'].to(device)
        noise_disc = torch.randn(batch_size, generator.latent_dim).to(device)
        fake_images_disc = generator(text_embeddings_gen_disc, noise_disc)
        fake_images_disc = F.interpolate(fake_images_disc, size=(512, 512), mode='bilinear', align_corners=False).to("cpu")
        text_embeddings_gen_disc = text_embeddings_gen_disc.to("cpu")
        fake_predictions_disc = discriminator(fake_images_disc, text_embeddings_gen_disc)
        fake_labels = torch.zeros(fake_images_disc.size(0), 1).to("cpu")
        fake_loss = adversarial_loss_discriminator(fake_predictions_disc, fake_labels)
        
        # 计算判别器损失并反向传播
        discriminator_loss = real_loss + fake_loss
        discriminator_loss.backward()
        discriminator_optimizer.step()
        
        # ---------------------
        #  训练生成器
        # ---------------------
        generator_optimizer.zero_grad()
        
        # 重新生成假图像用于生成器训练
        text_embeddings_gen = data['text_embedding'].to(device)
        noise_gen = torch.randn(batch_size, generator.latent_dim).to(device)
        fake_images_gen = generator(text_embeddings_gen, noise_gen)
        fake_images_gen = F.interpolate(fake_images_gen, size=(512, 512), mode='bilinear', align_corners=False).to("cpu")
        text_embeddings_gen = text_embeddings_gen.to("cpu")
        fake_predictions_gen = discriminator(fake_images_gen, text_embeddings_gen)
        
        # 计算生成器损失并反向传播
        real_labels_gen = torch.ones(batch_size, 1).to("cpu")
        generator_loss = adversarial_loss_generator(fake_predictions_gen, real_labels_gen).to(device)
        generator_loss.backward()
        generator_optimizer.step()
        
        # 统一清理内存
        del real_images, real_labels, text_embeddings_disc, real_predictions, real_loss
        del text_embeddings_gen_disc, noise_disc, fake_images_disc, fake_predictions_disc, fake_labels, fake_loss
        del text_embeddings_gen, noise_gen, fake_images_gen, fake_predictions_gen, real_labels_gen
        clear_memory()
        
        # 保存图像和模型(原逻辑保留)
        if i % 100 == 0:
            with torch.no_grad():
                noise = torch.randn(batch_size,generator.latent_dim).to(device)
                text_embeddings = data['text_embedding'].to(device)            
                fake_images = generator(text_embeddings, noise)
                fake_images = F.interpolate(fake_images, size=(512, 512), mode='bilinear', align_corners=False).to("cpu")
            save_image(fake_images, f"images/generated_images_epoch_{epoch}_batch_{i}.png", normalize=True, nrow=4)
            torch.save(generator.state_dict(), f"models/generator_checkpoint_epoch_{epoch}_batch_{i}.pt")
            torch.save(discriminator.state_dict(), f"models/discriminator_checkpoint_epoch_{epoch}_batch_{i}.pt")
            
    print(f"Epoch [{epoch+1}/{num_epochs}] Discriminator Loss: {discriminator_loss.item()}, Generator Loss: {generator_loss.item()}")

额外说明

  • 该方案无需设置retain_graph=True,生成器训练使用全新计算图,不依赖判别器阶段已释放的图结构,避免额外内存开销。
  • 设备迁移尽量一次性完成,减少计算图断裂风险。

内容的提问来源于stack exchange,提问作者Mohamed Amine

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 02:45:04