一、GAN - Generative Adversarial Networks 生成对抗网络
GAN 的本质是通过对抗博弈来隐式地学习数据的真实分布
-
结构: 由生成器 - Generator和判别器 - Discriminator两个神经网络组成。
-
原理: 生成器负责将一个低维的随机噪声向量映射到高维的数据空间中,试图拟合真实数据的概率分布;判别器则是一个二分类器,负责评估输入数据是来自真实数据集还是由生成器伪造的。
-
训练目标: 这是一个极大极小博弈 - Minimax Game过程。判别器努力最大化区分真伪数据的准确率,而生成器努力最小化判别器的准确率。当系统达到纳什均衡时,生成器所拟合的分布完美等价于真实数据分布。
I、训练生成手写文字的code时问题
- 出现了判别器过强导致纳什均衡别瞬间打破,导致无法收敛和梯度爆炸,生成模型并不能捕捉到有用特征,训练时的两个模型损失函数和生成的fake图像如下:

II、解决方案
- 降低判别器的学习率 打破两者的对称性,让判别器学得慢一点。
# 原始代码是统一的 lr = 0.0002# 修改为:lr_G = 0.0002lr_D = 0.00005 # 将判别器的学习率降低到原来的 1/4optimizer_G = optim.Adam(generator.parameters(), lr=lr_G)optimizer_D = optim.Adam(discriminator.parameters(), lr=lr_D)- 标签平滑(Label Smoothing) 不要让判别器过于自信。在计算 BCE Loss 时,不要给绝对的 和 ,而是给 和 。
# 原始代码:# real_labels = torch.ones(batch_size, 1)# fake_labels = torch.zeros(batch_size, 1)
# 修改为:real_labels = torch.ones(batch_size, 1) * 0.9 # 真实标签平滑fake_labels = torch.ones(batch_size, 1) * 0.1 # 假标签平滑- 改变训练频率 在每一个 Batch 的循环中,多训练几次生成器,少训练一次判别器。
# 在 for i, (imgs, _) in enumerate(dataloader): 循环内# 正常训练一次判别器 optimizer_D.step()
# 然后连续训练两次生成器:for _ in range(2): z = torch.randn(batch_size, latent_dim) fake_imgs = generator(z) outputs = discriminator(fake_imgs) g_loss = criterion(outputs, real_labels) # 注意这里目标依然是 real_labels (1.0 或 0.9)
optimizer_G.zero_grad() g_loss.backward() optimizer_G.step()-
削弱判别器的网络结构 可以在 Discriminator 的
nn.Sequential中加入nn.Dropout(0.3),人为增加判别器学习的难度。 -
sigmoid改用更稳定的 nn.BCEWithLogitsLoss
-
Adam 改为 GAN 常用 betas=(0.5, 0.999)
III、GAN中停止训练的时机

