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
相关产品推荐
相关产品推荐

