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

GAN生成结果波动不稳定问题求助(附完整实现代码)

GAN生成结果波动不稳定问题求助(附完整实现代码)

各位大佬好!我最近在训练一个类似StyleGAN结构的GAN,但遇到了生成结果波动特别不稳定的问题——有时候能生成勉强看得过去的图像,有时候直接崩成完全无意义的噪点。我把完整的模型实现和训练代码贴在下面,想请大家帮忙排查下问题,不管是网络结构、正则化策略、损失计算还是训练参数设置,有任何可以优化的方向都欢迎指点!

完整实现代码

import torch, os, torchvision
import torch.nn as nn
import torch.optim as optim
from torchvision import transforms, datasets, utils

class MappingNetwork(nn.Module):
    def __init__(self, latent_dim, style_dim):
        super(MappingNetwork, self).__init__()
        self.mapping = nn.Sequential(
            nn.Linear(latent_dim, style_dim),
            nn.ReLU(),
            nn.Linear(style_dim, style_dim)
        )

    def forward(self, z):
        return self.mapping(z)

class PixelNorm(nn.Module):
    def __init__(self, epsilon=1e-8):
        super(PixelNorm, self).__init__()
        self.epsilon = epsilon

    def forward(self, x):
        return x / torch.sqrt(torch.mean(x**2, dim=1, keepdim=True) + self.epsilon)

class NoiseInjection(nn.Module):
    def __init__(self):
        super(NoiseInjection, self).__init__()
        self.weight = nn.Parameter(torch.zeros(1))

    def forward(self, x, noise=None):
        if noise is None:
            batch, _, height, width = x.size()
            noise = torch.randn(batch, 1, height, width, device=x.device)
        return x + self.weight * noise

class MinibatchStdDev(nn.Module):
    def forward(self, x):
        batch_std = torch.std(x, dim=0, keepdim=True)
        batch_std = batch_std.mean().expand(x.size(0), 1, x.size(2), x.size(3))
        return torch.cat([x, batch_std], dim=1)

class MinibatchDiscrimination(nn.Module):
    def __init__(self, num_features, num_kernels, kernel_dim):
        super(MinibatchDiscrimination, self).__init__()
        self.T = nn.Parameter(torch.randn(num_features, num_kernels * kernel_dim))
        self.num_kernels = num_kernels
        self.kernel_dim = kernel_dim

    def forward(self, x):
        # Compute minibatch discrimination
        x = x @ self.T
        x = x.view(-1, self.num_kernels, self.kernel_dim)
        diffs = x.unsqueeze(0) - x.unsqueeze(1)
        abs_diffs = torch.abs(diffs).sum(-1)
        minibatch_features = torch.exp(-abs_diffs).sum(1)
        return minibatch_features

class Discriminator(nn.Module):
    def __init__(self, resolution, input_mask=False, minibatch_features=100, kernel_dim=5):
        super(Discriminator, self).__init__()
        self.input_mask = input_mask

        input_channels = 3 + 1 if self.input_mask else 3
        # Apply spectral normalization for stability
        self.from_rgb = nn.utils.spectral_norm(nn.Conv2d(input_channels, 256, kernel_size=1))

        self.blocks = nn.ModuleList()
        res = resolution
        while res > 4:
            self.blocks.append(
                nn.Sequential(
                    nn.utils.spectral_norm(nn.Conv2d(256, 256, kernel_size=3, padding=1)),
                    nn.LeakyReLU(0.2),
                    nn.AvgPool2d(kernel_size=2)
                )
            )
            res //= 2

        self.minibatch_stddev = MinibatchStdDev()

        self.intermediate_features = nn.Sequential(
            nn.Flatten(),
            nn.Linear(4 * 4 * (256 + 1), 128),  # +1 channel from minibatch stddev
            nn.LeakyReLU(0.2)
        )

        self.minibatch_discrimination = MinibatchDiscrimination(128, minibatch_features, kernel_dim)

        # Final Layer
        self.final = nn.Sequential(
            nn.Linear(128 + minibatch_features, 1),  # Append minibatch features
            nn.Tanh()  # Limit discriminator outputs to [-1, 1] for stability
        )

        self.apply(self.init_weights)

    def init_weights(self, m):
        if isinstance(m, nn.Conv2d) or isinstance(m, nn.Linear):
            nn.init.xavier_normal_(m.weight)
            if m.bias is not None:
                nn.init.constant_(m.bias, 0)

    def forward(self, img, mask=None):
        if self.input_mask and mask is not None:
            x = torch.cat((img, mask), dim=1)
        else:
            x = img

        x = self.from_rgb(x)

        for block in self.blocks:
            x = block(x)

        x = self.minibatch_stddev(x)
        x = self.intermediate_features(x)
        minibatch_features = self.minibatch_discrimination(x)
        x = torch.cat([x, minibatch_features], dim=1)
        return self.final(x)

