Skip to content

第 7 章 2013—2017:生成模型——VAE、GAN 与 DCGAN ​

学习目标 ​

  • 理解生成模型的目标:学习数据分布,而不是分类边界
  • 理解变分自编码器(VAE,2013)的重参数化技巧
  • 理解生成对抗网络(GAN,2014)的博弈思想
  • 实现并训练 VAE、GAN 与 DCGAN(2015)

7.1 时代背景:从判别到生成 ​

前几章的模型都是判别式的:给定输入,输出标签。但「理解数据」还有一个更高的目标——生成:从学到的分布里采出新的、看起来真实的样本。能生成,说明模型真正学到了数据的内在结构。

2013 年,Kingma 与 Welling 提出 VAE;2014 年,Goodfellow 等人提出 GAN;2015 年,Radford 等人提出 DCGAN,把卷积引入 GAN。这三者(加上 2015 年的 WGAN 改进)共同奠定了现代生成模型的基础,一路通向 2020 年后的扩散模型(第 11 章)。

7.2 VAE:把图片压缩进一个「潜空间」 ​

变分自编码器(VAE)的直觉:每张图片背后有一个低维的隐变量 z(比如「这件衣服的款式、颜色、袖长」)。编码器把图片变成 z 的分布(均值 μ、方差 σ²),解码器从 z 重建图片。

重参数化技巧是 VAE 的关键:不能直接从分布采样再反向传播(采样不可导),于是写成:

z = μ + σ·ε,   ε ~ N(0, I)

采样变成「确定性变换 + 随机噪声」,梯度可以穿过 μ、σ 传播。损失 = 重建误差(BCE)+ KL 散度(让 z 的分布接近标准正态):

KL = -0.5·Σ(1 + log σ² - μ² - σ²)

examples/ch07_vae.py 在 Fashion-MNIST 上训练 20 维潜变量 VAE:

VAE 参数量:439,480
epoch 1:loss = 302.386
epoch 2:loss = 262.775
epoch 3:loss = 253.979
重建 MSE = 0.02397
生成图片像素均值 = 0.214,标准差 = 0.249

重建 MSE 低说明编码-解码通路有效;直接从标准正态采样 z 也能生成合理分布的新图片。VAE 的优点是训练稳定、潜空间连续(可以插值),缺点是生成图像偏模糊——因为 MSE/BCE 鼓励「平均」而非「锐利」。

7.3 GAN:造假者与鉴定师的博弈 ​

生成对抗网络(GAN)由两个网络组成:

  • 生成器 G:把随机噪声 z 变成图片(造假者);
  • 判别器 D:判断一张图是真图还是生成图(鉴定师)。

两者玩一个极小极大博弈:

min_G max_D  E[log D(x)] + E[log(1 - D(G(z)))]

G 想骗过 D,D 想识破 G。理想平衡点:D 的准确率 = 50%(完全无法区分),此时 G 学到了真实分布。

examples/ch07_gan.py 用全连接 GAN 在 MNIST 上训练:

epoch 1:D loss = 0.769,G loss = 1.664,假图被判真的比例 = 32.7%
epoch 2:D loss = 0.773,G loss = 1.711,假图被判真的比例 = 19.8%
epoch 3:D loss = 0.924,G loss = 1.391,假图被判真的比例 = 24.4%
epoch 4:D loss = 0.573,G loss = 1.979,假图被判真的比例 = 9.0%
epoch 5:D loss = 0.313,G loss = 3.134,假图被判真的比例 = 3.6%
真实图片像素均值 = -0.744,生成图片像素均值 = -0.506

注意 D loss 与 G loss 此消彼长、并不单调下降——GAN 训练不稳定是常态,这正是它著名的难点。生成图片的像素统计逐渐接近真实分布(真实均值 -0.744,生成 -0.506),说明 G 在学,但 5 个 epoch 还不足以让两者达到平衡。

7.4 DCGAN:卷积版 GAN ​

2015 年,Radford 等人把卷积引入 GAN,提出 DCGAN,并给出四条经验法则:

  1. 判别器用卷积 + LeakyReLU,不用池化(用 stride 卷积下采样);
  2. 生成器用转置卷积上采样,加 BatchNorm;
  3. 全连接层尽量少(只在输入处);
  4. 生成器输出用 Tanh。

examples/ch07_dcgan.py 在 MNIST 上训练 DCGAN:

epoch 1:D loss = 0.352,G loss = 2.998,假图被判真比例 = 2.5%
epoch 2:D loss = 0.284,G loss = 2.556,假图被判真比例 = 1.5%
epoch 3:D loss = 0.247,G loss = 2.892,假图被判真比例 = 1.5%
epoch 4:D loss = 0.257,G loss = 2.922,假图被判真比例 = 1.8%
生成图片形状:(16, 1, 28, 28),像素范围 [-1.00, 1.00]

