GAN(生成对抗网络)深度学习入门:从“假钞制造”到代码实战(代码拿来就能用)

1. 引言:什么是 GAN?

想象一下,你正在玩一个"真假鉴定"游戏。一位画家(生成器)努力画出一幅足以乱真的《蒙娜丽莎》,而一位鉴定专家(判别器)则试图找出画中的破绽。随着游戏进行,画家的技巧越来越高超,鉴定专家的眼光也越来越毒辣。最终,画家画出的作品连专家都难辨真伪------这就是 GAN(Generative Adversarial Network,生成对抗网络)的核心思想。

GAN 是深度学习领域最具革命性的生成模型之一,由 Ian Goodfellow 等人于 2014 年提出。它通过让两个神经网络(生成器和判别器)相互对抗、共同进化,从而学会生成逼真的数据,如图像、音乐甚至文本。

为什么 GAN 如此重要?

  • 逼真的生成能力:能够生成以假乱真的图片、视频、语音。
  • 无监督学习:不需要人工标注的标签,直接从数据中学习分布。
  • 广泛应用:图像生成、风格迁移、数据增强、超分辨率、艺术创作等。

本文将从生活化的比喻出发,详细讲解 GAN 的原理,并结合一个完整的 PyTorch 代码示例(生成手写数字),带你一步步理解并实现自己的第一个 GAN。

2. GAN 原理解析:一场"造假"与"打假"的博弈

2.1 核心角色:生成器与判别器

  • 生成器(Generator,G):好比"假钞制造者"。它的目标是接收一个随机噪声(比如一堆乱码),通过神经网络"画"出一张逼真的图像,让判别器误以为这是真图。
  • 判别器(Discriminator,D):好比"警察"或"验钞机"。它的目标是判断输入图像是来自真实数据集(真钞)还是生成器伪造的(假钞)。

2.2 训练过程:对抗与进化

GAN 的训练是一个动态博弈过程,可以用以下流程图直观展示:
#mermaid-svg-glB0WJMARsEDI5n0{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-glB0WJMARsEDI5n0 .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-glB0WJMARsEDI5n0 .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-glB0WJMARsEDI5n0 .error-icon{fill:#552222;}#mermaid-svg-glB0WJMARsEDI5n0 .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-glB0WJMARsEDI5n0 .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-glB0WJMARsEDI5n0 .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-glB0WJMARsEDI5n0 .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-glB0WJMARsEDI5n0 .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-glB0WJMARsEDI5n0 .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-glB0WJMARsEDI5n0 .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-glB0WJMARsEDI5n0 .marker{fill:#333333;stroke:#333333;}#mermaid-svg-glB0WJMARsEDI5n0 .marker.cross{stroke:#333333;}#mermaid-svg-glB0WJMARsEDI5n0 svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-glB0WJMARsEDI5n0 p{margin:0;}#mermaid-svg-glB0WJMARsEDI5n0 .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-glB0WJMARsEDI5n0 .cluster-label text{fill:#333;}#mermaid-svg-glB0WJMARsEDI5n0 .cluster-label span{color:#333;}#mermaid-svg-glB0WJMARsEDI5n0 .cluster-label span p{background-color:transparent;}#mermaid-svg-glB0WJMARsEDI5n0 .label text,#mermaid-svg-glB0WJMARsEDI5n0 span{fill:#333;color:#333;}#mermaid-svg-glB0WJMARsEDI5n0 .node rect,#mermaid-svg-glB0WJMARsEDI5n0 .node circle,#mermaid-svg-glB0WJMARsEDI5n0 .node ellipse,#mermaid-svg-glB0WJMARsEDI5n0 .node polygon,#mermaid-svg-glB0WJMARsEDI5n0 .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-glB0WJMARsEDI5n0 .rough-node .label text,#mermaid-svg-glB0WJMARsEDI5n0 .node .label text,#mermaid-svg-glB0WJMARsEDI5n0 .image-shape .label,#mermaid-svg-glB0WJMARsEDI5n0 .icon-shape .label{text-anchor:middle;}#mermaid-svg-glB0WJMARsEDI5n0 .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-glB0WJMARsEDI5n0 .rough-node .label,#mermaid-svg-glB0WJMARsEDI5n0 .node .label,#mermaid-svg-glB0WJMARsEDI5n0 .image-shape .label,#mermaid-svg-glB0WJMARsEDI5n0 .icon-shape .label{text-align:center;}#mermaid-svg-glB0WJMARsEDI5n0 .node.clickable{cursor:pointer;}#mermaid-svg-glB0WJMARsEDI5n0 .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-glB0WJMARsEDI5n0 .arrowheadPath{fill:#333333;}#mermaid-svg-glB0WJMARsEDI5n0 .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-glB0WJMARsEDI5n0 .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-glB0WJMARsEDI5n0 .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-glB0WJMARsEDI5n0 .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-glB0WJMARsEDI5n0 .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-glB0WJMARsEDI5n0 .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-glB0WJMARsEDI5n0 .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-glB0WJMARsEDI5n0 .cluster text{fill:#333;}#mermaid-svg-glB0WJMARsEDI5n0 .cluster span{color:#333;}#mermaid-svg-glB0WJMARsEDI5n0 div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-glB0WJMARsEDI5n0 .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-glB0WJMARsEDI5n0 rect.text{fill:none;stroke-width:0;}#mermaid-svg-glB0WJMARsEDI5n0 .icon-shape,#mermaid-svg-glB0WJMARsEDI5n0 .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-glB0WJMARsEDI5n0 .icon-shape p,#mermaid-svg-glB0WJMARsEDI5n0 .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-glB0WJMARsEDI5n0 .icon-shape .label rect,#mermaid-svg-glB0WJMARsEDI5n0 .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-glB0WJMARsEDI5n0 .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-glB0WJMARsEDI5n0 .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-glB0WJMARsEDI5n0 :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 训练判别器D
训练生成器G