- GAN 的Loss 为什么不降?
在 GAN 的纳什均衡理想状态下,生成器造出的假图已经完美无瑕。此时,判别器无论怎么努力,都只能靠“瞎猜”,也就是对于任何一张图,它输出“真”的概率都是 50% ,把 0.5 代入二分类交叉熵(BCE)的公式中计算会得到一个神奇的常数:
-ln(0.5) ≈ 0.69。
- 健康的标志: 如果在训练的中后期,你的判别器 Loss(
D_loss)围绕着 0.69 上下波动(通常在 0.5 到 1.5 之间震荡),而生成器 Loss(G_loss)也保持在一个稳定的区间(通常在 0.7 到 2.0 之间),这意味着它们势均力敌,共同在进步。 - 不健康的标志: 如果某一方的 Loss 突然趋近于 0,说明平衡被打破,模型正在走向崩溃。
IV、Code
import torchimport torch.nn as nnimport torch.optim as optimimport torchvisionimport torchvision.transforms as transformsfrom torch.utils.data import DataLoaderfrom torch.utils.tensorboard import SummaryWriterimport torchvision.utils as vutils
# Basic setuptorch.manual_seed(42)device = torch.device("cuda" if torch.cuda.is_available() else "cpu")print(f"Using device: {device}")
writer = SummaryWriter("gan_logs")
# Normalize MNIST pixels to [-1, 1], matching the generator Tanh output.transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,)),])
dataset = torchvision.datasets.MNIST(root="./data", train=True, transform=transform, download=True)dataloader = DataLoader( dataset, batch_size=64, shuffle=True, drop_last=True, num_workers=0, pin_memory=(device.type == "cuda"),)
latent_dim = 100
class Generator(nn.Module): def __init__(self): super().__init__() self.model = nn.Sequential( nn.Linear(latent_dim, 256), nn.BatchNorm1d(256), nn.LeakyReLU(0.2, inplace=True), nn.Linear(256, 512), nn.BatchNorm1d(512), nn.LeakyReLU(0.2, inplace=True), nn.Linear(512, 1024), nn.BatchNorm1d(1024), nn.LeakyReLU(0.2, inplace=True), nn.Linear(1024, 28 * 28), nn.Tanh(), )
def forward(self, z): img = self.model(z) return img.view(img.size(0), 1, 28, 28)
class Discriminator(nn.Module): def __init__(self): super().__init__() self.model = nn.Sequential( nn.Linear(28 * 28, 512), nn.LeakyReLU(0.2, inplace=True), nn.Dropout(0.3), nn.Linear(512, 256), nn.LeakyReLU(0.2, inplace=True), nn.Dropout(0.3), nn.Linear(256, 1), # Return logits. Do not add Sigmoid here. )
def forward(self, img): img_flat = img.view(img.size(0), -1) return self.model(img_flat)
def weights_init(module): if isinstance(module, nn.Linear): nn.init.normal_(module.weight, mean=0.0, std=0.02) if module.bias is not None: nn.init.zeros_(module.bias)
def set_requires_grad(model, requires_grad): for param in model.parameters(): param.requires_grad_(requires_grad)
generator = Generator().to(device)discriminator = Discriminator().to(device)generator.apply(weights_init)discriminator.apply(weights_init)
# BCEWithLogitsLoss is more numerically stable than Sigmoid + BCELoss.criterion = nn.BCEWithLogitsLoss()lr_G = 2e-4lr_D = 1e-4optimizer_G = optim.Adam(generator.parameters(), lr=lr_G, betas=(0.5, 0.999))optimizer_D = optim.Adam(discriminator.parameters(), lr=lr_D, betas=(0.5, 0.999))
num_epochs = 50step = 0max_grad_norm = 1.0fixed_noise = torch.randn(64, latent_dim, device=device)
print("Start training. Run: tensorboard --logdir=gan_logs --port=6007")
try: for epoch in range(num_epochs): for i, (imgs, _) in enumerate(dataloader): real_imgs = imgs.to(device, non_blocking=True) batch_size = real_imgs.size(0)
# One-sided label smoothing: real=0.9, fake=0.0. real_labels = torch.full((batch_size, 1), 0.9, device=device) fake_labels = torch.zeros(batch_size, 1, device=device)
# 1) Train discriminator. Clear D gradients before backward. set_requires_grad(discriminator, True) optimizer_D.zero_grad(set_to_none=True)
real_logits = discriminator(real_imgs) d_loss_real = criterion(real_logits, real_labels)
z = torch.randn(batch_size, latent_dim, device=device) fake_imgs = generator(z) fake_logits = discriminator(fake_imgs.detach()) d_loss_fake = criterion(fake_logits, fake_labels)
d_loss = 0.5 * (d_loss_real + d_loss_fake) d_loss.backward() torch.nn.utils.clip_grad_norm_(discriminator.parameters(), max_grad_norm) optimizer_D.step()
# 2) Train generator with fresh noise. Only G is optimized here. set_requires_grad(discriminator, False) optimizer_G.zero_grad(set_to_none=True)
z = torch.randn(batch_size, latent_dim, device=device) gen_imgs = generator(z) gen_logits = discriminator(gen_imgs) g_loss = criterion(gen_logits, real_labels)
g_loss.backward() torch.nn.utils.clip_grad_norm_(generator.parameters(), max_grad_norm) optimizer_G.step()
if not torch.isfinite(d_loss) or not torch.isfinite(g_loss): raise RuntimeError( f"Non-finite loss detected at epoch={epoch}, batch={i}: " f"d_loss={d_loss.item()}, g_loss={g_loss.item()}" )
if i % 100 == 0: with torch.no_grad(): d_real_score = torch.sigmoid(real_logits).mean().item() d_fake_score = torch.sigmoid(fake_logits).mean().item()
print( f"[Epoch {epoch}/{num_epochs}] [Batch {i}/{len(dataloader)}] " f"[D loss: {d_loss.item():.4f}] [G loss: {g_loss.item():.4f}] " f"[D(real): {d_real_score:.3f}] [D(fake): {d_fake_score:.3f}]" )
writer.add_scalar("Loss/Discriminator", d_loss.item(), step) writer.add_scalar("Loss/Generator", g_loss.item(), step) writer.add_scalar("Score/D_real", d_real_score, step) writer.add_scalar("Score/D_fake", d_fake_score, step)
generator.eval() with torch.no_grad(): fake_visuals = generator(fixed_noise).detach().cpu() img_grid = vutils.make_grid(fake_visuals, normalize=True, value_range=(-1, 1)) writer.add_image("Generated_Fake_Images", img_grid, step) generator.train()
step += 1finally: set_requires_grad(discriminator, True) writer.close()V、拓展探索GAN的模式崩溃 - Mode Collapse
当生成器发现某一种特定的输出特别容易骗过判别器时,会放弃学习其他数据的特征,把所有的随机噪声都映射到这一个或几个固定的图像上。
1. 小批量判别 - Mini-batch Discrimination
原理: 传统的判别器是“单线程”工作的,它每次只看一张图片,判断这张是真还是假。这就导致如果生成器每次都拿同一张极其逼真的假图来骗它,它依然会给出高分。
小批量判别改变了规则:它强制判别器一次性审视一整个 Batch(比如 64 张)的图片,并计算这 64 张图片之间的相似度(方差)。
效果: 真实数据集的 64 张图片肯定长得各不相同(方差大);如果生成器发生了模式崩溃,它生成的 64 张图片会高度相似甚至一模一样(方差极小)。判别器一旦发现这批图片“长得太像了”,就会直接判定它们全是假图。这逼迫生成器必须生成多样化的图像。
2. 历史图像回放 - Experience Replay
原理: 模式崩溃有时是一种“动态崩溃”。比如生成器这会儿生成全是“1”,判别器学会了抓“1”;下一秒生成器为了逃避,又全变成了“7”,判别器又去抓“7”。生成器在几个模式之间反复横跳,就是不肯同时生成 1 到 9。
做法: 在代码中维护一个“历史假图池(Image Pool)”。每次训练判别器时,不要只拿生成器当前这一秒生成的假图,而是从历史池子里随机抽一半以前生成的假图。
效果: 判别器拥有了“长期记忆”,生成器就无法通过反复横跳来投机取巧,必须老老实实地覆盖所有模式。
3. 更换损失函数为 WGAN-GP
如果你目前的架构使用的是普通的交叉熵(BCE Loss),无论怎么调参,模式崩溃的风险都极高。目前工业界最标准的做法是直接放弃传统 GAN,改用 WGAN-GP (Wasserstein GAN with Gradient Penalty)。
原理: 传统的 JS 散度(BCE Loss 的底层本质)在衡量两个没有重叠的分布时,无法提供有效的梯度(这就是你之前遇到 100.000 的原因)。WGAN 引入了 Wasserstein 距离(也叫推土机距离),它衡量的是“把假数据分布这堆土,推到真数据分布那个坑里,需要消耗多少能量”。
效果: 彻底消灭了梯度消失问题。
-
生成器和判别器的 Loss 有了明确的指示意义(Loss 越低,图像质量真实且多样性越好)。
-
极大地缓解了模式崩溃。
4. 调整网络容量与隐变量维度 - Latent Dimension
-
隐变量维度()太小: 如果你只给生成器 10 个维度的随机噪声,它可能没有足够的“信息容量”去映射复杂的真实世界多样性。尝试把维度提升到 100 或 256。
-
生成器网络太弱: 如果判别器很复杂,而生成器只有简单的两层全连接网络,生成器在被“逼急了”的情况下只能选择模式崩溃来自保。尝试加深生成器的网络层数,或者加入 Batch Normalization。
二、VAE - Variational Autoencoder 变分自编码器
-
结构: 包含编码器 - Encoder和解码器 - Decoder

-
原理: 传统的自编码器是将数据压缩成固定向量再还原,容易导致隐空间- Latent Space不连续。VAE 不直接输出固定的隐向量,而是让编码器输出隐变量的概率分布参数。然后从该分布中采样出一个向量,交给解码器还原。

- 核心机制: 为了保证模型能生成新样本,VAE 引入了KL 散度作为正则化项,强制让编码器输出的分布逼近标准正态分布 。在推理生成阶段,我们直接从标准正态分布中随机采样,输入给解码器即可生成符合真实数据特征的新样本。
# VAE model and lossclass VAE(nn.Module): def __init__(self, latent_dim=20, hidden_dim=400): super().__init__() self.latent_dim = latent_dim self.hidden_dim = hidden_dim
self.fc1 = nn.Linear(28 * 28, hidden_dim) self.fc_mu = nn.Linear(hidden_dim, latent_dim) self.fc_logvar = nn.Linear(hidden_dim, latent_dim)
self.fc3 = nn.Linear(latent_dim, hidden_dim) self.fc4 = nn.Linear(hidden_dim, 28 * 28)
def encode(self, x): h1 = F.relu(self.fc1(x)) return self.fc_mu(h1), self.fc_logvar(h1)
def reparameterize(self, mu, logvar): std = torch.exp(0.5 * logvar) eps = torch.randn_like(std) return mu + eps * std
def decode(self, z): h3 = F.relu(self.fc3(z)) return torch.sigmoid(self.fc4(h3))
def forward(self, x): x_flat = x.view(-1, 28 * 28) mu, logvar = self.encode(x_flat) z = self.reparameterize(mu, logvar) recon_x = self.decode(z) return recon_x, mu, logvar
def vae_loss(recon_x, x, mu, logvar): bce = F.binary_cross_entropy(recon_x, x.view(-1, 28 * 28), reduction="sum") kld = -0.5 * torch.sum(1 + logvar - mu.pow(2) - logvar.exp()) return bce + kld, bce, kld
三、Diffusion Models - 扩散模型
扩散模型的核心是非平衡热力学中的马尔可夫链 - Markov Chain。
-
原理: 分为前向过程 - Forward Process和反向过程 - Reverse Process。
-
前向过程 - 加噪: 对真实数据分布经过 步连续添加高斯噪声。随着步数增加,数据逐渐失去原有结构,最终在第 步变成纯高斯噪声。这是一个确定的过程,不需要训练。
-
反向过程 - 去噪: 训练一个神经网络,输入带有噪声的图像和当前的时间步 ,预测出在这一步被添加的噪声量。通过逐步减去网络预测出的噪声,从纯高斯噪声中一步步还原出清晰的数据。由于每一步去噪都引入了条件概率,模型最终能生成极高质量且多样化的数据。
# Diffusion schedule and forward noising processdef linear_beta_schedule(timesteps, beta_start=1e-4, beta_end=0.02): return torch.linspace(beta_start, beta_end, timesteps)
def extract(values, t, x_shape): batch_size = t.shape[0] out = values.gather(-1, t.cpu()).to(t.device) return out.reshape(batch_size, *((1,) * (len(x_shape) - 1)))
betas = linear_beta_schedule(timesteps)alphas = 1.0 - betasalphas_cumprod = torch.cumprod(alphas, dim=0)alphas_cumprod_prev = F.pad(alphas_cumprod[:-1], (1, 0), value=1.0)
sqrt_recip_alphas = torch.sqrt(1.0 / alphas)sqrt_alphas_cumprod = torch.sqrt(alphas_cumprod)sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - alphas_cumprod)posterior_variance = betas * (1.0 - alphas_cumprod_prev) / (1.0 - alphas_cumprod)
def q_sample(x_start, t, noise=None): if noise is None: noise = torch.randn_like(x_start)
sqrt_alpha = extract(sqrt_alphas_cumprod, t, x_start.shape) sqrt_one_minus_alpha = extract(sqrt_one_minus_alphas_cumprod, t, x_start.shape) return sqrt_alpha * x_start + sqrt_one_minus_alpha * noise
def denormalize_images(x): return (x.clamp(-1, 1) + 1) / 2
# Denoising model. It predicts the noise added at timestep t.class SinusoidalPositionEmbeddings(nn.Module): def __init__(self, dim): super().__init__() self.dim = dim
def forward(self, time): device = time.device half_dim = self.dim // 2 embeddings = math.log(10000) / (half_dim - 1) embeddings = torch.exp(torch.arange(half_dim, device=device) * -embeddings) embeddings = time[:, None] * embeddings[None, :] embeddings = torch.cat((embeddings.sin(), embeddings.cos()), dim=-1) return embeddings
class ConvBlock(nn.Module): def __init__(self, in_channels, out_channels, time_emb_dim): super().__init__() self.time_mlp = nn.Linear(time_emb_dim, out_channels) self.block1 = nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1), nn.GroupNorm(8, out_channels), nn.SiLU(), ) self.block2 = nn.Sequential( nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1), nn.GroupNorm(8, out_channels), nn.SiLU(), )
def forward(self, x, t_emb): h = self.block1(x) time_bias = self.time_mlp(t_emb)[:, :, None, None] h = h + time_bias return self.block2(h)
class SimpleDenoiser(nn.Module): def __init__(self, image_channels=1, base_channels=64, time_emb_dim=128): super().__init__() self.time_mlp = nn.Sequential( SinusoidalPositionEmbeddings(time_emb_dim), nn.Linear(time_emb_dim, time_emb_dim), nn.SiLU(), nn.Linear(time_emb_dim, time_emb_dim), )
self.init_conv = nn.Conv2d(image_channels, base_channels, kernel_size=3, padding=1) self.down1 = ConvBlock(base_channels, base_channels, time_emb_dim) self.downsample1 = nn.Conv2d(base_channels, base_channels * 2, kernel_size=4, stride=2, padding=1) self.down2 = ConvBlock(base_channels * 2, base_channels * 2, time_emb_dim) self.downsample2 = nn.Conv2d(base_channels * 2, base_channels * 4, kernel_size=4, stride=2, padding=1)
self.mid = ConvBlock(base_channels * 4, base_channels * 4, time_emb_dim)
self.upsample1 = nn.ConvTranspose2d(base_channels * 4, base_channels * 2, kernel_size=4, stride=2, padding=1) self.up1 = ConvBlock(base_channels * 4, base_channels * 2, time_emb_dim) self.upsample2 = nn.ConvTranspose2d(base_channels * 2, base_channels, kernel_size=4, stride=2, padding=1) self.up2 = ConvBlock(base_channels * 2, base_channels, time_emb_dim)
self.out = nn.Conv2d(base_channels, image_channels, kernel_size=1)
def forward(self, x, t): t_emb = self.time_mlp(t)
x0 = self.init_conv(x) d1 = self.down1(x0, t_emb) d2 = self.down2(self.downsample1(d1), t_emb) mid = self.mid(self.downsample2(d2), t_emb)
u1 = self.upsample1(mid) u1 = torch.cat([u1, d2], dim=1) u1 = self.up1(u1, t_emb)
u2 = self.upsample2(u1) u2 = torch.cat([u2, d1], dim=1) u2 = self.up2(u2, t_emb)
return self.out(u2)