判别器收敛得很快(前 4 轮就把假图判真比例压到 2% 左右),生成器还在追赶——这是 GAN 训练中经典的「D 领先」阶段。延长训练、调低 D 的学习率、使用 WGAN 的 Wasserstein 损失,都能改善平衡。

7.5 三者的定位 ​

模型核心思想优点缺点
VAE编码-解码 + 潜变量训练稳定、潜空间连续图像模糊
GAN生成器与判别器博弈图像锐利训练不稳定、模式崩塌
DCGAN卷积 + 博弈图像质量更高仍需小心调参

它们之后的融合思路(如 VAE-GAN、WGAN-GP、StyleGAN)以及 2020 年扩散模型的崛起,都建立在这三块基石之上。

完整代码 ​

本章用到的完整示例代码:

examples/ch07_vae.py ​

py
"""第 7 章示例(2013):变分自编码器(VAE)。

编码器把图片压缩成隐变量 z 的分布(均值+方差),解码器从 z 重建图片。
KL 项让 z 的分布接近标准正态,从而可以直接从噪声采样生成新图。

运行方式:
    uv run python examples/ch07_vae.py
"""
import torch
from torch import nn
from torch.utils.data import DataLoader
from torchvision import datasets, transforms

torch.manual_seed(0)
device = "cuda" if torch.cuda.is_available() else "cpu"

transform = transforms.ToTensor()
train = datasets.FashionMNIST("data", train=True, download=True, transform=transform)
test = datasets.FashionMNIST("data", train=False, download=True, transform=transform)
loader = DataLoader(train, batch_size=128, shuffle=True)


class VAE(nn.Module):
    def __init__(self, latent=20):
        super().__init__()
        self.encoder = nn.Sequential(
            nn.Flatten(), nn.Linear(784, 256), nn.ReLU(), nn.Linear(256, 64), nn.ReLU(),
        )
        self.mu = nn.Linear(64, latent)
        self.logvar = nn.Linear(64, latent)
        self.decoder = nn.Sequential(
            nn.Linear(latent, 64), nn.ReLU(),
            nn.Linear(64, 256), nn.ReLU(),
            nn.Linear(256, 784), nn.Sigmoid(),
        )

    def reparameterize(self, mu, logvar):
        std = logvar.exp().sqrt()
        return mu + std * torch.randn_like(std)

    def forward(self, x):
        h = self.encoder(x)
        mu, logvar = self.mu(h), self.logvar(h)
        z = self.reparameterize(mu, logvar)
        return self.decoder(z), mu, logvar


def loss_fn(recon, x, mu, logvar):
    bce = nn.functional.binary_cross_entropy(recon, x.view(-1, 784), reduction="sum")
    kl = -0.5 * torch.sum(1 + logvar - mu**2 - logvar.exp())
    return (bce + kl) / x.shape[0]


model = VAE().to(device)
opt = torch.optim.Adam(model.parameters(), lr=1e-3)
print(f"VAE 参数量:{sum(p.numel() for p in model.parameters()):,}")

for epoch in range(3):
    total_loss = 0
    for xb, _ in loader:
        xb = xb.to(device)
        recon, mu, logvar = model(xb)
        loss = loss_fn(recon, xb, mu, logvar)
        opt.zero_grad()
        loss.backward()
        opt.step()
        total_loss += loss.item() * len(xb)
    print(f"epoch {epoch + 1}:loss = {total_loss / len(train):.3f}")

# 重建质量:取 8 张测试图
xb = torch.stack([test[i][0] for i in range(8)]).to(device)
recon, _, _ = model(xb)
mse = nn.functional.mse_loss(recon, xb.view(8, -1)).item()
print(f"重建 MSE = {mse:.5f}")

# 生成:从标准正态采样 z
z = torch.randn(8, 20, device=device)
gen = model.decoder(z)
print(f"生成图片像素均值 = {gen.mean().item():.3f},标准差 = {gen.std().item():.3f}")

examples/ch07_gan.py ​

py
"""第 7 章示例(2014):生成对抗网络(GAN)在 MNIST 上训练。

生成器把随机噪声变成图片,判别器判断图片真假;两者对抗训练。
注意 GAN 训练不稳定是常态,D/G 损失此消彼长。

运行方式:
    uv run python examples/ch07_gan.py
"""
import torch
from torch import nn
from torch.utils.data import DataLoader
from torchvision import datasets, transforms

torch.manual_seed(0)
device = "cuda" if torch.cuda.is_available() else "cpu"

transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize((0.5,), (0.5,)),   # 归一化到 [-1, 1]
])
train = datasets.MNIST("data", train=True, download=True, transform=transform)
loader = DataLoader(train, batch_size=128, shuffle=True)