开始训练
初始化生成器G与判别器D
训练阶段
固定生成器G
输入真实图像x
判别器D(x) → 真/假
计算D损失: L_D = log(D(x)) + log(1-D(G(z)))
更新判别器参数
固定判别器D
输入随机噪声z
生成器G(z) → 假图像
判别器D(G(z)) → 真/假
计算G损失: L_G = log(1-D(G(z)))
更新生成器参数
达到收敛条件?
训练完成

生成器可生成逼真图像

GAN 的训练是一个动态博弈过程:

  1. 固定生成器,训练判别器

    • 给判别器看一批真实图片,告诉它"这些是真的(标签为 1)"。
    • 给判别器看一批生成器造的假图片,告诉它"这些是假的(标签为 0)"。
    • 判别器通过比较预测结果与真实标签,更新自己的参数,提高鉴别能力。
  2. 固定判别器,训练生成器

    • 生成器造出一批假图片,但这次我们希望判别器把它们判断为"真(标签为 1)"。
    • 生成器根据判别器的"误判"程度来更新自己的参数,让自己画得更逼真。
  3. 循环往复

    • 两者不断对抗、相互提升。理想情况下,最终生成器能生成与真实数据分布几乎一致的样本,而判别器则无法区分真假(输出概率接近 0.5)。

2.3 数学本质:最小化 JS 散度

从数学上看,GAN 的训练目标是一个极小极大博弈(Minimax Game)

min⁡Gmax⁡DV(D,G)=Ex∼pdata(x)log⁡D(x)+Ez∼pz(z)log⁡(1−D(G(z))) \min_G \max_D V(D, G) = \mathbb{E}{x \sim p{data}(x)}\\log D(x) + \mathbb{E}_{z \sim p_z(z)}\\log (1 - D(G(z))) GminDmaxV(D,G)=Ex∼pdata(x)logD(x)+Ez∼pz(z)log(1−D(G(z)))

  • D(x)D(x)D(x):判别器认为真实样本 xxx 为真的概率。
  • G(z)G(z)G(z):生成器根据噪声 zzz 生成的假样本。
  • D(G(z))D(G(z))D(G(z)):判别器认为假样本为真的概率。

生成器 G 希望 D(G(z))D(G(z))D(G(z)) 越大越好(假图被判为真),即最小化 log⁡(1−D(G(z)))\log(1 - D(G(z)))log(1−D(G(z)))。

判别器 D 希望 D(x)D(x)D(x) 越大越好(真图被判为真),且 D(G(z))D(G(z))D(G(z)) 越小越好(假图被判为假),即最大化整个式子。

通过这种对抗,生成器最终学会逼近真实数据的分布。

3. 生活中的 GAN 例子

3.1 艺术品伪造与鉴定

正如开头提到的画家与鉴定专家,历史上许多赝品正是通过不断模仿、被专家指出破绽、再改进的过程,最终达到几乎无法识别的程度。GAN 的学习过程与此高度相似。

3.2 游戏 NPC 的进化

在一些对抗性游戏中,AI 控制的角色(生成器)会不断尝试新的策略来击败玩家(判别器),而玩家也会适应 AI 的新策略。两者在对抗中共同进化,使得游戏体验更加丰富。

3.3 反欺诈系统

银行的反欺诈模型(判别器)需要识别异常交易(假样本),而欺诈者(生成器)会不断设计新的欺诈手段。两者的对抗促使反欺诈系统越来越智能。

4. 代码实战:用 PyTorch 实现手写数字生成

我们将使用 PyTorch 框架,基于 sklearn 内置的 digits 数据集(8×8 手写数字),构建一个完整的 GAN 模型,生成 64×64 的手写数字图像。