输入:一张被加噪后的图片 noisy_img,以及当前时间步 t 输出:模型预测这张图片里被加进去的噪声 predicted_noise 训练时不是让模型直接生成图片,而是让它学会:这张 noisy image 里面的噪声是什么?
整体结构
这个模型由三部分组成:
- SinusoidalPositionEmbeddings 把时间步 t 编码成一个向量。
- ConvBlock 卷积模块,每个模块都会接收图片特征和时间信息。
- SimpleDenoiser 真正的去噪网络,结构类似一个简化版 U-Net。
class SinusoidalPositionEmbeddings(nn.Module):
定义一个 PyTorch 模块,用来把时间步 t 转成时间 embedding。同一个 noisy image,在不同 t 下需要不同的去噪方式,因为不同时刻的噪声多少不同。
- 创建一个 nn.Module
- 让模型接收两个输入:noisy image 和 timestep
- 用 CNN 处理图片
- 用 embedding 处理 timestep
- 把 timestep 信息注入 CNN 特征
- 输出和图片一样形状的 predicted noise
class SimpleDenoiser(nn.Module): def __init__(self): super().__init__() # 定义时间编码 # 定义卷积网络 # 定义下采样 # 定义上采样 # 定义输出层
def forward(self, x, t): # 编码时间 t # 编码 noisy image # 下采样提取语义 # 上采样恢复分辨率 # 输出 predicted noisepredicted_noise = model(noisy_imgs, t):给模型一张加噪图片 noisy_imgs,再告诉模型当前是第几个扩散时间步 t,模型输出它认为图片中包含的噪声 predicted_noise, 让模型预测的噪声尽量接近真实加进去的噪声:loss = F.mse_loss(predicted_noise, noise)
四、CVAE - Conditional VAE
普通的 VAE 只能随机生成数据,无法控制它生成具体哪一个数字。为了能够控制生成的结果会使用 CVAE。
-
做法: 如果想生成“7”,会将“7”转换成长度为 10 的 One-hot 向量 “。
-
这个 One-hot 向量会直接拼接到原始输入数据以及潜在变量 的后面,以此来指导模型生成特定类别的结果。