class Generator(nn.Module):
    def __init__(self, latent=100):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(latent, 256), nn.LeakyReLU(0.2),
            nn.Linear(256, 512), nn.LeakyReLU(0.2),
            nn.Linear(512, 784), nn.Tanh(),
        )

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


class Discriminator(nn.Module):
    def __init__(self):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(784, 512), nn.LeakyReLU(0.2), nn.Dropout(0.3),
            nn.Linear(512, 256), nn.LeakyReLU(0.2), nn.Dropout(0.3),
            nn.Linear(256, 1),
        )

    def forward(self, x):
        return self.net(x)


G = Generator().to(device)
D = Discriminator().to(device)
opt_g = torch.optim.Adam(G.parameters(), lr=1e-4)
opt_d = torch.optim.Adam(D.parameters(), lr=1e-4)
loss_fn = nn.BCEWithLogitsLoss()

for epoch in range(5):
    g_loss_sum = d_loss_sum = fake_ratio = 0
    for xb, _ in loader:
        real = xb.view(-1, 784).to(device)
        z = torch.randn(xb.shape[0], 100, device=device)
        fake = G(z)

        opt_d.zero_grad()
        d_real = D(real)
        d_fake = D(fake.detach())
        d_loss = loss_fn(d_real, torch.ones_like(d_real)) + loss_fn(d_fake, torch.zeros_like(d_fake))
        d_loss.backward()
        opt_d.step()

        opt_g.zero_grad()
        g_loss = loss_fn(D(fake), torch.ones_like(fake[:, :1]))
        g_loss.backward()
        opt_g.step()

        g_loss_sum += g_loss.item() * len(xb)
        d_loss_sum += d_loss.item() * len(xb)
        fake_ratio += (torch.sigmoid(d_fake) > 0.5).float().mean().item() * len(xb)

    n = len(train)
    print(f"epoch {epoch + 1}:D loss = {d_loss_sum / n:.3f},G loss = {g_loss_sum / n:.3f},"
          f"假图被判真的比例 = {fake_ratio / n * 100:.1f}%")

# 采样对比:真实图片与生成图片的像素统计
real_mean = (train.data[:1000].float() / 127.5 - 1).mean().item()
with torch.no_grad():
    samples = G(torch.randn(1000, 100, device=device))
print(f"真实图片像素均值 = {real_mean:.3f},生成图片像素均值 = {samples.mean().item():.3f}")

examples/ch07_dcgan.py ​

py
"""第 7 章示例(2015):DCGAN——用卷积搭建的 GAN。

生成器和判别器都用卷积/转置卷积,结构更稳、图像质量更好,
是 GAN 从玩具走向图像生成的关键一步。

运行方式:
    uv run python examples/ch07_dcgan.py
"""
import torch
from torch import nn
from torch.utils.data import DataLoader
from torchvision import datasets, transforms

torch.manual_seed(0)
device = "cuda" if torch.cuda.is_available() else "cpu"

transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize((0.5,), (0.5,)),
])
train = datasets.MNIST("data", train=True, download=True, transform=transform)
loader = DataLoader(train, batch_size=128, shuffle=True)


class Generator(nn.Module):
    def __init__(self, latent=100):
        super().__init__()
        self.fc = nn.Sequential(
            nn.Linear(latent, 128 * 7 * 7), nn.BatchNorm1d(128 * 7 * 7), nn.ReLU(),
        )
        self.deconv = nn.Sequential(
            nn.ConvTranspose2d(128, 64, 4, 2, 1), nn.BatchNorm2d(64), nn.ReLU(),   # 7→14
            nn.ConvTranspose2d(64, 1, 4, 2, 1), nn.Tanh(),                          # 14→28
        )

    def forward(self, z):
        return self.deconv(self.fc(z).view(-1, 128, 7, 7))


class Discriminator(nn.Module):
    def __init__(self):
        super().__init__()
        self.net = nn.Sequential(
            nn.Conv2d(1, 64, 4, 2, 1), nn.LeakyReLU(0.2),                           # 28→14
            nn.Conv2d(64, 128, 4, 2, 1), nn.BatchNorm2d(128), nn.LeakyReLU(0.2),    # 14→7
            nn.Flatten(), nn.Linear(128 * 7 * 7, 1),
        )

    def forward(self, x):
        return self.net(x)


G = Generator().to(device)
D = Discriminator().to(device)
opt_g = torch.optim.Adam(G.parameters(), lr=2e-4, betas=(0.5, 0.999))
opt_d = torch.optim.Adam(D.parameters(), lr=2e-4, betas=(0.5, 0.999))
loss_fn = nn.BCEWithLogitsLoss()