GAN 整体架构图:
#mermaid-svg-Emko39JYPZ0vO86Q{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-Emko39JYPZ0vO86Q .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-Emko39JYPZ0vO86Q .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-Emko39JYPZ0vO86Q .error-icon{fill:#552222;}#mermaid-svg-Emko39JYPZ0vO86Q .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-Emko39JYPZ0vO86Q .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-Emko39JYPZ0vO86Q .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-Emko39JYPZ0vO86Q .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-Emko39JYPZ0vO86Q .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-Emko39JYPZ0vO86Q .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-Emko39JYPZ0vO86Q .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-Emko39JYPZ0vO86Q .marker{fill:#333333;stroke:#333333;}#mermaid-svg-Emko39JYPZ0vO86Q .marker.cross{stroke:#333333;}#mermaid-svg-Emko39JYPZ0vO86Q svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-Emko39JYPZ0vO86Q p{margin:0;}#mermaid-svg-Emko39JYPZ0vO86Q .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-Emko39JYPZ0vO86Q .cluster-label text{fill:#333;}#mermaid-svg-Emko39JYPZ0vO86Q .cluster-label span{color:#333;}#mermaid-svg-Emko39JYPZ0vO86Q .cluster-label span p{background-color:transparent;}#mermaid-svg-Emko39JYPZ0vO86Q .label text,#mermaid-svg-Emko39JYPZ0vO86Q span{fill:#333;color:#333;}#mermaid-svg-Emko39JYPZ0vO86Q .node rect,#mermaid-svg-Emko39JYPZ0vO86Q .node circle,#mermaid-svg-Emko39JYPZ0vO86Q .node ellipse,#mermaid-svg-Emko39JYPZ0vO86Q .node polygon,#mermaid-svg-Emko39JYPZ0vO86Q .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-Emko39JYPZ0vO86Q .rough-node .label text,#mermaid-svg-Emko39JYPZ0vO86Q .node .label text,#mermaid-svg-Emko39JYPZ0vO86Q .image-shape .label,#mermaid-svg-Emko39JYPZ0vO86Q .icon-shape .label{text-anchor:middle;}#mermaid-svg-Emko39JYPZ0vO86Q .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-Emko39JYPZ0vO86Q .rough-node .label,#mermaid-svg-Emko39JYPZ0vO86Q .node .label,#mermaid-svg-Emko39JYPZ0vO86Q .image-shape .label,#mermaid-svg-Emko39JYPZ0vO86Q .icon-shape .label{text-align:center;}#mermaid-svg-Emko39JYPZ0vO86Q .node.clickable{cursor:pointer;}#mermaid-svg-Emko39JYPZ0vO86Q .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-Emko39JYPZ0vO86Q .arrowheadPath{fill:#333333;}#mermaid-svg-Emko39JYPZ0vO86Q .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-Emko39JYPZ0vO86Q .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-Emko39JYPZ0vO86Q .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-Emko39JYPZ0vO86Q .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-Emko39JYPZ0vO86Q .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-Emko39JYPZ0vO86Q .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-Emko39JYPZ0vO86Q .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-Emko39JYPZ0vO86Q .cluster text{fill:#333;}#mermaid-svg-Emko39JYPZ0vO86Q .cluster span{color:#333;}#mermaid-svg-Emko39JYPZ0vO86Q div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-Emko39JYPZ0vO86Q .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-Emko39JYPZ0vO86Q rect.text{fill:none;stroke-width:0;}#mermaid-svg-Emko39JYPZ0vO86Q .icon-shape,#mermaid-svg-Emko39JYPZ0vO86Q .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-Emko39JYPZ0vO86Q .icon-shape p,#mermaid-svg-Emko39JYPZ0vO86Q .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-Emko39JYPZ0vO86Q .icon-shape .label rect,#mermaid-svg-Emko39JYPZ0vO86Q .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-Emko39JYPZ0vO86Q .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-Emko39JYPZ0vO86Q .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-Emko39JYPZ0vO86Q :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 判别器 Discriminator
生成器 Generator
随机噪声 z

~N(0,1)
全连接层 100→256

BatchNorm + LeakyReLU
全连接层 256→512

BatchNorm + LeakyReLU
全连接层 512→1024

BatchNorm + LeakyReLU
全连接层 1024→4096

Tanh激活
Reshape

64×64×1
生成图像 G(z)
真实图像 x
Flatten

4096→1
Flatten

4096→1
全连接层 4096→512

LeakyReLU + Dropout
全连接层 4096→512

LeakyReLU + Dropout
全连接层 512→256

LeakyReLU + Dropout
全连接层 512→256

LeakyReLU + Dropout
全连接层 256→1

Sigmoid
全连接层 256→1

Sigmoid
D(x) ∈ 0,1
D(G(z)) ∈ 0,1
判别器损失

L_D = -log(D(x)) + log(1-D(G(z)))
生成器损失

L_G = -log(D(G(z)))

我们将使用 PyTorch 框架,基于 sklearn 内置的 digits 数据集(8×8 手写数字),构建一个完整的 GAN 模型,生成 64×64 的手写数字图像。

4.1 环境准备与数据加载

首先安装必要的库:

bash 复制代码
pip install torch torchvision matplotlib scikit-learn numpy

加载并预处理数据:

python 复制代码
import numpy as np
import torch
from sklearn.datasets import load_digits

def load_digits_dataset():
    """
    加载 sklearn 内置的 digits 手写数字数据集。
    返回:归一化并上采样到 64×64 的图像张量。
    """
    digits = load_digits()
    data = digits.images  # (1797, 8, 8)
    # 归一化到 [-1, 1](适配 Tanh 输出)
    data_norm = (data.astype(np.float32) - 8.0) / 8.0
    # 上采样到 64×64
    n = len(data_norm)
    upsampled = np.zeros((n, 1, 64, 64), dtype=np.float32)
    for i in range(n):
        upsampled[i, 0] = data_norm[i, 0].repeat(8, axis=0).repeat(8, axis=1)
    return torch.FloatTensor(upsampled)

代码拆解

  • load_digits():加载 sklearn 自带的 1797 张 8×8 手写数字灰度图。
  • 归一化:将像素值从 0,16 映射到 -1,1,因为生成器的输出层使用 Tanh 激活函数,其值域为 -1,1
  • 上采样:通过 repeat 将 8×8 图像放大 8 倍到 64×64,便于可视化。

4.2 构建生成器(Generator)

生成器接收一个 100 维的随机噪声,通过全连接层逐步"画"出 64×64 的图像。

python 复制代码
import torch.nn as nn

