PyTorch1.5+运行旧图像生成模型出现inplace操作梯度报错如何解决
PyTorch版本兼容报错修复方案
错误根因
PyTorch 1.5 版本新增了张量版本校验机制,对inplace修改计算图依赖张量的行为检查更严格。你当前代码的问题出在backward函数的执行顺序:
- 你在
process函数中已经基于初始状态的判别器参数,同时计算出了dis_loss和gen_loss - 原有
backward逻辑先执行dis_optimizer.step(),这个操作会inplace修改判别器的参数张量,导致gen_loss对应的计算图中依赖的判别器张量版本发生变化 - 后续执行
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
相关产品推荐
相关产品推荐

