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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.21 04:30:05