class Generator(nn.Module):
    def __init__(self, latent_dim=100, img_size=64):
        super(Generator, self).__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, img_size * img_size),
            nn.Tanh()  # 输出范围 [-1, 1]
        )
        self.img_size = img_size

    def forward(self, z):
        img = self.model(z)
        img = img.view(img.size(0), 1, self.img_size, self.img_size)
        return img

代码拆解

  • nn.Linear:全连接层,将噪声向量逐步映射到更高维度。
  • nn.BatchNorm1d:批量归一化,加速训练并稳定 GAN 的训练过程。
  • nn.LeakyReLU:带泄露的 ReLU,防止梯度消失(负值也有小梯度)。
  • nn.Tanh:将输出压缩到 -1,1,与归一化后的真实数据范围一致。
  • view:将一维向量重塑为图像形状 (batch, channel, height, width)。

4.3 构建判别器(Discriminator)

判别器接收一张 64×64 的图像,输出一个 0~1 的概率值,表示其为真实图像的可信度。

python 复制代码
class Discriminator(nn.Module):
    def __init__(self, img_size=64):
        super(Discriminator, self).__init__()
        self.model = nn.Sequential(
            nn.Linear(img_size * img_size, 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),
            nn.Sigmoid()  # 输出概率值 [0, 1]
        )

    def forward(self, img):
        img_flat = img.view(img.size(0), -1)  # 展平
        validity = self.model(img_flat)
        return validity

代码拆解

  • nn.Dropout:随机丢弃部分神经元,防止判别器过早过拟合(变得太强),保持与生成器的对抗平衡。
  • nn.Sigmoid:将输出映射到 0,1,表示"真"的概率。
  • view(img.size(0), -1):将图像展平为一维向量,输入全连接网络。

4.4 训练循环:对抗博弈的实现

训练过程严格按照 2.2 节描述的步骤进行:

python 复制代码
def train_gan(generator, discriminator, dataloader, epochs=200):
    adversarial_loss = nn.BCELoss()  # 二分类交叉熵损失
    optimizer_G = torch.optim.Adam(generator.parameters(), lr=0.0002, betas=(0.5, 0.999))
    optimizer_D = torch.optim.Adam(discriminator.parameters(), lr=0.0002, betas=(0.5, 0.999))

    for epoch in range(epochs):
        for i, (real_imgs, _) in enumerate(dataloader):
            batch_size = real_imgs.size(0)
            real_labels = torch.ones((batch_size, 1), device=real_imgs.device)
            fake_labels = torch.zeros((batch_size, 1), device=real_imgs.device)

            # ---------------------
            # 训练判别器
            # ---------------------
            optimizer_D.zero_grad()
            # 真实图像的损失
            real_pred = discriminator(real_imgs)
            d_real_loss = adversarial_loss(real_pred, real_labels)
            # 假图像的损失
            z = torch.randn((batch_size, 100), device=real_imgs.device)
            fake_imgs = generator(z).detach()  # detach 阻止梯度传到生成器
            fake_pred = discriminator(fake_imgs)
            d_fake_loss = adversarial_loss(fake_pred, fake_labels)
            # 判别器总损失
            d_loss = (d_real_loss + d_fake_loss) / 2
            d_loss.backward()
            optimizer_D.step()

            # ---------------------
            # 训练生成器
            # ---------------------
            optimizer_G.zero_grad()
            z = torch.randn((batch_size, 100), device=real_imgs.device)
            gen_imgs = generator(z)
            gen_pred = discriminator(gen_imgs)
            # 生成器希望假图被判别为真
            g_loss = adversarial_loss(gen_pred, real_labels)
            g_loss.backward()
            optimizer_G.step()

关键点

  • detach():在训练判别器时,假图像从生成器产生后要 detach(),防止判别器的梯度影响生成器。
  • 标签:真实图像标签为 1,假图像标签为 0。训练生成器时,虽然输入是假图像,但我们希望判别器输出 1(即骗过判别器),所以使用 real_labels
  • 损失函数:二分类交叉熵(BCELoss),衡量预测概率与真实标签的差距。

4.5 可视化训练结果

训练完成后,我们可以查看生成效果和损失曲线:

python 复制代码
import matplotlib.pyplot as plt

def visualize_results(generator, dataloader, g_losses, d_losses):
    generator.eval()
    with torch.no_grad():
        z = torch.randn((16, 100), device=next(generator.parameters()).device)
        gen_imgs = generator(z).cpu().numpy()

    fig, axes = plt.subplots(1, 2, figsize=(12, 5))
    # 生成样本
    for i in range(16):
        ax = axes[0].imshow(gen_imgs[i, 0], cmap='gray', vmin=-1, vmax=1)
        axes[0].axis('off')
    axes[0].set_title('Generated Handwritten Digits')
    # 损失曲线
    axes[1].plot(d_losses, label='Discriminator Loss', color='red')
    axes[1].plot(g_losses, label='Generator Loss', color='green')
    axes[1].set_xlabel('Epoch')
    axes[1].set_ylabel('Loss')
    axes[1].legend()
    axes[1].grid(True, alpha=0.3)
    plt.show()

5. 常见问题与调优技巧

5.1 GAN 训练不稳定的原因

