一对多条件GAN(One-to-Many CGAN)训练不收敛问题求助
One-to-Many CGAN在MNIST上的收敛故障
问题背景
刚接触GAN,正在MNIST数据集上实现One-to-Many CGAN:目标是根据一个数字总和,生成对应数量的数字图像序列(生成器数量等于输出图像数)。例如,用4个生成器+1个判别器的模型,输入噪声序列并给定条件“11”时,应输出如[3,2,6]这样的数字图像序列。
当前故障
训练过程中,判别器损失持续下降至接近0,生成器损失不断上升,最终生成器仅输出无意义的噪声图像。尝试过给判别器添加Dropout层、减少滤波器数量来削弱判别器,但情况没有改善。即使将生成器数量设为1(等价于标准CGAN),模型依然无法收敛。
实现细节
自定义数据集
数据集包含X(维度为(data_size, digit_num, channels, height, width),digit_num为输出序列的数字个数)和Y(共digit_num*9+1种可能的总和标签),代码如下:
digit_num = 3 label_num = 9 * digit_num + 1 data_size = 120000 dataset = SumMNISTDataset( "mnist", 0, digit_num, data_size, transforms.Compose([transforms.Grayscale(), transforms.Normalize(127.5, 127.5)]), ) dataloader = DataLoader(dataset, batch_size, True, drop_last=True)
生成器与判别器架构
针对MNIST简化了原论文的深层网络结构,代码如下:
class Generator(nn.Module): def __init__(self, latent_dim, filter_num, label_num, embed_num=50, bias=False): super().__init__() self.pre_main = nn.Sequential( # 7 x 7 x 128 nn.ConvTranspose2d(latent_dim, filter_num * 4, 7, 1, 0, bias=bias), nn.BatchNorm2d(filter_num * 4), nn.LeakyReLU(0.2), ) self.condition = nn.Sequential( # 1 x 50 nn.Embedding(label_num, embed_num), nn.Linear(embed_num, 49, bias=bias), ) self.main = nn.Sequential( # 14 x 14 x 64 nn.ConvTranspose2d(filter_num * 4 + 1, filter_num * 2, 4, 2, 1, bias=bias), nn.BatchNorm2d(filter_num * 2), nn.LeakyReLU(0.2), # 28 x 28 x 1 nn.ConvTranspose2d(filter_num * 2, 1, 4, 2, 1, bias=bias), nn.Tanh(), ) def forward(self, x, y): y = self.condition(y).reshape(-1, 1, 7, 7) x = self.pre_main(x) x = torch.cat((x, y), dim=1) x = self.main(x) return x class Discriminator(nn.Module): def __init__(self, filter_num, label_num, embed_num=50, bias=True): super().__init__() self.condition = nn.Sequential( # 28 x 28 x 50 nn.Embedding(label_num, embed_num), nn.Linear(embed_num, 28 * 28, bias=bias), ) self.main = nn.Sequential( # 14 x 14 x 64 nn.Conv2d(2, filter_num, 3, 2, 1, bias=bias), nn.BatchNorm2d(filter_num), nn.LeakyReLU(0.2), # 7 x 7 x 128 nn.Conv2d(filter_num, filter_num * 2, 3, 2, 1, bias=bias), nn.BatchNorm2d(filter_num * 2), nn.LeakyReLU(0.2), # Dense nn.Flatten(), nn.Linear(7 * 7 * filter_num * 2, 1, bias=bias), ) def forward(self, x, y): y = self.condition(y).reshape(-1, 1, 28, 28) x = torch.cat((x, y), dim=1) x = self.main(x) return x
模型初始化
生成器数量等于输出序列的数字个数,初始化代码如下:
learning_rate = 0.0002 beta_1 = 0.5 latent_dim = 100 filter_num = 32 generator_num = digit_num omega = 1 / generator_num def weight_ini_G(model): if type(model) == nn.Linear: nn.init.constant_(model.weight.data, 1 / generator_num) elif type(model) == nn.BatchNorm2d: nn.init.constant_(model.weight.data, 1 / generator_num) nn.init.constant_(model.bias.data, 0) def weight_ini_D(model): if type(model) == nn.Linear: nn.init.normal_(model.weight.data, 0.0, 0.2) elif type(model) == nn.BatchNorm2d: nn.init.normal_(model.weight.data, 1.0, 0.2) nn.init.constant_(model.bias.data, 0) Gs = [ Generator(latent_dim, filter_num, label_num).to(device).apply(weight_ini_G) for _ in range(generator_num) ] D = Discriminator(filter_num, label_num).to(device).apply(weight_ini_D) G_optimizers = [ optim.Adam(G.parameters(), learning_rate, betas=(beta_1, 0.999)) for G in Gs ] D_optimizer = optim.Adam(D.parameters(), learning_rate, betas=(beta_1, 0.999)) bce = nn.BCEWithLogitsLoss() l1 = nn.L1Loss()
辅助函数
generate_hybrid函数用于计算序列维度上所有图像的均值,用于训练过程:
def generate_fake(): rand_labels = torch.randint(0, label_num, (batch_size, 1), device=device) images = [ Gs[g](torch.randn((batch_size, latent_dim, 1, 1), device=device), rand_labels) for g in range(generator_num) ] images = torch.stack(images, axis=1).detach_() return images, rand_labels def generate_real(): images, labels = next(iter(dataloader)) return images.to(device), labels.to(device) def generate_hybrid(images): if images.shape[0] == digit_num: images = torch.mean(images, dim=0) elif images.shape[1] == digit_num: images = torch.mean(images, dim=1) return images
生成器更新逻辑
def update_generators(real, fake): ones = torch.ones((batch_size, 1), device=device) f_images, f_labels = fake r_images, _ = real total_loss = 0 for g in range(generator_num): hybrid_fake = generate_hybrid(f_images) # r_image = r_images[:, g, :, :] preds = D(hybrid_fake, f_labels) bce_loss = bce(preds, ones) # l1_loss = l1(f_image, r_image) loss = bce_loss Gs[g].zero_grad() loss.backward() G_optimizers[g].step() total_loss += loss.item() return total_loss / generator_num
判别器更新逻辑
def update_discriminator(real, fake): half_batch_size = batch_size // 2 zeros = torch.zeros((half_batch_size, 1), device=device) ones = torch.ones((half_batch_size, 1), device=device) f_images, f_labels = fake r_images, r_labels = real f_images = f_images[:half_batch_size] f_labels = f_labels[:half_batch_size] r_images = r_images[:half_batch_size] r_labels = r_labels[:half_batch_size] total_loss = 0 # Train on Real hybrid_real = generate_hybrid(r_images) real_preds = D(hybrid_real, r_labels) bce_r_loss = bce(real_preds, ones) D.zero_grad() bce_r_loss.backward() # Train of Fake hybrid_fake = generate_hybrid(f_images) fake_preds = D(hybrid_fake, f_labels) bce_f_loss = bce(fake_preds, zeros) bce_f_loss.backward() D_optimizer.step() total_loss = (bce_f_loss.item() + bce_r_loss.item()) / 2 return total_loss
训练循环
D_losses = [] G_losses = [] epochs = 5 fixed_noise = torch.randn((4, latent_dim, 1, 1), device=device) fixed_label = torch.randint(0, label_num, (4,), device=device) for epoch in range(epochs): print(f"Epoch {epoch + 1}:") for batch in range(data_size // batch_size): # Generate Fake Images fake = generate_fake() # Generate Real Images real = generate_real() D_loss = update_discriminator(real, fake) fake = generate_fake() G_loss = update_generators(real, fake) if batch % 100 == 0: print( f"[Batch: {(batch + 1) * batch_size :7d}/{data_size} D_Loss: {D_loss} G_Loss: {G_loss}]" ) generate_image(epoch, batch, fixed_noise, fixed_label) D_losses.append(D_loss) G_losses.append(G_loss)
训练表现
训练约1个epoch后停止,生成器损失持续上升,判别器损失快速降至接近0,生成图像完全无意义。
内容的提问来源于stack exchange,提问作者Long Trần
相关产品推荐
相关产品推荐