class Generator(nn.Module):
    def __init__(self, latent_dim, style_dim, resolution, output_mask=False):
        super(Generator, self).__init__()
        self.output_mask = output_mask
        self.mapping = MappingNetwork(latent_dim, style_dim)

        self.initial = nn.Sequential(
            nn.Linear(style_dim, 4 * 4 * 512),
            nn.LeakyReLU(0.2),
            nn.Unflatten(1, (512, 4, 4)),
            PixelNorm()  # Pixel normalization for stability
        )

        self.blocks = nn.ModuleList()
        self.noise_injections = nn.ModuleList()
        self.to_rgb = nn.Sequential(
            nn.Conv2d(256, 3, kernel_size=1),
            nn.Tanh()  # Scale outputs to [-1, 1]
        )
        if self.output_mask:
            self.to_mask = nn.Conv2d(256, 1, kernel_size=1)

        in_channels = 512
        res = 4
        while res < resolution:
            out_channels = max(256, in_channels // 2)
            self.blocks.append(
                nn.Sequential(
                    nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1),
                    nn.LeakyReLU(0.2),
                    nn.Upsample(scale_factor=2),
                    PixelNorm()  # Add PixelNorm
                )
            )
            self.noise_injections.append(NoiseInjection())
            in_channels = out_channels
            res *= 2

        self.apply(self.init_weights)

    def init_weights(self, m):
        if isinstance(m, nn.Conv2d) or isinstance(m, nn.Linear):
            nn.init.kaiming_normal_(m.weight, nonlinearity='leaky_relu')
            if m.bias is not None:
                nn.init.constant_(m.bias, 0)

    def forward(self, z):
        style = self.mapping(z)
        x = self.initial(style)
        for block, noise in zip(self.blocks, self.noise_injections):
            x = block(x)
            x = noise(x)
        img = self.to_rgb(x)
        if self.output_mask:
            mask = self.to_mask(x)
            return img, mask
        return img

# Hyperparameters
latent_dim = 128
style_dim = 512
image_size = 64  # Resolution of images (e.g., 64x64)
batch_size = 16
num_epochs = 50
learning_rate_gen = 2e-4
learning_rate_disc = 1e-4
ema_decay = 0.999  # Decay rate for EMA
gp_weight = 0.5 # Weight for gradient penalty
lambda_r1 = 0.1  # Regularization weight for generator
initial_noise_std = 0.1  # Initial discriminator input noise
final_noise_std = 0.01  # Final noise after decay
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')

generator = Generator(latent_dim, style_dim, resolution=image_size).to(device)
ema_generator = Generator(latent_dim, style_dim, resolution=image_size).to(device)
ema_generator.load_state_dict(generator.state_dict())
discriminator = Discriminator(resolution=image_size).to(device)

optimizer_G = optim.Adam(generator.parameters(), lr=learning_rate_gen, betas=(0.0, 0.99))
optimizer_D = optim.Adam(discriminator.parameters(), lr=learning_rate_disc, betas=(0.0, 0.99))

scheduler_G = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer_G, T_max=150)
scheduler_D = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer_D, T_max=150)

transform = transforms.Compose([
        transforms.Resize((image_size, image_size)),
        transforms.RandomHorizontalFlip(),
        transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1),
        transforms.RandomAffine(degrees=15, translate=(0.1, 0.1)),
        transforms.RandomCrop(image_size, padding=4),
        transforms.GaussianBlur(kernel_size=(3, 3)),
        transforms.ToTensor(),
        transforms.Normalize([0.5] * 3, [0.5] * 3),
        transforms.RandomErasing(p=0.5, scale=(0.02, 0.2), ratio=(0.3, 3.3))
])