GAN 训练过程中常见的不稳定问题及其相互关系:
#mermaid-svg-MZjNQViTEANPcTM6{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-MZjNQViTEANPcTM6 .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-MZjNQViTEANPcTM6 .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-MZjNQViTEANPcTM6 .error-icon{fill:#552222;}#mermaid-svg-MZjNQViTEANPcTM6 .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-MZjNQViTEANPcTM6 .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-MZjNQViTEANPcTM6 .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-MZjNQViTEANPcTM6 .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-MZjNQViTEANPcTM6 .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-MZjNQViTEANPcTM6 .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-MZjNQViTEANPcTM6 .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-MZjNQViTEANPcTM6 .marker{fill:#333333;stroke:#333333;}#mermaid-svg-MZjNQViTEANPcTM6 .marker.cross{stroke:#333333;}#mermaid-svg-MZjNQViTEANPcTM6 svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-MZjNQViTEANPcTM6 p{margin:0;}#mermaid-svg-MZjNQViTEANPcTM6 .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-MZjNQViTEANPcTM6 .cluster-label text{fill:#333;}#mermaid-svg-MZjNQViTEANPcTM6 .cluster-label span{color:#333;}#mermaid-svg-MZjNQViTEANPcTM6 .cluster-label span p{background-color:transparent;}#mermaid-svg-MZjNQViTEANPcTM6 .label text,#mermaid-svg-MZjNQViTEANPcTM6 span{fill:#333;color:#333;}#mermaid-svg-MZjNQViTEANPcTM6 .node rect,#mermaid-svg-MZjNQViTEANPcTM6 .node circle,#mermaid-svg-MZjNQViTEANPcTM6 .node ellipse,#mermaid-svg-MZjNQViTEANPcTM6 .node polygon,#mermaid-svg-MZjNQViTEANPcTM6 .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-MZjNQViTEANPcTM6 .rough-node .label text,#mermaid-svg-MZjNQViTEANPcTM6 .node .label text,#mermaid-svg-MZjNQViTEANPcTM6 .image-shape .label,#mermaid-svg-MZjNQViTEANPcTM6 .icon-shape .label{text-anchor:middle;}#mermaid-svg-MZjNQViTEANPcTM6 .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-MZjNQViTEANPcTM6 .rough-node .label,#mermaid-svg-MZjNQViTEANPcTM6 .node .label,#mermaid-svg-MZjNQViTEANPcTM6 .image-shape .label,#mermaid-svg-MZjNQViTEANPcTM6 .icon-shape .label{text-align:center;}#mermaid-svg-MZjNQViTEANPcTM6 .node.clickable{cursor:pointer;}#mermaid-svg-MZjNQViTEANPcTM6 .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-MZjNQViTEANPcTM6 .arrowheadPath{fill:#333333;}#mermaid-svg-MZjNQViTEANPcTM6 .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-MZjNQViTEANPcTM6 .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-MZjNQViTEANPcTM6 .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-MZjNQViTEANPcTM6 .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-MZjNQViTEANPcTM6 .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-MZjNQViTEANPcTM6 .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-MZjNQViTEANPcTM6 .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-MZjNQViTEANPcTM6 .cluster text{fill:#333;}#mermaid-svg-MZjNQViTEANPcTM6 .cluster span{color:#333;}#mermaid-svg-MZjNQViTEANPcTM6 div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-MZjNQViTEANPcTM6 .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-MZjNQViTEANPcTM6 rect.text{fill:none;stroke-width:0;}#mermaid-svg-MZjNQViTEANPcTM6 .icon-shape,#mermaid-svg-MZjNQViTEANPcTM6 .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-MZjNQViTEANPcTM6 .icon-shape p,#mermaid-svg-MZjNQViTEANPcTM6 .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-MZjNQViTEANPcTM6 .icon-shape .label rect,#mermaid-svg-MZjNQViTEANPcTM6 .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-MZjNQViTEANPcTM6 .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-MZjNQViTEANPcTM6 .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-MZjNQViTEANPcTM6 :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} GAN 训练不稳定
模式崩溃 Mode Collapse
梯度消失/爆炸
训练不平衡
生成器只产生少数几种样本
缺乏多样性
判别器过早收敛
判别器太强
生成器梯度接近零
生成器无法更新
生成器太强
判别器无法学习
判别器收敛过快
生成器无法进步
生成器收敛过快
判别器无法区分真假
纳什均衡难以达到
训练失败或效果差

5.1 GAN 训练不稳定的原因

  • 模式崩溃(Mode Collapse):生成器只学会生成少数几种样本,缺乏多样性。
  • 梯度消失:判别器太强,导致生成器梯度几乎为零,无法更新。
  • 训练不平衡:判别器或生成器一方过早收敛。

5.2 实用调优技巧

  1. 使用 BatchNorm:在生成器中加入批量归一化,稳定训练。
  2. 使用 LeakyReLU:避免梯度消失,尤其在判别器中。
  3. 标签平滑(Label Smoothing):将真实标签从 1 改为 0.9,假标签从 0 改为 0.1,防止判别器过于自信。
  4. 调整学习率:GAN 通常需要较小的学习率(如 0.0002)。
  5. 交替训练次数:可以尝试让判别器多训练几次(如 D_STEPS=5),再训练一次生成器。

6. 总结

GAN 通过生成器与判别器的对抗博弈,实现了从噪声中生成逼真数据的能力。本文从生活化的比喻入手,详细剖析了 GAN 的原理、训练过程,并提供了一个完整的 PyTorch 实现,生成手写数字图像。