for epoch in range(4):
    g_sum = d_sum = fake_ratio = 0
    for xb, _ in loader:
        real = xb.to(device)
        z = torch.randn(xb.shape[0], 100, device=device)
        fake = G(z)

        opt_d.zero_grad()
        d_loss = loss_fn(D(real), torch.ones_like(D(real))) \
               + loss_fn(D(fake.detach()), torch.zeros_like(D(fake.detach())))
        d_loss.backward()
        opt_d.step()

        opt_g.zero_grad()
        g_loss = loss_fn(D(fake), torch.ones_like(D(fake)))
        g_loss.backward()
        opt_g.step()

        g_sum += g_loss.item() * len(xb)
        d_sum += d_loss.item() * len(xb)
        fake_ratio += (torch.sigmoid(D(fake.detach())) > 0.5).float().mean().item() * len(xb)

    n = len(train)
    print(f"epoch {epoch + 1}:D loss = {d_sum / n:.3f},G loss = {g_sum / n:.3f},"
          f"假图被判真比例 = {fake_ratio / n * 100:.1f}%")

with torch.no_grad():
    samples = G(torch.randn(16, 100, device=device))
print(f"生成图片形状:{tuple(samples.shape)},像素范围 [{samples.min().item():.2f}, {samples.max().item():.2f}]")

动手实践 ​

  1. 运行三个示例,比较 VAE 的重建 MSE 与 GAN 的生成像素统计。
  2. 把 VAE 的潜变量维度从 20 改成 2,训练后观察重建质量;再改成 64,对比。
  3. 在 examples/ch07_gan.py 中把判别器学习率调成生成器的 1/2,观察「假图被判真比例」是否更平稳。
  4. 给 VAE 的 KL 项乘 0.1(减弱正则),观察生成样本的多样性变化。

常见错误 ​

错误写法现象原因
GAN 里 fake.detach() 忘记G 和 D 的梯度互相污染训练 D 时假图要 detach,否则 G 也被更新
BCEWithLogitsLoss 输入未压缩结果发散该损失自带 sigmoid;别在 D 里再套 sigmoid
VAE 采样直接 torch.randn 不进图梯度断掉必须用重参数化 μ + σ·ε
生成器输出没有 Tanh像素范围不对G 输出应匹配归一化后的 [-1, 1]
D 判真比例长期 0% 或 100%训练死了检查学习率、BN、标签平滑
只用 loss 判断 GAN 好坏误判看生成的样本图与像素统计

章末练习 ​

基础

  1. 判别式模型与生成式模型的区别是什么?
  2. 重参数化技巧解决了什么问题?
  3. GAN 的均衡点为什么是 D 准确率 50%?

提高

  1. 推导 VAE 的 KL 项公式(N(μ,σ²) 与 N(0,1) 之间的 KL 散度)。
  2. 把 DCGAN 的生成器改为 3 层转置卷积(7→14→28→56),说明输出尺寸变化,并调整判别器适配。
  3. 实验:GAN 训练中 D 领先(G loss 上升)时,降低 D 学习率,观察是否恢复。

挑战

  1. 实现 WGAN-GP(用 Wasserstein 距离 + 梯度惩罚),与标准 GAN 对比训练稳定性。
  2. 实现 VAE 潜空间插值:取两张图片的 z,线性插值后解码,观察中间样本是否自然。

章末自测 ​

  1. VAE 的隐变量编码是?
    • A. 一个确定向量
    • B. 一个分布(μ, σ²)
    • C. 一张图片
    • D. 一个标签
  2. 重参数化技巧中 z 的表达式是?
    • A. z = μ
    • B. z = μ + σ·ε
    • C. z = σ·ε
    • D. z = μ·σ
  3. GAN 由哪两个网络组成?
    • A. 编码器与解码器
    • B. 生成器与判别器
    • C. Q 网络与 V 网络
    • D. Teacher 与 Student
  4. 训练判别器时,假图应该?
    • A. 直接用于更新生成器
    • B. detach 后使用
    • C. 不用假图
    • D. 加倍权重
  5. DCGAN 中生成器用什么上采样?
    • A. 插值
    • B. 转置卷积
    • C. 反池化
    • D. 全连接
  6. VAE 损失包含哪两部分?
    • A. 分类损失 + 正则
    • B. 重建损失 + KL 散度
    • C. MSE + 交叉熵
    • D. 对抗损失 + 重建
  7. GAN 训练不稳定的常见表现?
    • A. 损失单调下降
    • B. D/G 此消彼长、模式崩塌
    • C. 永远收敛
    • D. 无法开始训练
  8. DCGAN 判别器下采样用什么?
    • A. 池化
    • B. stride 卷积
    • C. 全连接
    • D. 转置卷积
  9. 生成器输出常用激活是?
    • A. ReLU
    • B. Tanh
    • C. Sigmoid
    • D. Softmax
  10. 2014 年提出 GAN 的论文作者是?
    • A. Kingma 与 Welling
    • B. Goodfellow 等人
    • C. Radford 等人
    • D. He 等人