dataset_path = os.path.join(os.getcwd(), 'dataset', 'sub-data')
dataset = datasets.ImageFolder(root=dataset_path, transform=transform)
dataloader = torch.utils.data.DataLoader(dataset, batch_size=batch_size, shuffle=True)

output_dir = os.path.join(os.getcwd(), 'output')

def gradient_penalty(discriminator, real_images, fake_images, device):
    alpha = torch.rand(real_images.size(0), 1, 1, 1).to(device)
    interpolates = (alpha * real_images + (1 - alpha) * fake_images).requires_grad_(True)
    disc_interpolates = discriminator(interpolates)
    grad_outputs = torch.ones_like(disc_interpolates)
    gradients = torch.autograd.grad(
        outputs=disc_interpolates,
        inputs=interpolates,
        grad_outputs=grad_outputs,
        create_graph=True,
        retain_graph=True,
        only_inputs=True,
    )[0]
    gradients = gradients.view(gradients.size(0), -1)
    penalty = ((gradients.norm(2, dim=1) - 1) ** 2).mean()
    return penalty

def r1_regularization(discriminator, real_images, device):
    real_images.requires_grad_(True)
    outputs = discriminator(real_images)
    grad_outputs = torch.ones_like(outputs, device=device)
    gradients = torch.autograd.grad(
        outputs=outputs,
        inputs=real_images,
        grad_outputs=grad_outputs,
        create_graph=True,
        retain_graph=True,
        only_inputs=True,
    )[0]
    penalty = gradients.pow(2).sum(dim=(1, 2, 3)).mean()
    return penalty

g_losses, d_losses = [], []

for epoch in range(num_epochs):
    noise_std = initial_noise_std * (1 - epoch / num_epochs) + final_noise_std * (epoch / num_epochs)  # Gradual noise decay
    for i, (real_images, _) in enumerate(dataloader):
        real_images = real_images.to(device)
        if real_images.size(0) == 0:  # Skip empty batches
            continue

        # Train Discriminator
        optimizer_D.zero_grad()
        real_images.requires_grad_(True)
        real_images = real_images + torch.randn_like(real_images) * noise_std

        z1 = torch.randn(real_images.size(0), latent_dim, device=device)
        fake_images = generator(z1)
        real_pred = discriminator(real_images)
        fake_pred = discriminator(fake_images.detach())
        
        # 判别器损失计算(Wasserstein损失)
        d_loss = torch.mean(fake_pred) - torch.mean(real_pred)
        # 加入梯度惩罚
        gp = gradient_penalty(discriminator, real_images, fake_images, device)
        d_loss += gp_weight * gp
        # 加入R1正则
        r1_penalty = r1_regularization(discriminator, real_images, device)
        d_loss += lambda_r1 * r1_penalty
        
        d_loss.backward()
        optimizer_D.step()
        
        # Train Generator
        optimizer_G.zero_grad()
        z2 = torch.randn(real_images.size(0), latent_dim, device=device)
        fake_images = generator(z2)
        fake_pred = discriminator(fake_images)
        g_loss = -torch.mean(fake_pred)
        
        g_loss.backward()
        optimizer_G.step()
        
        # 更新EMA生成器
        for ema_param, param in zip(ema_generator.parameters(), generator.parameters()):
            ema_param.data = ema_param.data * ema_decay + param.data * (1 - ema_decay)
        
        # 记录损失
        g_losses.append(g_loss.item())
        d_losses.append(d_loss.item())
        
        # 学习率调度
        scheduler_G.step()
        scheduler_D.step()

补充说明

原代码的训练循环部分存在截断,我根据常规GAN训练逻辑补全了损失计算和参数更新的核心步骤,但不确定补全的部分是否合理。目前我的主要困惑点:

  1. 同时使用了梯度惩罚(GP)和R1正则,两者的权重设置是否合理?会不会造成正则过度?
  2. 生成器的噪声注入和PixelNorm的使用是否符合StyleGAN的标准逻辑?有没有可能是这部分导致的不稳定?
  3. 判别器同时叠加了MinibatchStdDev和MinibatchDiscrimination,会不会造成功能冗余,反而干扰判别器的稳定性?
  4. EMA生成器的更新逻辑是否正确?训练过程中是否应该用EMA生成器来采样评估?

真心希望各位能帮忙看看,给点改进建议,谢谢大家!

备注:内容来源于stack exchange,提问作者mchd

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.14 17:58:01