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

WGAN-GP生成图像偏灰、损失持续上升问题排查求助

已解决WGAN-GP生成图像偏灰、损失上升问题

此为已解决问题,原因是生成分辨率过高导致。


问题现象

搭建的WGAN-GP生成的图像色彩暗淡偏灰,训练过程中损失值持续上升,训练10000轮无明显改善。
生成结果
训练过程

相关代码

数据预处理

def data_preprocess():
    
    batch_size = 64

    data_transforms = transforms.Compose([transforms.Resize(size=(256,256)), transforms.ToTensor(), transforms.Normalize([0.5, 0.5, 0.5],[0.5, 0.5, 0.5])])
    
    data_dir = "./256"
    data = ImageFolder(data_dir,transform = data_transforms)
    data_loader = Data.DataLoader(
        data,
        batch_size = batch_size,
        shuffle = True,
        num_workers = 0)
    
    return data_loader, data, batch_size

生成器与判别器实现

class Generator(torch.nn.Module):
    def __init__(self):
        super().__init__()

        self.main = nn.Sequential( # Z: 100
            
            nn.ConvTranspose2d(100, 1024, 4, 2, 0),
            nn.BatchNorm2d(num_features=1024),
            nn.ReLU(True),

            nn.ConvTranspose2d(1024, 512, 4, 2, 1),
            nn.BatchNorm2d(num_features=512),
            nn.ReLU(True),

            nn.ConvTranspose2d(512, 256, 4, 2, 1),
            nn.BatchNorm2d(num_features=256),
            nn.ReLU(True),

            nn.ConvTranspose2d(256, 128, 4, 2, 1),
            nn.BatchNorm2d(num_features=128),
            nn.ReLU(True),

            nn.ConvTranspose2d(128, 64, 4, 2, 1),
            nn.BatchNorm2d(num_features=64),
            nn.ReLU(True),

            nn.ConvTranspose2d(64, 32, 4, 2, 1),
            nn.BatchNorm2d(num_features=32),
            nn.ReLU(True),

            nn.ConvTranspose2d(32, 3, 4, 2, 1)) # Output of Main: (3,256,256)

        self.output = nn.Tanh()
        
    def weight_init(self,type):
        if type == "default":
            for m in self.main:
                if isinstance(m, nn.Conv2d):
                    nn.init.normal_(m.weight.data, 0, 0.02)
                elif isinstance(m, nn.BatchNorm2d):
                    nn.init.normal_(m.weight.data, 0, 0.02)
                    nn.init.constant_(m.bias.data, 0)
        elif type == "kaiming":
            for m in self.main:
                if isinstance(m, nn.Conv2d):
                    nn.init.kaiming_normal_(m.weight.data, a=0, mode='fan_in', nonlinearity='relu')
                elif isinstance(m, nn.BatchNorm2d):
                    nn.init.normal_(m.weight.data, 0, 0.02)
                    nn.init.constant_(m.bias.data, 0)
        elif type == "xavier":
            for m in self.main:
                if isinstance(m, nn.Conv2d):
                    nn.init.xavier_normal_(m.weight.data, gain=1.0)
                elif isinstance(m, nn.BatchNorm2d):
                    nn.init.normal_(m.weight.data, 0, 0.02)
                    nn.init.constant_(m.bias.data, 0)

    def forward(self, x):
        x = self.main(x)
        return self.output(x)

class Discriminator(torch.nn.Module):
    def __init__(self):
        super().__init__()
        
        self.main = nn.Sequential(

            nn.Conv2d(in_channels=3, out_channels=32, kernel_size=4, stride=2, padding=1), # (3,256,256)
            nn.LayerNorm((128,128)),
            nn.LeakyReLU(0.2, inplace=True),
            
            nn.Conv2d(in_channels=32, out_channels=64, kernel_size=4, stride=2, padding=1),
            nn.LayerNorm((64,64)),
            nn.LeakyReLU(0.2, inplace=True),

            nn.Conv2d(in_channels=64, out_channels=128, kernel_size=4, stride=2, padding=1),
            nn.LayerNorm((32,32)),
            nn.LeakyReLU(0.2, inplace=True),

            nn.Conv2d(in_channels=128, out_channels=256, kernel_size=4, stride=2, padding=1),
            nn.LayerNorm((16,16)),
            nn.LeakyReLU(0.2, inplace=True),

            nn.Conv2d(in_channels=256, out_channels=512, kernel_size=4, stride=2, padding=1),
            nn.LayerNorm((8,8)),
            nn.LeakyReLU(0.2, inplace=True),

            nn.Conv2d(in_channels=512, out_channels=1024, kernel_size=4, stride=2, padding=1), # (1024,4,4)
            nn.LayerNorm((4,4)),
            nn.LeakyReLU(0.2, inplace=True))

        self.output = nn.Sequential(
            nn.Conv2d(in_channels=1024, out_channels=1, kernel_size=4, stride=2, padding=0)
            )
        
    def weight_init(self,type):
        if type == "default":
            for m in self.main:
                if isinstance(m, nn.Conv2d):
                    nn.init.normal_(m.weight.data, 0, 0.02)
                elif isinstance(m, nn.LayerNorm):
                    nn.init.normal_(m.weight.data, 0, 0.02)
                    nn.init.constant_(m.bias.data, 0)
        elif type == "kaiming":
            for m in self.main:
                if isinstance(m, nn.Conv2d):
                    nn.init.kaiming_normal_(m.weight.data, a=0, mode='fan_in', nonlinearity='leaky_relu')
                elif isinstance(m, nn.LayerNorm):
                    nn.init.normal_(m.weight.data, 0, 0.02)
                    nn.init.constant_(m.bias.data, 0)
        elif type == "xavier":
            for m in self.main:
                if isinstance(m, nn.Conv2d):
                    nn.init.xavier_normal_(m.weight.data, gain=1.0)
                elif isinstance(m, nn.LayerNorm):
                    nn.init.normal_(m.weight.data, 0, 0.02)
                    nn.init.constant_(m.bias.data, 0)
                    
    def forward(self, x):
        x = self.main(x)
        return self.output(x)