核心要点回顾

  1. 生成器:从噪声生成数据,目标是骗过判别器。
  2. 判别器:区分真实数据与生成数据,目标是准确鉴别。
  3. 对抗训练:两者在博弈中共同进化,最终达到纳什均衡。
  4. 数学本质:最小化生成数据分布与真实数据分布之间的 JS 散度。

GAN 是生成模型的里程碑,后续还衍生出 DCGAN、WGAN、StyleGAN 等更强大的变体。掌握基础 GAN 是深入理解现代生成式 AI 的重要第一步。

附录:完整可运行代码

以下代码整合了上述所有模块,可直接复制到 PyTorch 环境中运行(需要安装 torch, matplotlib, scikit-learn, numpy)。

python 复制代码
# -*- coding: utf-8 -*-
"""
==============================================================================
GAN 初学者教程 ------ 使用生成对抗网络生成手写数字
==============================================================================
模型:  GAN(生成对抗网络),包含生成器 Generator 与判别器 Discriminator
框架:  PyTorch
数据集:sklearn 内置的 digits 数据集(8×8 灰度手写数字,无需联网下载)
任务:  训练一个 GAN,使其能够生成以假乱真的手写数字图像

GAN 核心思想(大白话):
    - 生成器就像"假钞制造者",试图做出逼真的假币
    - 判别器就像"警察",努力区分真币和假币
    - 两者相互博弈,最终生成器能造出连警察都分不出的"假钞"

运行方式:
    pip install -r requirements.txt
    python gan_demo.py
==============================================================================
"""

import numpy as np
import matplotlib.pyplot as plt
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader, TensorDataset
from sklearn.datasets import load_digits

# 设置中文字体(解决 matplotlib 中文显示问题)
plt.rcParams["font.sans-serif"] = ["SimHei", "Microsoft YaHei", "DejaVu Sans"]
plt.rcParams["axes.unicode_minus"] = False

# ============================================================================
# 0. 全局配置(初学者可以调整这些超参数来观察效果)
# ============================================================================
LATENT_DIM = 100        # 噪声向量维度(生成器的输入,越大网络表达能力越强)
IMG_SIZE = 64           # 生成图像的目标尺寸(8×8 → 上采样到 64×64)
HIDDEN_G = 128          # 生成器隐藏层神经元数
HIDDEN_D = 128          # 判别器隐藏层神经元数
BATCH_SIZE = 64         # 每批训练样本数
LEARNING_RATE = 0.0002  # 学习率(GAN 通常用较小的学习率)
BETAS = (0.5, 0.999)    # Adam 优化器参数
EPOCHS = 200            # 训练轮数
D_STEPS = 1             # 每轮判别器训练次数
G_STEPS = 1             # 每轮生成器训练次数
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"

print(f"[设备] 使用设备: {DEVICE}")
print(f"[版本] PyTorch 版本: {torch.__version__}")


# ============================================================================
# 1. 加载数据集 ------ sklearn 内置 digits,零网络开销
# ============================================================================
def load_digits_dataset():
    """
    加载 sklearn 内置的 digits 手写数字数据集。
    数据说明:
        - 共 1797 张 8×8 的灰度手写数字图像(0~9)
        - 每张图像像素值范围 [0, 16]
        - sklearn 自带,完全本地加载,无需联网下载
    """
    print("\n[数据] 正在加载 sklearn digits 数据集(本地,无需下载)...")
    digits = load_digits()
    data = digits.images        # (1797, 8, 8)
    labels = digits.target      # (1797,)

    print(f"   [OK] 加载成功!")
    print(f"   [统计] 图像数量: {len(data)} 张")
    print(f"   [尺寸] 原始尺寸: 8x8 像素")
    print(f"   [类别] 类别: 0 ~ 9(共 10 个数字)")
    print(f"   [范围] 像素范围: [{data.min()}, {data.max()}]")

    # 将像素值归一化到 [-1, 1](GAN 中通常使用 tanh 输出,对应 [-1, 1])
    data_norm = (data.astype(np.float32) - 8.0) / 8.0  # [0,16] → [-1,1]
    # 添加通道维度: (N, 8, 8) → (N, 1, 8, 8)
    data_norm = data_norm.reshape(-1, 1, 8, 8)

    # 上采样到 64×64(使生成图像更清晰,也方便可视化)
    # 使用简单的 repeat + 平均池化的方式
    n = len(data_norm)
    upsampled = np.zeros((n, 1, IMG_SIZE, IMG_SIZE), dtype=np.float32)
    for i in range(n):
        upsampled[i, 0] = data_norm[i, 0].repeat(8, axis=0).repeat(8, axis=1)

    print(f"   [尺寸] 上采样后尺寸: {IMG_SIZE}x{IMG_SIZE} 像素")
    return torch.FloatTensor(upsampled), labels


