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)"。
- 给判别器看一批生成器造的假图片,告诉它"这些是假的(标签为 0)"。
- 判别器通过比较预测结果与真实标签,更新自己的参数,提高鉴别能力。
-
固定判别器,训练生成器:
- 生成器造出一批假图片,但这次我们希望判别器把它们判断为"真(标签为 1)"。
- 生成器根据判别器的"误判"程度来更新自己的参数,让自己画得更逼真。
-
循环往复:
- 两者不断对抗、相互提升。理想情况下,最终生成器能生成与真实数据分布几乎一致的样本,而判别器则无法区分真假(输出概率接近 0.5)。
2.3 数学本质:最小化 JS 散度
从数学上看,GAN 的训练目标是一个极小极大博弈(Minimax Game):
minGmaxDV(D,G)=Ex∼pdata(x)logD(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 实用调优技巧
- 使用 BatchNorm:在生成器中加入批量归一化,稳定训练。
- 使用 LeakyReLU:避免梯度消失,尤其在判别器中。
- 标签平滑(Label Smoothing):将真实标签从 1 改为 0.9,假标签从 0 改为 0.1,防止判别器过于自信。
- 调整学习率:GAN 通常需要较小的学习率(如 0.0002)。
- 交替训练次数:可以尝试让判别器多训练几次(如 D_STEPS=5),再训练一次生成器。
6. 总结
GAN 通过生成器与判别器的对抗博弈,实现了从噪声中生成逼真数据的能力。本文从生活化的比喻入手,详细剖析了 GAN 的原理、训练过程,并提供了一个完整的 PyTorch 实现,生成手写数字图像。
核心要点回顾:
- 生成器:从噪声生成数据,目标是骗过判别器。
- 判别器:区分真实数据与生成数据,目标是准确鉴别。
- 对抗训练:两者在博弈中共同进化,最终达到纳什均衡。
- 数学本质:最小化生成数据分布与真实数据分布之间的 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)