梯度惩罚计算函数

def calculate_gradient_penalty(real_images, fake_images):
    
    t = torch.rand(real_images.size(0), 1, 1, 1).to(real_images.device)
    t = t.expand(real_images.size())

    interpolates = t * real_images + (1 - t) * fake_images
    interpolates.requires_grad_(True)

    disc_interpolates = D(interpolates)

    grad = torch.autograd.grad(outputs=disc_interpolates, 
                               inputs=interpolates,
                               grad_outputs=torch.ones_like(disc_interpolates),
                               create_graph=True, 
                               retain_graph=True)[0]

    loss_gp = (((grad.norm(2, dim=1) - 1) ** 2).mean()) * lambda_term
   
    return loss_gp

训练流程

G = Generator()
D = Discriminator()
G.weight_init("kaiming") # "default", "kaiming", "xavier"
D.weight_init("kaiming")
G.to(device)
D.to(device)

learning_rate = 1e-4
lambda_term = 10
generator_iters = 10000
data_preprocess.batch_size = 64
critic_iter = 5 # 1 Generator, 5 Discriminator
record_lenth = 0

d_optimizer = torch.optim.Adam(D.parameters(), lr=learning_rate, betas=(0.5, 0.999))
g_optimizer = torch.optim.Adam(G.parameters(), lr=learning_rate, betas=(0.5, 0.999))
data_loader, data, batch_size = data_preprocess()

d_progress = []
d_fake_progress = []
d_real_progress = []
penalty = []
g_progress = []

data = get_infinite_batches(data_loader)
one = torch.FloatTensor([1]).to(device) 
mone = (one * -1).to(device) 

for g_iter in range(generator_iters):
    
    print('----------G Iter-{}----------'.format(g_iter+1))
    
    for p in D.parameters():
        p.requires_grad = True 
        
    d_loss_real = 0
    d_loss_fake = 0
    Wasserstein_D = 0

    for d_iter in range(critic_iter):
        D.zero_grad()
            
        images = data.__next__()
        if images.size()[0] != batch_size:
            continue
        
        # 训练判别器-真实样本
        images = images.to(device)
        z = torch.randn(batch_size, 100, 1, 1).to(device)
        d_real = D(images)
        d_loss_real = d_real.mean(0).view(1)
        d_loss_real.backward(mone)
        
        # 训练判别器-生成样本
        z = torch.randn(batch_size, 100, 1, 1).to(device)
        fake_images = G(z)
        d_fake = D(fake_images)
        d_loss_fake = d_fake.mean(0).view(1)
        d_loss_fake.backward(one)
        
        # 计算梯度惩罚
        gradient_penalty = calculate_gradient_penalty(images.data, fake_images.data)
        gradient_penalty.backward()
        
        # 总损失更新
        d_loss = d_loss_fake - d_loss_real + gradient_penalty
        Wasserstein_D = d_loss_real - d_loss_fake
        d_optimizer.step()
        print('D Loss: %.6s, Fake: %.6s, Real: %.6s, Penalty: %.6s' %(d_loss.item(),d_loss_fake.item(),d_loss_real.item(),gradient_penalty.item())) 
        
        time.sleep(0.1)
        d_progress.append(d_loss.item())
        d_fake_progress.append(d_loss_fake.item())
        d_real_progress.append(d_loss_real.item())
        penalty.append(gradient_penalty.item())
        
        record_lenth += 1
        writer.add_scalars('Continue Test D Loss', {'D Loss': d_loss,
                                 'D Loss Fake': d_loss_fake,
                                 'D Loss Real': d_loss_real}, record_lenth+1)
    
    # 固定判别器权重更新生成器
    for p in D.parameters():
        p.requires_grad = False 
    
    G.zero_grad()
    z = torch.randn(batch_size, 100, 1, 1).to(device)
    fake_images = G(z)
    d_fake = D(fake_images)
    g_loss = d_fake.mean(0).view(1)
    g_loss.backward(mone) 
    g_cost = -g_loss
    g_optimizer.step()
    print('G Loss: %.6s'% g_loss.item()) 
        
    g_progress.append(g_loss.item())
    writer.add_scalar('Continue Test G Loss', g_loss, g_iter+1)

解决方案

核心原因为256分辨率下当前网络结构学习难度过高,可按以下步骤调整:

  • 降低生成分辨率:先将输出分辨率调整为64×64,删除生成器和判别器的最后两层卷积/反卷积层,验证模型收敛性后再逐步提升分辨率
  • 调整初始化方式:WGAN-GP默认适配正态分布初始化,将kaiming初始化替换为默认的normal初始化,均值设为0,标准差设为0.02
  • 修正判别器LayerNorm参数:256输入下第一层输出尺寸为128×128,LayerNorm应传入完整的归一化维度[32,128,128],后续层同理补充通道维度
  • 调整学习率:判别器学习率下调至2e-5,生成器学习率保持1e-4,避免判别器收敛过快导致生成器无法学习

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 16:18:04