# CVAE model and lossclass CVAE(nn.Module): def __init__(self, latent_dim=20, hidden_dim=400, num_classes=10): super().__init__() self.latent_dim = latent_dim self.hidden_dim = hidden_dim self.num_classes = num_classes
# Encoder input is image pixels plus the condition label. self.fc1 = nn.Linear(28 * 28 + num_classes, hidden_dim) self.fc_mu = nn.Linear(hidden_dim, latent_dim) self.fc_logvar = nn.Linear(hidden_dim, latent_dim)
# Decoder input is latent noise plus the condition label. self.fc3 = nn.Linear(latent_dim + num_classes, hidden_dim) self.fc4 = nn.Linear(hidden_dim, 28 * 28)
def labels_to_one_hot(self, labels): # Convert class ids like 3 into one-hot vectors like [0,0,0,1,0,...]. # This vector is the explicit condition used by the CVAE. return F.one_hot(labels, num_classes=self.num_classes).float()
def encode(self, x, labels): label_one_hot = self.labels_to_one_hot(labels).to(x.device) # CVAE difference from VAE: encode q(z | x, y), not q(z | x). # The encoder sees both the image and its label condition. h1 = F.relu(self.fc1(torch.cat([x, label_one_hot], dim=1))) return self.fc_mu(h1), self.fc_logvar(h1)
def reparameterize(self, mu, logvar): std = torch.exp(0.5 * logvar) eps = torch.randn_like(std) return mu + eps * std
def decode(self, z, labels): label_one_hot = self.labels_to_one_hot(labels).to(z.device) # This is the main conditional generation step: p(x | z, y). # Keeping z fixed but changing labels asks the decoder for different digits. h3 = F.relu(self.fc3(torch.cat([z, label_one_hot], dim=1))) return torch.sigmoid(self.fc4(h3))
def forward(self, x, labels): x_flat = x.view(-1, 28 * 28) mu, logvar = self.encode(x_flat, labels) z = self.reparameterize(mu, logvar) recon_x = self.decode(z, labels) return recon_x, mu, logvar
def cvae_loss(recon_x, x, mu, logvar): bce = F.binary_cross_entropy(recon_x, x.view(-1, 28 * 28), reduction="sum") kld = -0.5 * torch.sum(1 + logvar - mu.pow(2) - logvar.exp()) return bce + kld, bce, kld正在加载评论...
链上评论区