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

PyTorch Lightning训练GAN生成器梯度恒为0问题排查

GAN生成器梯度恒为0问题排查

问题现象

训练GAN模型时出现梯度异常,通过如下代码打印梯度值,发现生成器梯度始终为0,但判别器训练表现正常:

print(self.ganout.ganout1[0].weight.grad)

梯度打印结果参考:
生成器梯度打印结果

附完整实现代码如下:

import os
import time
from collections import OrderedDict
import numpy
import numpy as np
import torch
import torch.nn.functional as F
import torch.nn as nn
import pytorch_lightning as pyl
from torch.autograd import Variable
from torch.utils.data import Dataset, DataLoader
import matplotlib.pyplot as plt
import pandas as pd
import cv2
from PIL import Image
class cc_block(nn.Module):
    def __init__(self,in_channels, out_channels, kernel_size, stride):
        super().__init__()
        self.inn=0
        self.convv =nn.Conv2d(in_channels=in_channels, out_channels=out_channels, kernel_size=kernel_size, stride=stride)

    def forward(self, x):
        self.inn=x
        aa =self.convv(x)
        cc = torch.add(self.inn, aa)
        return cc

class ResBlock(nn.Module):
    def __init__(self, n_chans,n_chans1):
        super(ResBlock, self).__init__()
        self.conv = nn.Conv2d(n_chans, n_chans1,
                              kernel_size=(8, 8), stride=(2, 2),
                              bias=False)
        self.batch_norm = nn.BatchNorm2d(num_features=n_chans)
        torch.nn.init.kaiming_normal_(self.conv.weight,
                                      nonlinearity='relu')
        torch.nn.init.constant_(self.batch_norm.weight, 0.5)
        torch.nn.init.zeros_(self.batch_norm.bias)

    def forward(self, x):
        out = self.conv(x)
        out = self.batch_norm(out)
        out = torch.relu(out)
        return out + x

class jian(nn.Module):
    def __init__(self, ):
        super().__init__()
        self.ganj = nn.Sequential(nn.Conv2d(in_channels=3, out_channels=300, kernel_size=(8, 8), stride=(2, 2)),
                                  nn.BatchNorm2d(300), nn.LeakyReLU(0.2),
                                  nn.Conv2d(in_channels=300, out_channels=300, kernel_size=(8, 8), stride=(2, 2)),
                                  nn.BatchNorm2d(300), nn.LeakyReLU(0.2),
                                  nn.Conv2d(in_channels=300, out_channels=3, kernel_size=(8, 8), stride=(2, 2)),
                                  nn.LeakyReLU(0.2))

        self.ganj2 = nn.Sequential(nn.Linear(3 * 22 * 17, 1), nn.Sigmoid())
        self.optimizer = torch.optim.Adam(self.parameters(),lr=0.00001)

    def forward(self, x):
        a = self.ganj(x)
        b = self.ganj2(a.view(3 * 22 * 17))
        return b
class grent(nn.Module):
    def __init__(self,):
        super().__init__()
        self.ganout1 = nn.Sequential(nn.Linear(100, 6 * 22 * 17), nn.SELU())
        self.optimizer = torch.optim.Adam(self.parameters(), lr=0.00001)
        self.ganout = nn.Sequential(
            nn.ConvTranspose2d(in_channels=6, out_channels=400, kernel_size=(8, 8), stride=(2, 2)),
            nn.SELU(), nn.ConvTranspose2d(in_channels=400, out_channels=400, kernel_size=(8, 8), stride=(2, 2)),
            nn.SELU(),
            nn.ConvTranspose2d(in_channels=400, out_channels=3, kernel_size=(8, 8), stride=(2, 2)),
            nn.Sigmoid())
    def forward(self, x):
        a = self.ganout1(x)
        b = self.ganout(a.view((1,6 , 22, 17)))
        return b


