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

GAN训练触发RuntimeError:梯度变量被原地操作修改问题排查

GAN训练梯度计算原地操作错误排查

问题背景

构建GAN时使用如下Generator类:

class Generator(nn.Module):
    def __init__(self, nz=128, num_classes=120, channels=3, nfilt=64):
        super(Generator, self).__init__()
        self.nz = nz
        self.num_classes = num_classes
        self.channels = channels
        
        self.label_emb = nn.Embedding(num_classes, nz)
        self.pixelnorm = PixelwiseNorm()
        self.upconv1 = UpConvBlock(2*nz, nfilt*16, num_classes, k_size=4, stride=1, padding=0, norm="cbn", dropout_p=0.15)
        self.upconv2 = UpConvBlock(nfilt*16, nfilt*8, num_classes, k_size=4, stride=2, padding=1, norm="cbn", dropout_p=0.10)
        self.upconv3 = UpConvBlock(nfilt*8, nfilt*4, num_classes, k_size=4, stride=2, padding=1, norm="cbn", dropout_p=0.05)
        self.upconv4 = UpConvBlock(nfilt*4, nfilt*2, num_classes, k_size=4, stride=2, padding=1, norm="cbn", dropout_p=0.05)
        self.upconv5 = UpConvBlock(nfilt*2, nfilt, num_classes, k_size=4, stride=2, padding=1, norm="cbn", dropout_p=0.05)
        self.self_attn = Self_Attn(nfilt)
        self.upconv6 = UpConvBlock(nfilt, 3, num_classes, k_size=3, stride=1, padding=1, norm="cbn")
        self.out_conv = spectral_norm(nn.Conv2d(3, 3, 3, 1, 1, bias=False))  
        self.out_activ = nn.Tanh()
        
    def forward(self, inputs):
        z, labels = inputs
        
        enc = self.label_emb(labels).view((-1, self.nz, 1, 1))
        enc = F.normalize(enc, p=2, dim=1)
        x = torch.cat((z, enc), 1)
        
        x = self.upconv1((x, labels))
        x = self.upconv2((x, labels))
        x = self.upconv3((x, labels))
        x = self.upconv4((x, labels))
        x = self.upconv5((x, labels))
        x = self.self_attn(x)
        x = self.upconv6((x, labels))
        x = self.out_conv(x)
        img = self.out_activ(x)           
        return img

实例化代码:

netG = Generator(nz, num_classes=len(encoded_labels), nfilt=64).to(device)

训练时触发错误:

/usr/local/lib/python3.7/dist-packages/ipykernel_launcher.py:21: UserWarning: To copy construct from a tensor, it is recommended to use sourceTensor.clone().detach() or sourceTensor.clone().detach().requires_grad_(True), rather than torch.tensor(sourceTensor).
---------------------------------------------------------------------------
RuntimeError                              Traceback (most recent call last)
<ipython-input-12-1a0c4f899e6b> in <module>
     41         errG = (torch.mean((outputR - torch.mean(outputF) + real_labels) ** 2) +
     42                 torch.mean((outputF - torch.mean(outputR) - real_labels) ** 2))/2
---&gt; 43         errG.backward()
     44         optimizerG.step()
     45 

1 frames
/usr/local/lib/python3.7/dist-packages/torch/autograd/__init__.py in backward(tensors, grad_tensors, retain_graph, create_graph, grad_variables, inputs)
    173     Variable._execution_engine.run_backward(  # Calls into the C++ engine to run the backward pass
    174         tensors, grad_tensors_, retain_graph, create_graph, inputs,
--&gt; 175         allow_unreachable=True, accumulate_grad=True)  # Calls into the C++ engine to run the backward pass
    176 
    177 def grad(

RuntimeError: one of the variables needed for gradient computation has been modified by an inplace operation: [torch.FloatTensor [1, 513, 4, 4]] is at version 2; expected version 1 instead. Hint: enable anomaly detection to find the operation that failed to compute its gradient, with torch.autograd.set_detect_anomaly(True).

已知错误含义,以下是代码中可能触发错误的点:

可能的触发点分析

  • 自定义模块的原地操作:错误提示的张量维度指向Generator前向传播中的中间特征,重点排查UpConvBlock和Self_Attn两个自定义模块:
    • 若UpConvBlock里的条件批量归一化(CBN)使用了x += ...、x.div_()这类带下划线的原地修改方法,会直接改变计算图中张量的版本,导致反向传播时版本不匹配。
    • Self_Attn模块的注意力计算过程中,如果对权重张量做了原地修改,也会触发该错误。
  • 谱归一化的隐性问题:out_conv使用的spectral_norm,在部分旧版PyTorch实现中,更新权重奇异值时可能存在隐性原地操作,干扰梯度计算。可以临时移除spectral_norm验证是否解决问题。
  • 损失计算的链式操作:损失函数errG的链式运算中,若real_labels、outputR、outputF之前被其他操作原地修改,也会导致反向传播出错。可以尝试拆解损失计算步骤,避免链式操作的隐性问题。

验证建议

  1. 在训练代码开头添加torch.autograd.set_detect_anomaly(True),运行后会打印出具体触发原地修改的操作,精准定位问题。
  2. 逐个注释自定义模块(比如先注释self_attn,再依次注释upconv系列),排查哪个模块引入了问题。
  3. 检查自定义模块实现,将所有带下划线的原地操作方法(如add_、mul_)替换为非原地版本(如x = x + ...、x = x * ...)。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.15 19:30:53