# ============================================================================
# 2. 定义生成器 Generator
# ============================================================================
class Generator(nn.Module):
    """
    生成器:从随机噪声中"画"出逼真的手写数字。

    输入: [batch, LATENT_DIM] 的随机噪声
    输出: [batch, 1, 64, 64] 的灰度图像(像素值范围 [-1, 1])
    """
    def __init__(self):
        super(Generator, self).__init__()
        self.model = nn.Sequential(
            # 第 1 层: 噪声 → 256 个神经元
            nn.Linear(LATENT_DIM, 256),
            nn.BatchNorm1d(256),
            nn.LeakyReLU(0.2, inplace=True),

            # 第 2 层: 256 → 512
            nn.Linear(256, 512),
            nn.BatchNorm1d(512),
            nn.LeakyReLU(0.2, inplace=True),

            # 第 3 层: 512 → 1024
            nn.Linear(512, 1024),
            nn.BatchNorm1d(1024),
            nn.LeakyReLU(0.2, inplace=True),

            # 输出层: 1024 → 64*64 = 4096
            nn.Linear(1024, IMG_SIZE * IMG_SIZE),
            nn.Tanh()   # Tanh 使输出落在 [-1, 1],与真实数据范围一致
        )

    def forward(self, z):
        """
        前向传播。
        z: 随机噪声,shape = [batch, LATENT_DIM]
        """
        img = self.model(z)
        img = img.view(img.size(0), 1, IMG_SIZE, IMG_SIZE)  # 重塑为图像形状
        return img


# ============================================================================
# 3. 定义判别器 Discriminator
# ============================================================================
class Discriminator(nn.Module):
    """
    判别器:判断一张图像是"真实手写数字"还是"生成器伪造的"。

    输入: [batch, 1, 64, 64] 的图像
    输出: [batch, 1] 的概率值(0=假,1=真)
    """
    def __init__(self):
        super(Discriminator, self).__init__()
        self.model = nn.Sequential(
            # 输入层: 64*64 → 512
            nn.Linear(IMG_SIZE * IMG_SIZE, 512),
            nn.LeakyReLU(0.2, inplace=True),
            nn.Dropout(0.3),      # 防止判别器过强,保持对抗平衡

            # 隐藏层: 512 → 256
            nn.Linear(512, 256),
            nn.LeakyReLU(0.2, inplace=True),
            nn.Dropout(0.3),

            # 输出层: 256 → 1
            nn.Linear(256, 1),
            nn.Sigmoid()          # Sigmoid 输出概率值 [0, 1]
        )

    def forward(self, img):
        """
        前向传播。
        img: 图像,shape = [batch, 1, 64, 64]
        """
        img_flat = img.view(img.size(0), -1)  # 展平为向量
        validity = self.model(img_flat)
        return validity


# ============================================================================
# 4. 权重初始化(帮助 GAN 稳定训练)
# ============================================================================
def weights_init_normal(m):
    """使用正态分布初始化权重,均值为 0,标准差为 0.02"""
    classname = m.__class__.__name__
    if classname.find("Linear") != -1:
        nn.init.normal_(m.weight.data, 0.0, 0.02)
        if m.bias is not None:
            nn.init.constant_(m.bias.data, 0)


# ============================================================================
# 5. 训练函数
# ============================================================================
def train_gan(generator, discriminator, dataloader, epochs):
    """
    训练 GAN 的核心逻辑。

    训练流程:
        1. 训练判别器:给真图→判为真,给假图→判为假
        2. 训练生成器:生成假图→骗过判别器→判为真
        3. 不断重复,直到生成器可以以假乱真
    """
    # 损失函数:二分类交叉熵(BCE)
    adversarial_loss = nn.BCELoss()

    # 优化器:两个网络各用各的
    optimizer_G = optim.Adam(generator.parameters(), lr=LEARNING_RATE, betas=BETAS)
    optimizer_D = optim.Adam(discriminator.parameters(), lr=LEARNING_RATE, betas=BETAS)

    # 记录损失历史(用于画图)
    g_losses = []
    d_losses = []

    print(f"\n[训练] 开始训练 GAN(共 {epochs} 轮)...")
    print("=" * 60)

    for epoch in range(epochs):
        epoch_g_loss = 0.0
        epoch_d_loss = 0.0
        n_batches = 0

        for i, (imgs, _) in enumerate(dataloader):
            batch_size_actual = imgs.size(0)
            n_batches += 1

            # ---------------------
            # 真实图像和标签
            # ---------------------
            real_imgs = imgs.to(DEVICE)
            real_labels = torch.ones((batch_size_actual, 1), device=DEVICE)  # 真=1
            fake_labels = torch.zeros((batch_size_actual, 1), device=DEVICE) # 假=0

            # ---------------------
            # 训练判别器 Discriminator
            # ---------------------
            for _ in range(D_STEPS):
                optimizer_D.zero_grad()

                # 1) 用真实图像训练判别器 ------ 希望判别器输出 1(真)
                real_pred = discriminator(real_imgs)
                d_real_loss = adversarial_loss(real_pred, real_labels)

                # 2) 用假图像训练判别器 ------ 希望判别器输出 0(假)
                z = torch.randn((batch_size_actual, LATENT_DIM), device=DEVICE)
                fake_imgs = generator(z).detach()  # detach 防止梯度传到生成器
                fake_pred = discriminator(fake_imgs)
                d_fake_loss = adversarial_loss(fake_pred, fake_labels)

                # 判别器总损失 = 真损失 + 假损失
                d_loss = (d_real_loss + d_fake_loss) / 2
                d_loss.backward()
                optimizer_D.step()

            # ---------------------
            # 训练生成器 Generator
            # ---------------------
            for _ in range(G_STEPS):
                optimizer_G.zero_grad()

                # 生成一批假图像 → 希望判别器判为真(输出 1)
                z = torch.randn((batch_size_actual, LATENT_DIM), device=DEVICE)
                gen_imgs = generator(z)
                gen_pred = discriminator(gen_imgs)
                g_loss = adversarial_loss(gen_pred, real_labels)  # 注意:标签用"真"

                g_loss.backward()
                optimizer_G.step()

            epoch_g_loss += g_loss.item()
            epoch_d_loss += d_loss.item()

        # 记录平均损失
        avg_g = epoch_g_loss / n_batches
        avg_d = epoch_d_loss / n_batches
        g_losses.append(avg_g)
        d_losses.append(avg_d)

        # 每 20 轮打印一次进度
        if (epoch + 1) % 20 == 0:
            print(f"[进度] Epoch {epoch+1:4d}/{epochs}  "
                  f"|  D Loss: {avg_d:.4f}  |  G Loss: {avg_g:.4f}")

    print("=" * 60)
    print("[OK] 训练完成!")
    return g_losses, d_losses


