PyTorch Lightning GAN训练报错:缺少optimizer_idx参数排查求助
问题排查:PyTorch Lightning GAN训练中
optimizer_idx参数缺失错误 问题场景
训练基于PyTorch Lightning的GAN模型时,执行以下代码:
trainer = pl.Trainer(max_epochs=20, devices=AVAIL_GPUS, accelerator='gpu') trainer.fit(GAN(), MNISTDataModule())
出现错误:
TypeError: GAN.training_step() missing 1 required positional argument: 'optimizer_idx'
但training_step函数中已定义optimizer_idx参数,相关GAN模型代码如下:
# GAN model using PyTorch lightning class GAN(pl.LightningModule): # learning rate 0.002 (tweak this) # latent dimension 100 def __init__(self, latent_dim=100, lr=0.002): super().__init__() self.save_hyperparameters() # save self.hparams self.automatic_optimization = False # activates manual optimization self.generator = Generator(latent_dim=self.hparams.latent_dim) self.discriminator = Discriminator() # random noise self.validation_z = torch.randn(6, self.hparams.latent_dim) # 6 images # forward pass # - input tensor z def forward(self, z): return self.generator(z) # loss function # - predicted label y_hat # - actual label y def adversarial_loss(self, y_hat, y): return F.binary_cross_entropy(y_hat, y) def training_step(self, batch, batch_idx, optimizer_idx): # tensor real_imgs real_imgs, labels = batch # sample noise z = torch.randn(real_imgs.shape[0], self.hparams.latent_dim) z = z.type_as(real_imgs) # to use GPU # train generator: max log(D(G(z))) where z is random noise / fake images if optimizer_idx == 0: fake_imgs = self(z) y_hat = self.discriminator(fake_imgs) y = torch.ones(real_imgs.size(0), 1) y = y.type_as(real_imgs) g_loss = self.adversarial_loss(y_hat, y) log_dict = { "g_loss": g_loss } return { "loss": g_loss, "progress_bar": log_dict} # train discriminator: max log(D(x)) + log(1 - D(G(z))) if optimizer_idx == 1: # how well can discriminator label as real y_hat_real = self.discriminator(real_imgs) y_real = torch.ones(real_imgs.size(0), 1) y_real = y_real.type_as(real_imgs) real_loss = self.adversarial_loss(y_hat_real, y_real) # how well can discriminator label as fake y_hat_fake = self.discriminator(self(z).detach()) # detach: creates a new tensor that is detached from computational graph (since we already do fake_imgs = self(z)) y_fake = torch.zeros(real_imgs.size(0), 1) y_fake = y_fake.type_as(real_imgs) fake_loss = self.adversarial_loss(y_hat_fake, y_fake) d_loss = (real_loss + fake_loss) / 2 log_dict = { "d_loss": d_loss } return { "loss": d_loss, "progress_bar": log_dict, "log": log_dict } # log in case we want to use tensorboard def configure_optimizers(self): lr = self.hparams.lr opt_generator = torch.optim.Adam(self.generator.parameters(), lr=lr) opt_discriminator = torch.optim.Adam(self.discriminator.parameters(), lr=lr) # return empty list [] in case we use scheduler return [opt_generator, opt_discriminator]
错误原因
核心矛盾在于手动优化模式与多优化器参数的冲突:
- 你在
__init__中设置了self.automatic_optimization = False,开启了手动优化模式。 - 在手动优化模式下,PyTorch Lightning不会自动向
training_step传入optimizer_idx参数,因为该模式要求开发者手动获取优化器并执行更新操作,框架不再负责多优化器的调度。 - 但你的
training_step函数仍保留了optimizer_idx参数,导致框架调用时因参数不匹配抛出错误。
解决方案
提供两种可行修复方案,根据需求选择:
方案1:保持手动优化,修改training_step
移除optimizer_idx参数,手动获取优化器并执行训练逻辑:
def training_step(self, batch, batch_idx): real_imgs, labels = batch z = torch.randn(real_imgs.shape[0], self.hparams.latent_dim) z = z.type_as(real_imgs) # 获取两个优化器 opt_g, opt_d = self.optimizers() # 训练生成器 fake_imgs = self(z) y_hat = self.discriminator(fake_imgs) y = torch.ones(real_imgs.size(0), 1).type_as(real_imgs) g_loss = self.adversarial_loss(y_hat, y) # 手动反向传播+更新 self.manual_backward(g_loss) opt_g.step() opt_g.zero_grad() # 训练判别器 y_hat_real = self.discriminator(real_imgs) y_real = torch.ones(real_imgs.size(0), 1).type_as(real_imgs) real_loss = self.adversarial_loss(y_hat_real, y_real) y_hat_fake = self.discriminator(self(z).detach()) y_fake = torch.zeros(real_imgs.size(0), 1).type_as(real_imgs) fake_loss = self.adversarial_loss(y_hat_fake, y_fake) d_loss = (real_loss + fake_loss) / 2 # 手动反向传播+更新 self.manual_backward(d_loss) opt_d.step() opt_d.zero_grad() # 记录日志 self.log("g_loss", g_loss, prog_bar=True) self.log("d_loss", d_loss, prog_bar=True, logger=True)
方案2:关闭手动优化,保留optimizer_idx
删除self.automatic_optimization = False这一行,让PyTorch Lightning自动处理多优化器调度:
def __init__(self, latent_dim=100, lr=0.002): super().__init__() self.save_hyperparameters() # 移除手动优化设置 # self.automatic_optimization = False self.generator = Generator(latent_dim=self.hparams.latent_dim) self.discriminator = Discriminator() self.validation_z = torch.randn(6, self.hparams.latent_dim)
此时原有的training_step代码无需修改,框架会自动传入optimizer_idx参数,按顺序调度两个优化器。
内容的提问来源于stack exchange,提问作者Jpark9061
相关产品推荐
相关产品推荐

