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训练逻辑补全了损失计算和参数更新的核心步骤,但不确定补全的部分是否合理。目前我的主要困惑点:
- 同时使用了梯度惩罚(GP)和R1正则,两者的权重设置是否合理?会不会造成正则过度?
- 生成器的噪声注入和PixelNorm的使用是否符合StyleGAN的标准逻辑?有没有可能是这部分导致的不稳定?
- 判别器同时叠加了MinibatchStdDev和MinibatchDiscrimination,会不会造成功能冗余,反而干扰判别器的稳定性?
- EMA生成器的更新逻辑是否正确?训练过程中是否应该用EMA生成器来采样评估?
真心希望各位能帮忙看看,给点改进建议,谢谢大家!
备注:内容来源于stack exchange,提问作者mchd
相关产品推荐
相关产品推荐