# ============================================================================
# 6. 可视化 ------ 生成样本 & 损失曲线
# ============================================================================
def visualize_results(generator, dataloader, g_losses, d_losses):
    """
    展示训练成果:
    - 左图:生成器产生的 16 张手写数字
    - 右图:生成器 & 判别器的损失变化曲线
    """
    generator.eval()

    # ---------- 生成 16 张样本 ----------
    with torch.no_grad():
        z = torch.randn((16, LATENT_DIM), device=DEVICE)
        gen_imgs = generator(z).cpu().numpy()  # (16, 1, 64, 64)

    fig, axes = plt.subplots(2, 2, figsize=(14, 12))
    fig.suptitle("GAN 训练结果 --- 手写数字生成", fontsize=16, fontweight="bold")

    # ---- 子图 1-2:生成的假图像 ----
    for i in range(16):
        row, col = i // 4, i % 4
        axes[0, 0].imshow(gen_imgs[i, 0], cmap="gray", vmin=-1, vmax=1)
        axes[0, 0].axis("off")
    axes[0, 0].set_title("生成器生成的样本(假图像)", fontsize=13)

    # ---- 子图 1-2:真实图像对比 ----
    real_imgs, _ = next(iter(dataloader))
    real_imgs = real_imgs[:16].cpu().numpy()
    for i in range(16):
        row, col = i // 4, i % 4
        axes[0, 1].imshow(real_imgs[i, 0], cmap="gray", vmin=-1, vmax=1)
        axes[0, 1].axis("off")
    axes[0, 1].set_title("真实手写数字(训练集样本)", fontsize=13)

    # ---- 子图 2-1:损失曲线 ----
    axes[1, 0].plot(d_losses, label="判别器损失 (D Loss)", color="#e74c3c", linewidth=1.5)
    axes[1, 0].plot(g_losses, label="生成器损失 (G Loss)", color="#2ecc71", linewidth=1.5)
    axes[1, 0].set_xlabel("Epoch")
    axes[1, 0].set_ylabel("Loss")
    axes[1, 0].set_title("损失变化曲线", fontsize=13)
    axes[1, 0].legend()
    axes[1, 0].grid(True, alpha=0.3)

    # ---- 子图 2-2:训练过程 GIF 式的展示(额外生成一组新样本) ----
    with torch.no_grad():
        z2 = torch.randn((16, LATENT_DIM), device=DEVICE)
        gen_imgs2 = generator(z2).cpu().numpy()
    for i in range(16):
        row, col = i // 4, i % 4
        axes[1, 1].imshow(gen_imgs2[i, 0], cmap="gray", vmin=-1, vmax=1)
        axes[1, 1].axis("off")
    axes[1, 1].set_title("再生成一组样本(验证稳定性)", fontsize=13)

    plt.tight_layout()
    plt.savefig("gan_result.png", dpi=150, bbox_inches="tight")
    plt.show()
    print("\n[图片] 结果图已保存为: gan_result.png")


# ============================================================================
# 7. 主程序入口
# ============================================================================
if __name__ == "__main__":
    print("\n" + "=" * 60)
    print("  GAN 初学者教程 -- 生成手写数字")
    print("=" * 60)

    # 7.1 加载数据
    tensor_data, labels = load_digits_dataset()
    dataset = TensorDataset(tensor_data, torch.LongTensor(labels))
    dataloader = DataLoader(dataset, batch_size=BATCH_SIZE, shuffle=True)
    print(f"   [批次] 批次数量: {len(dataloader)} (batch_size={BATCH_SIZE})")

    # 7.2 创建模型
    generator = Generator().to(DEVICE)
    discriminator = Discriminator().to(DEVICE)

    # 初始化权重
    generator.apply(weights_init_normal)
    discriminator.apply(weights_init_normal)

    print(f"\n[模型] 生成器参数量: {sum(p.numel() for p in generator.parameters()):,}")
    print(f"[模型] 判别器参数量: {sum(p.numel() for p in discriminator.parameters()):,}")

    # 7.3 训练
    g_losses, d_losses = train_gan(generator, discriminator, dataloader, EPOCHS)

    # 7.4 可视化结果
    visualize_results(generator, dataloader, g_losses, d_losses)

    print("\n" + "=" * 60)
    print("[完成] 程序运行完毕!")
    print("[提示] 如果生成效果不理想,可以尝试:")
    print("   1. 增加 EPOCHS(如 300~500)")
    print("   2. 调整 LEARNING_RATE(如 0.0001)")
    print("   3. 修改 HIDDEN_G / HIDDEN_D(如 256)")
    print("   4. 增大 LATENT_DIM(如 200)")
    print("=" * 60)