文本到图像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,这个张量的计算图已经被销毁,导致生成器反向传播时无法追溯到生成器的参数。同时代码中大量不必要的设备迁移操作,不仅增加开销,还可能破坏计算图连续性。
具体修复步骤
- 重新计算生成器训练所需的判别器输出
训练生成器时不能复用判别器阶段的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()
减少不必要的设备迁移
避免反复将同一个张量(比如text_embeddings)在GPU和CPU之间来回移动,尽量一次性完成设备分配,降低计算图被破坏的风险。优化内存清理时机
不要过于频繁调用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