class main_modle(pyl.LightningModule):
    def __init__(self, ):
        super().__init__()
        self.ganout = grent()
        self.ganj = jian()

    def forward(self, inputs):
        a =self.ganout(inputs)
        return a

    def configure_optimizers(self):
        optimizer = jian().optimizer
        optimizer2 = grent().optimizer
        return optimizer2,optimizer

    def adversarial_loss(self, y_hat, y):
        return F.binary_cross_entropy(y_hat, y)
    def training_step(self, batch, batch_idx,optimizer_idx):
        self.loss_function = nn.BCELoss()
        self.zero_grad()
        x = batch[0]
        if optimizer_idx == 1:
            if batch_idx%1==0:
                outt=self.ganj(x)
                if torch.isnan(outt).any():
                    outt =Variable(torch.tensor([1.0]),requires_grad=True).type_as(x)
                aa = torch.as_tensor(torch.Tensor([1.0])).type_as(x)
                realloss=self.adversarial_loss(outt,aa)
                aaas = torch.round(self.forward(torch.randn([100]).type_as(x) ).detach() *255)
                outt2 = self.ganj(aaas)
                if torch.isnan(outt2).any():
                    outt2 =Variable(torch.tensor([0.0]),requires_grad=True).type_as(x)
                aa2 = torch.as_tensor(torch.Tensor([0.0])).type_as(x)
                fuckloss = self.adversarial_loss(outt2, aa2)
                mainloss = (realloss +fuckloss) /2
                tqdm_dict = {"d_loss": mainloss}
                self.log_dict(tqdm_dict)
                return  mainloss
        if optimizer_idx == 0:
            self.loss_function2 = nn.MSELoss()
            outt2 = self.ganj(torch.round(self.ganout(torch.randn([100]).type_as(x))*255))
            if torch.isnan(outt2).any():
                outt2 =Variable(torch.tensor([0.0]),requires_grad=True).type_as(x)
            losss = self.adversarial_loss(outt2, torch.Tensor([1.0]).type_as(x))
            if batch_idx%100 ==0:
                aaas = torch.round(self.ganout(torch.randn([100]).type_as(x)) * 255)
                print(self.ganout.ganout1[0].weight.grad)
                img_1 = aaas[0].cpu().detach().numpy()
                img_1 = img_1.astype('uint8')
                img_1 = np.transpose(img_1, (1, 2, 0))
                cv2.imwrite('./114/'+str(batch_idx)+'.png',img_1)
            tqdm_dict = {"g_loss": losss}
            self.log_dict(tqdm_dict)
            return losss

    def test_step(self, batch, batch_idx):
        pass


class dataset2(Dataset):
    def __init__(self, csc_file):
        self.filll = os.listdir(csc_file)
        self.num =len(self.filll)
        if csc_file[-1] == '/':
            self.path = csc_file
        else:
            self.path = csc_file + '/'

    def __len__(self):
        return self.num

    def __getitem__(self, index):
        if index <= self.num:
            filllsd = self.filll[index]
            img = cv2.imread(self.path + filllsd)
            imgg = torch.FloatTensor(img)
            img = numpy.transpose(imgg, (2, 0, 1))
            input1 = torch.FloatTensor(img)
        else:
            acc = torch.rand([ 3, 218, 178]) * 255
            acc = torch.round(acc)
            input1 = acc
        return input1,


aaa =dataset2(r'K:\aaaaa\Img\img_align_celeba')
from pytorch_lightning.loggers import TensorBoardLogger
logger = TensorBoardLogger('tb_logs', name='my_model')
trainer = pyl.Trainer(gpus=1,logger =logger)
trainer.fit(model=main_modle(),train_dataloader=DataLoader(aaa,shuffle=True,batch_size=1))

错误原因

代码存在3个直接导致生成器梯度为0的问题,附带其他影响训练稳定性的bug:

  • 优化器绑定错误(核心原因)
    configure_optimizers方法中重新实例化了全新的jian()和grent()对象,优化器绑定的是这两个临时创建的新实例的参数,和当前训练用的self.ganj、self.ganout完全无关。反向传播计算出的梯度根本不会被对应的优化器接收,参数自然不会更新。
    修正写法:
def configure_optimizers(self):
    opt_d = torch.optim.Adam(self.ganj.parameters(), lr=0.00001)
    opt_g = torch.optim.Adam(self.ganout.parameters(), lr=0.00001)
    return opt_g, opt_d

同时删除jian和grent类内部定义的self.optimizer,优化器不需要写在子模块里。

  • 梯度链路被不可微操作切断
    训练生成器时,代码对生成器输出做了torch.round(...*255)操作,round是离散的不可微运算,直接把生成器输出到判别器损失之间的梯度传播路径截断,反向传播的梯度根本回不到生成器参数上。
    修正方式:训练阶段直接把生成器输出的0-1范围浮点值传给判别器,只有在保存图片做可视化的时候,再做乘255、round、转uint8的量化操作。

  • NaN处理逻辑错误
    当判别器输出出现NaN时,代码直接创建了一个新的常量张量替换输出,这一操作会直接断掉当前步的计算图,不仅会导致梯度消失,还会让训练过程极不稳定。出现NaN应该先排查前向传播的数值溢出问题(比如学习率过高、激活函数输出饱和),而不是直接替换输出值。

其他潜在bug

  • 判别器jian的forward方法中,a.view(3*22*17)写死了维度,没有保留batch维度,当前batch_size=1可以运行,一旦修改batch_size会直接报维度不匹配错误,应该改成a.view(a.shape[0], -1)。
  • 生成器grent的forward方法中,a.view((1,6,22,17))同样写死了batch维度为1,batch_size变更后会报错,应该改成a.view(a.shape[0], 6, 22, 17)。
  • 训练步中手动调用self.zero_grad()是多余的,PyTorch Lightning会在每步反向传播前自动处理梯度清零,手动调用可能会打乱梯度累积的逻辑。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.03 07:15:53