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

一对多条件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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.19 13:49:50