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

PyTorch1.5+运行旧图像生成模型出现inplace操作梯度报错如何解决

PyTorch版本兼容报错修复方案

错误根因

PyTorch 1.5 版本新增了张量版本校验机制,对inplace修改计算图依赖张量的行为检查更严格。你当前代码的问题出在backward函数的执行顺序:

  1. 你在process函数中已经基于初始状态的判别器参数,同时计算出了dis_loss和gen_loss
  2. 原有backward逻辑先执行dis_optimizer.step(),这个操作会inplace修改判别器的参数张量,导致gen_loss对应的计算图中依赖的判别器张量版本发生变化
  3. 后续执行gen_loss.backward()时,计算图找不到对应版本的张量,触发报错

修复方案

方案1:调整反向传播与参数更新顺序(推荐)

将所有损失的反向传播操作全部执行完毕后,再统一执行优化器的参数更新,避免计算图被中途修改:

def backward(self, gen_loss=None, dis_loss=None):
    if dis_loss is not None:
        # 保留计算图,供后续生成器反向传播使用
        dis_loss.backward(retain_graph=True)
    if gen_loss is not None:
        gen_loss.backward()
    # 所有反向传播完成后再更新参数
    if dis_loss is not None:
        self.dis_optimizer.step()
    if gen_loss is not None:
        self.gen_optimizer.step()

方案2:拆分计算流程

如果不想保留计算图节省显存,可以调整训练流程,先完成判别器的全流程训练,再执行生成器的训练:

# process函数调整为分阶段计算损失
def process(self, images, edges, masks):
    self.iteration += 1
    outputs = self(images, edges, masks)
    logs = []

    # 判别器训练全流程
    self.dis_optimizer.zero_grad()
    dis_input_real = torch.cat((images, edges), dim=1)
    dis_input_fake = torch.cat((images, outputs.detach()), dim=1)
    dis_real, dis_real_feat = self.discriminator(dis_input_real)
    dis_fake, _ = self.discriminator(dis_input_fake)
    dis_real_loss = self.adversarial_loss(dis_real, True, True)
    dis_fake_loss = self.adversarial_loss(dis_fake, False, True)
    dis_loss = (dis_real_loss + dis_fake_loss) / 2
    dis_loss.backward()
    self.dis_optimizer.step()
    logs.append(("l_d1", dis_loss.item()))

    # 生成器训练全流程
    self.gen_optimizer.zero_grad()
    gen_input_fake = torch.cat((images, outputs), dim=1)
    gen_fake, gen_fake_feat = self.discriminator(gen_input_fake)
    gen_gan_loss = self.adversarial_loss(gen_fake, True, False)
    gen_fm_loss = 0
    for i in range(len(dis_real_feat)):
        gen_fm_loss += self.l1_loss(gen_fake_feat[i], dis_real_feat[i].detach())
    gen_fm_loss = gen_fm_loss * self.config.FM_LOSS_WEIGHT
    gen_loss = gen_gan_loss + gen_fm_loss
    gen_loss.backward()
    self.gen_optimizer.step()
    logs.extend([
        ("l_g1", gen_gan_loss.item()),
        ("l_fm", gen_fm_loss.item()),
    ])
    return outputs, gen_loss, dis_loss, logs

额外排查点

  • 检查生成器、判别器的网络结构中是否使用了带inplace=True的操作(比如nn.ReLU(inplace=True)),这类操作也会触发同类型报错,将inplace参数改为False即可
  • 如果修改后仍未定位问题,可在代码入口添加torch.autograd.set_detect_anomaly(True),运行后会直接打印触发inplace修改的具体代码位置

内容的提问来源于stack exchange,提问作者Linux Penguin

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 12:51:03