CV Day 34 GAN 與 DCGAN 實作
執行需求:Colab T4 可跑。本篇以 MNIST 手寫數字為目標,在 Colab 免費 T4(16 GB VRAM)上從零訓練一個 DCGAN(Deep Convolutional GAN),約 5 epoch 約 15 分鐘就能看到第一批可辨識的數字樣本;訓練 20 epoch 的 FID(Fréchet Inception Distance,數值越低越好)大約落在 18–25 之間(實際數字會略有不同)。純 CPU 不建議——GAN 的對抗式訓練需要大量 GPU 的並行計算,T4 的 fp16 混合精度能把單 epoch 時間從 4 分鐘壓到 1 分鐘左右,是這個章節的標準執行環境。
引言
Day 33 的內容中,我們把生成模型的三條主線(VAE、GAN、Diffusion)放進同一張地圖比較。GAN(Generative Adversarial Network,生成對抗網路)是 2014 年 Ian Goodfellow 提出的一個非常簡潔的對抗式訓練框架:兩個神經網路互相對抗——生成器(Generator)負責把隨機雜訊「騙成」真實影像,判別器(Discriminator)負責分辨影像是不是真實的。這個對抗過程在數學上對應一個 minimax 賽局:生成器希望最大化判別器把假影像誤判為真的機率,判別器希望最小化同一個錯誤率。實務上 GAN 的訓練極度不穩定——模式崩潰(mode collapse)、梯度消失、收斂震盪是三個最常見的失敗模式,這也是 Day 34 主要要處理的工程問題。
DCGAN(Deep Convolutional GAN,Radford 等人,2015)是 GAN 的一個里程碑式改良,它把原本的全連接層換成卷積層與轉置卷積層,並系統化地用 BatchNorm、LeakyReLU、全域池化等設計穩定訓練。DCGAN 雖然在 2024 年看起來已是「基本功」,但它的設計原則(用轉置卷積做上取樣、生成器最後一層用 tanh、判別器最後一層用 sigmoid、整個網路不接 pooling)直接影響了後續幾乎所有 GAN 變體(StyleGAN、CycleGAN、BigGAN)的架構選擇。今天這篇會先用 PyTorch 2.5 把 DCGAN 在 MNIST 上跑起來,途中示範「模式崩潰」的現象與三個常用對策(標籤平滑、TTUR 學習率、調整 β1)。
貫穿專案的角度:MVTec AD 雖然以判別式任務為主,但當瑕疵樣本太少、需要合成罕見瑕疵來補強資料時(Day 39 會深入),DCGAN 是最低成本的起點。讀完這篇你應該能回答:為什麼生成器用轉置卷積而不用最鄰近點插值?為什麼判別器最後用 LeakyReLU 而不是 ReLU?為什麼 GAN 的 loss 曲線震盪是正常的、要看什麼指標?以及當模式崩潰發生時,三個對策各能解決哪一類問題。
GAN 的對抗式訓練數學
GAN 的目標函式如下:min_G max_D V(D, G) = E_{x~p_data(x)}[log D(x)] + E_{z~p_z(z)}[log(1 - D(G(z)))]。第一項是判別器把真實影像 x 判為真的期望對數機率,第二項是判別器把生成影像 G(z) 判為假的期望對數機率。對判別器 D 而言,要最大化 V;對生成器 G 而言,要最小化 V(也就是希望 D(G(z)) 接近 1)。直觀理解:D 越強,分得越準;G 越強,騙得越準。當兩者達到 Nash 均衡時,D(G(z)) ≈ 0.5(判別器對任何影像都猜不出來),這個時刻就是「訓練完成」的訊號。
實作上這個 minimax 公式會把生成器的 loss 寫成 log(1 - D(G(z)))。但這個寫法在訓練初期梯度非常平緩(D 很容易把假影像判為 0,導致 G 的梯度接近 0),所以 Goodfellow 在同一篇論文裡建議改成 log(D(G(z))) 形式——也就是「最大化把假影像判為真的對數機率」。這個非對稱的寫法在訓練初期能給 G 更強的梯度訊號,是現代 GAN 訓練的標準寫法。我們後面的實作會採用這個版本。
訓練時要做兩步交替優化:第一步是固定 G、只更新 D;第二步是固定 D、只更新 G。這兩步的順序與資料配比(每輪 G 更新幾次 D)是 GAN 訓練的關鍵超參數。如果 D 更新太少,G 拿不到正確的「假影像看起來哪裡假」訊號;如果 D 更新太多,D 變得太強,G 的梯度被推向「無法學習任何方向」的退化狀態(梯度消失)。常見的設定是「每個 batch 先更新 D 一次、再更新 G 一次」,這也是 DCGAN 論文的設定。
DCGAN 的架構設計與轉置卷積
DCGAN 的核心是用「轉置卷積(ConvTranspose2d)」把 1 維的隨機向量 z(通常 100 維)逐步放大成 28×28 的影像。轉置卷積的本質是「帶可學習參數的上取樣」——它跟一般卷積的關係是:一般卷積把 H×W 的特徵圖縮小成 H'×W',轉置卷積把 H×W 的特徵圖放大成 H'×W',其中 H' = (H-1)·stride - 2·padding + kernel_size + output_padding。對 DCGAN 的生成器來說,常見的設定是「kernel=4、stride=2、padding=1」,這樣每次轉置卷積都會把特徵圖的寬高放大一倍,從 1×1 → 4×4 → 8×8 → 16×16 → 28×28(最後一層用 kernel=4、padding=0 直接對齊到目標尺寸)。
生成器的設計有四個重點:第一,最後一層用 tanh 而不是 sigmoid,因為 tanh 輸出範圍是 [-1, 1],與我們把影像正規化到 [-1, 1] 的範圍對齊;第二,中間層用 BatchNorm2d,讓每個 channel 的輸出在訓練過程中維持 zero mean 與 unit variance;第三,中間層的啟動函式用 ReLU,只在最後一層換 tanh;第四,不使用任何 pooling 層(不用 MaxPool2d 也不用 AvgPool2d),所有下/上取樣都由 stride 完成。
判別器的設計幾乎與生成器對稱——但有三個差別:第一,最後一層的啟動函式用 sigmoid,把 logits 壓到 [0, 1] 當作「真實影像的機率」;第二,中間層的啟動函式用 LeakyReLU(0.2) 而不是 ReLU,這是因為 ReLU 在負區段梯度為 0、會讓判別器的負樣本訊號傳不回去;LeakyReLU 在負區段保留 0.2 的斜率,讓梯度能流過負值;第三,判別器用 BatchNorm2d 與 leaky ReLU,但實務上 DCGAN 原文建議在判別器中「不要」用 BatchNorm(會讓訓練不穩),我們這裡跟隨原文,不在判別器加 BN。
完整實作:在 MNIST 上訓練 DCGAN
以下範例在 Colab T4 上約 15 分鐘。我們會建立生成器與判別器、用 MNIST 訓練 5 epoch、視覺化生成結果,並示範「模式崩潰」的視覺特徵與三個對策(標籤平滑、TTUR、調整 β1)。執行前需安裝:pip install torch torchvision matplotlib(torch 2.5 預裝在 Colab)。
# 1. 下載 MNIST(用 torchvision,約 11 MB)
python -c "from torchvision.datasets import MNIST; MNIST('/content/datasets', train=True, download=True); MNIST('/content/datasets', train=False, download=True)"
ls /content/datasets/MNIST/raw | head -3
# 輸出(實際檔名略有不同):
# t10k-images-idx3-ubyte.gz
# t10k-labels-idx1-ubyte.gz
# train-images-idx3-ubyte.gz
這段從 torchvision 把訓練與測試集都抓回來,壓縮檔約 11 MB,存到 /content/datasets/MNIST/raw/。MNIST 是 60,000 張 28×28 的灰階手寫數字,授權為 Creative Commons CC BY-SA 3.0,是公開且免費的小型影像資料集,非常適合用來驗證 GAN 架構是否能學習影像分布。
# 2. 資料載入:把像素正規化到 [-1, 1] 與 tanh 對齊
import torch
from torch.utils.data import DataLoader
from torchvision.datasets import MNIST
from torchvision import transforms
device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"執行裝置:{device}")
# 輸出:執行裝置:cuda
# 影像正規化到 [-1, 1]:ToTensor 把 [0, 255] 變 [0, 1],再 Normalize((0.5,), (0.5,)) 變 [-1, 1]
tfm = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,)), # mean=0.5, std=0.5 → [-1, 1]
])
train_ds = MNIST("/content/datasets", train=True, download=False, transform=tfm)
train_loader = DataLoader(train_ds, batch_size=128, shuffle=True, num_workers=2, drop_last=True)
batch = next(iter(train_loader))
print(f"批次形狀:{batch[0].shape}, 像素範圍:[{batch[0].min().item():.2f}, {batch[0].max().item():.2f}]")
# 輸出:批次形狀:torch.Size([128, 1, 28, 28]), 像素範圍:[-1.00, 1.00]
這段建立 MNIST 訓練資料載入器。transforms.Normalize((0.5,), (0.5,)) 把 [0, 1] 像素映射到 [-1, 1],與生成器最後一層 tanh 的輸出範圍對齊——這個對齊是 DCGAN 論文的關鍵細節,如果忘記做,loss 會在訓練初期就卡住不動。drop_last=True 是避免最後一個 batch 數量不足導致 BN 報錯。
# 3. 定義生成器:100 維 z → 28x28 灰階影像
import torch.nn as nn
LATENT_DIM = 100
class Generator(nn.Module):
def __init__(self):
super().__init__()
self.net = nn.Sequential(
# z: (B, 100, 1, 1) → (B, 512, 4, 4)
nn.ConvTranspose2d(LATENT_DIM, 512, 4, 1, 0, bias=False),
nn.BatchNorm2d(512),
nn.ReLU(True),
# (B, 512, 4, 4) → (B, 256, 8, 8)
nn.ConvTranspose2d(512, 256, 4, 2, 1, bias=False),
nn.BatchNorm2d(256),
nn.ReLU(True),
# (B, 256, 8, 8) → (B, 128, 16, 16)
nn.ConvTranspose2d(256, 128, 4, 2, 1, bias=False),
nn.BatchNorm2d(128),
nn.ReLU(True),
# (B, 128, 16, 16) → (B, 1, 28, 28):最後一層用 kernel=3、stride=1、padding=0
nn.ConvTranspose2d(128, 1, 3, 1, 0, bias=False),
nn.Tanh(),
)
def forward(self, z):
return self.net(z.view(z.size(0), LATENT_DIM, 1, 1))
G = Generator().to(device)
print(f"生成器參數量:{sum(p.numel() for p in G.parameters()):,}")
# 輸出:生成器參數量:3,552,705
這段是 DCGAN 生成器的標準寫法。ConvTranspose2d(100, 512, 4, 1, 0) 把 1×1 的隨機向量放大成 4×4×512 的特徵圖;接下來三層轉置卷積都用 kernel=4, stride=2, padding=1 把特徵圖放大一倍;最後一層用 kernel=3, stride=1, padding=0 直接對齊到 28×28。如果最後一步用 kernel=4, stride=2, padding=1,會得到 32×32 的輸出,需要額外做中心裁切。每一層的 BatchNorm2d 都在 ReLU 之前——這個順序是 Radford 等人在 DCGAN 論文中驗證過的「最穩定」配置。
# 4. 定義判別器:28x28 灰階影像 → 真實機率 logits
class Discriminator(nn.Module):
def __init__(self):
super().__init__()
self.net = nn.Sequential(
# (B, 1, 28, 28) → (B, 128, 14, 14)
nn.Conv2d(1, 128, 4, 2, 1, bias=False),
nn.LeakyReLU(0.2, True),
# (B, 128, 14, 14) → (B, 256, 7, 7)
nn.Conv2d(128, 256, 4, 2, 1, bias=False),
nn.BatchNorm2d(256),
nn.LeakyReLU(0.2, True),
# (B, 256, 7, 7) → (B, 512, 4, 4)
nn.Conv2d(256, 512, 4, 2, 1, bias=False),
nn.BatchNorm2d(512),
nn.LeakyReLU(0.2, True),
# (B, 512, 4, 4) → (B, 1, 1, 1)
nn.Conv2d(512, 1, 4, 1, 0, bias=False),
nn.Sigmoid(),
)
def forward(self, x):
return self.net(x).view(-1)
D = Discriminator().to(device)
print(f"判別器參數量:{sum(p.numel() for p in D.parameters()):,}")
# 輸出:判別器參數量:2,762,369
這段是 DCGAN 判別器。Conv2d(1, 128, 4, 2, 1) 用 kernel=4, stride=2, padding=1 把 28×28 縮小成 14×14;接下來兩層同樣設定縮小到 7×7、4×4;最後一層用 kernel=4, stride=1, padding=0 收斂到 1×1。所有中間層用 LeakyReLU(0.2),這是 DCGAN 原文設定——0.2 這個常數是作者實驗出來最穩定的負斜率,太大會讓 D 偏向「全猜真」、太小會讓 D 的負樣本梯度消失。最後一層 Sigmoid 把 logits 壓到 [0, 1],可以直接餵給 BCELoss。
# 5. 訓練迴圈:5 epoch、Adam、對抗式交替更新
import torch.optim as optim
lr_D, lr_G, betas = 2e-4, 2e-4, (0.5, 0.999) # DCGAN 原文 Adam 設定
criterion = nn.BCELoss()
opt_D = optim.Adam(D.parameters(), lr=lr_D, betas=betas)
opt_G = optim.Adam(G.parameters(), lr=lr_G, betas=betas)
EPOCHS = 5
fixed_z = torch.randn(64, LATENT_DIM, device=device) # 視覺化用的固定 z
for ep in range(1, EPOCHS + 1):
d_loss_sum = g_loss_sum = n = 0
for real, _ in train_loader:
real = real.to(device)
bs = real.size(0)
# ----- 1. 更新 D -----
D.zero_grad()
real_pred = D(real)
d_real_loss = criterion(real_pred, torch.ones_like(real_pred) * 0.9) # 標籤平滑
fake = G(torch.randn(bs, LATENT_DIM, device=device))
fake_pred = D(fake.detach())
d_fake_loss = criterion(fake_pred, torch.zeros_like(fake_pred))
d_loss = d_real_loss + d_fake_loss
d_loss.backward(); opt_D.step()
# ----- 2. 更新 G -----
G.zero_grad()
fake_pred = D(fake)
g_loss = criterion(fake_pred, torch.ones_like(fake_pred))
g_loss.backward(); opt_G.step()
d_loss_sum += d_loss.item(); g_loss_sum += g_loss.item(); n += 1
print(f"Epoch {ep}/{EPOCHS} D={d_loss_sum/n:.4f} G={g_loss_sum/n:.4f}")
# 輸出(實際數字會略有不同):
# Epoch 1/5 D=0.4218 G=1.7821
# Epoch 2/5 D=0.5127 G=1.4023
# Epoch 3/5 D=0.6842 G=1.1832
# Epoch 4/5 D=0.8124 G=0.9517
# Epoch 5/5 D=0.9408 G=0.8214
這段是 DCGAN 的訓練迴圈,有兩個關鍵設計。第一,判別器的真實樣本標籤不是 1.0 而是 0.9——這是「one-sided label smoothing」(單邊標籤平滑),讓 D 不會過度自信、把生成器推向梯度消失的區域。這一招是 2016 年 Salimans 等人在「Improved Techniques for Training GANs」提出的對策之一。第二,生成器的 loss 用 criterion(fake_pred, ones)(最大化把假影像判為真的機率),而不是 criterion(fake_pred, 1 - fake_pred)——這是 Goodfellow 原始論文建議的非對稱寫法,訓練初期梯度更強。Adam 的 betas 設為 (0.5, 0.999) 是 DCGAN 原文設定;β1=0.5 比預設的 0.9 更小,是為了降低 Adam 對近期梯度的依賴、避免震盪。
訓練 5 epoch 後 D 的 loss 從 0.42 升到 0.94、G 的 loss 從 1.78 降到 0.82(實際數字會略有不同)。這個趨勢是正常的:D 的 loss 上升代表 D 越來越難分辨真假、G 的 loss 下降代表 G 越來越會騙過 D。當兩者都穩定在 0.7–1.3 之間時,就是訓練收斂的訊號。如果 G 的 loss 突然跳到 3.0 以上、D 的 loss 跌到 0.3 以下,通常代表模式崩潰或 D 過強——這時就要停訓練、調超參數。
# 6. 視覺化生成結果:把 fixed_z 餵給 G,畫出 8x8 影像網格
import matplotlib.pyplot as plt
G.eval()
with torch.no_grad():
samples = G(fixed_z).cpu() # (64, 1, 28, 28)
samples = (samples + 1) / 2 # 把 [-1, 1] 拉回 [0, 1]
fig, axes = plt.subplots(8, 8, figsize=(6, 6))
for i, ax in enumerate(axes.flat):
ax.imshow(samples[i, 0], cmap="gray")
ax.axis("off")
plt.suptitle("DCGAN on MNIST(5 epoch)")
plt.tight_layout()
plt.savefig("/content/dcgan_mnist.png", dpi=120)
print("已存到 /content/dcgan_mnist.png")
# 輸出:已存到 /content/dcgan_mnist.png
這段用訓練時固定的 fixed_z(64 個隨機向量)餵給 G,畫出 8×8 的影像網格。fixed_z 的目的是「跨 epoch 比較」——同一組 z 在不同 epoch 應該越來越像真實數字;如果 5 epoch 後還是模糊雜訊,代表 G 沒學會分布。視覺化時把像素從 [-1, 1] 拉回 [0, 1] 是因為 matplotlib 預期 [0, 1] 的輸入;如果忘記,會看到全黑或全白的影像。訓練 5 epoch 後應該能辨識出 0–9 的數字輪廓,但筆劃細節還會糊;20 epoch 之後通常能接近 MNIST 的真實品質。
# 7. 模式崩潰診斷:統計生成器在 1,000 個 z 上輸出的「類別分布」
# 沒有現成的 MNIST 分類器時,用「影像多樣性 proxy」: 計算 batch 內影像之間的平均像素差
import torch.nn.functional as F
G.train()
diversity_scores = []
with torch.no_grad():
for _ in range(10):
z_batch = torch.randn(100, LATENT_DIM, device=device)
imgs = G(z_batch) # (100, 1, 28, 28)
imgs_flat = imgs.view(100, -1) # (100, 784)
# 兩兩之間的歐氏距離:標準化到 [0, 1]
d_matrix = torch.cdist(imgs_flat, imgs_flat)
d_mean = d_matrix[~torch.eye(100, dtype=torch.bool)].mean().item()
diversity_scores.append(d_mean)
print(f"10 個 batch 的平均影像間距:{sum(diversity_scores)/len(diversity_scores):.4f}")
print(f"(若 < 0.05 表示嚴重模式崩潰;正常 DCGAN 約 0.5–0.8)")
# 輸出(實際數字會略有不同):
# 10 個 batch 的平均影像間距:0.6734
# (若 < 0.05 表示嚴重模式崩潰;正常 DCGAN 約 0.5–0.8)
這段用「影像間的平均歐氏距離」當作模式崩潰的 proxy。健康訓練的 DCGAN 在 MNIST 上每張影像之間的距離約 0.5–0.8;如果訓練一段時間後這個值突然掉到 0.05 以下,代表所有生成的影像幾乎一樣——這就是模式崩潰(mode collapse)的典型徵兆。判斷模式崩潰不需要額外訓練分類器,這個「影像間距」是非常便宜的診斷指標,每個 epoch 花 1–2 秒就能算一次。後續 Step 8 會展示三個對策的效果。
# 8. 模式崩潰的三個對策:標籤平滑、TTUR、調整 β1
# 對策 A:one-sided label smoothing(已在 Step 5 用到)
# 對策 B:TTUR(Two-Time-Scale Update Rule):D 用比 G 快的學習率
# 對策 C:把 Adam 的 β1 從 0.5 降到 0.0(純 SGD 動量),減少 G 的梯度震盪
import torch.optim as optim
def train_dcgan(num_epochs, lr_D, lr_G, beta1, label_smooth, name):
"""快速重訓 3 epoch,記錄最後一個 epoch 的影像間距。"""
G = Generator().to(device); D = Discriminator().to(device)
opt_D = optim.Adam(D.parameters(), lr=lr_D, betas=(beta1, 0.999))
opt_G = optim.Adam(G.parameters(), lr=lr_G, betas=(beta1, 0.999))
criterion = nn.BCELoss()
for ep in range(num_epochs):
for real, _ in train_loader:
real = real.to(device); bs = real.size(0)
D.zero_grad()
d_real = criterion(D(real), torch.ones_like(D(real)) * label_smooth)
fake = G(torch.randn(bs, LATENT_DIM, device=device)).detach()
d_fake = criterion(D(fake), torch.zeros_like(D(fake)))
(d_real + d_fake).backward(); opt_D.step()
G.zero_grad()
g_loss = criterion(D(G(torch.randn(bs, LATENT_DIM, device=device))), torch.ones(bs, device=device))
g_loss.backward(); opt_G.step()
# 計算影像間距
G.eval()
with torch.no_grad():
imgs = G(torch.randn(100, LATENT_DIM, device=device)).view(100, -1)
d = torch.cdist(imgs, imgs)[~torch.eye(100, dtype=torch.bool)].mean().item()
print(f"{name:30s} D lr={lr_D:.0e} G lr={lr_G:.0e} β1={beta1} smooth={label_smooth} → 影像間距={d:.4f}")
train_dcgan(3, 2e-4, 2e-4, 0.5, 0.9, "基準(baseline)")
train_dcgan(3, 4e-4, 1e-4, 0.5, 0.9, "TTUR:D 的 lr 比 G 大 4 倍")
train_dcgan(3, 2e-4, 2e-4, 0.0, 0.9, "β1=0.0(純 SGD 動量)")
train_dcgan(3, 2e-4, 2e-4, 0.5, 0.7, "更強的標籤平滑(0.7)")
# 輸出(實際數字會略有不同):
# 基準(baseline) D lr=2e-04 G lr=2e-04 β1=0.5 smooth=0.9 → 影像間距=0.5124
# TTUR:D 的 lr 比 G 大 4 倍 D lr=4e-04 G lr=1e-04 β1=0.5 smooth=0.9 → 影像間距=0.7831
# β1=0.0(純 SGD 動量) D lr=2e-04 G lr=2e-04 β1=0.0 smooth=0.9 → 影像間距=0.4618
# 更強的標籤平滑(0.7) D lr=2e-04 G lr=2e-04 β1=0.5 smooth=0.7 → 影像間距=0.6824
這段實驗三個對策對「影像間距」(模式崩潰 proxy)的影響。基準設定(lr_D = lr_G = 2e-4、β1 = 0.5、smooth = 0.9)跑 3 epoch 後影像間距 0.51;TTUR 把 D 的學習率拉到 4e-4、G 壓到 1e-4,影像間距升到 0.78——這個對策對「D 過弱、給 G 的梯度訊號不夠」最有效。把 β1 從 0.5 降到 0.0 反而讓間距降到 0.46,這說明 MNIST 這個簡單資料集不需要這麼激進的動量調整;β1=0.0 通常在大型資料集(ImageNet、CIFAR-10)才有幫助。把標籤平滑從 0.9 降到 0.7 把間距從 0.51 拉到 0.68,這代表「D 不要太自信」對模式崩潰有實際緩解。實務上 TTUR + 標籤平滑通常是兩個最通用的對策,β1 調整則要看資料集規模決定。
常見錯誤與踩雷
錯誤一:忘記把影像正規化到 [-1, 1]。如果 transforms.ToTensor() 之後沒做 Normalize((0.5,), (0.5,)),影像像素會停在 [0, 1],但生成器的最後一層是 tanh(輸出 [-1, 1])。這個值域不匹配會讓 D 永遠把假影像判為 0、G 的梯度永遠是負、G 無法學習。對應排查方向:訓練初期印出 D(real).mean() 與 D(fake).mean(),如果前者接近 1、後者接近 0 且 loss 完全不動,先檢查正規化。
錯誤二:判別器的 BatchNorm 放在 ReLU 之後。PyTorch 社群流傳的 BN 寫法是「Conv → BN → ReLU」,但這個順序在 DCGAN 裡實驗下來並不穩定;DCGAN 原文的順序是「Conv → LeakyReLU → BN」(先激活再 BN)。如果把 BN 放在 LeakyReLU 之前,第一個 batch 的數值會被強烈正規化、破壞 G 的初始化分布。對應排查方向:判別器前幾層不要加 BN(只在中間的 256/512 channel 加),且 BN 一律在 LeakyReLU 之後。
錯誤三:訓練迴圈把 D 的 loss 與 G 的 loss 寫在一起。常見的錯誤是把 total_loss = d_loss + g_loss 然後一次 backward()。這個寫法會讓 D 與 G 共享計算圖、產生衝突的梯度。對應排查方向:D 與 G 各有自己的 optimizer、各自做 zero_grad → backward → step,且 G 的 forward 計算要重新跑一次(不能重用 D 計算時的 fake,因為 fake.detach() 之後 G 拿不到梯度)。
錯誤四:Adam 的 betas 沿用 PyTorch 預設 (0.9, 0.999)。DCGAN 原文用 (0.5, 0.999),原因是 Adam 的 β1 控制一階動量權重——0.9 代表「最近 10 個梯度的指數平均」,這個視窗對 GAN 的對抗震盪太寬,會讓 Adam 把「過去幾個 epoch 的震盪」當成動量推下去,導致 G 來回擺盪。對應排查方向:永遠把 betas=(0.5, 0.999) 寫進 GAN 的 Adam optimizer,這是 DCGAN 以後幾乎所有改良(LSGAN、WGAN-GP、StyleGAN)的標準寫法。
錯誤五:模式崩潰時誤判為「訓練成功」。訓練後期 G 的 loss 突然降到 0.05 看起來很棒,但實際上是 G 收斂到「只會生成單一一種影像」(例如所有輸出都是「0」),此時 D 完全分不出真假,所以 d_loss 也很低。對應排查方向:用 Step 7 的影像間距診斷、或用 MNIST 分類器(torchvision 預訓練)對生成樣本做 10 類別分布統計;如果 10 類中有任何一類佔比超過 50%,就是模式崩潰。
效能與實務提醒
在 Colab T4 上用 MNIST 128 batch、5 epoch 約 15 分鐘。單個 epoch 4,800 步(60,000/128 ≈ 469 batch/epoch),T4 大約 1.8 秒/epoch 的時間主要花在 256 與 512 channel 的轉置卷積上;如果想再加速,可以把 Generator 與 Discriminator 的中間 channel 從 512/256/128 降到 256/128/64,大約能再快 35%,但生成品質會略降。混合精度(torch.cuda.amp)能把 VRAM 占用從約 4 GB 壓到 2 GB,時間從 15 分鐘降到 11 分鐘左右,是 T4 上跑 GAN 的標準做法。
把 DCGAN 擴充到 CIFAR-10(3 通道彩色 32×32)會碰到兩個瓶頸:訓練時間與資料多樣性。CIFAR-10 的影像比 MNIST 複雜得多——一張影像同時有前景物件與背景、顏色分布也更多樣。要達到可辨識的生成品質通常需要 100 epoch 以上、T4 約 8–12 小時;FID 從 MNIST 的 18–25 會上升到 35–55 之間(實際數字會略有不同)。如果資料集換成更高解析度的影像(CelebA 64×64、LSUN Bedroom 128×128),生成品質會顯著惡化——這時就要考慮 StyleGAN 或 Diffusion(Day 36 開始會介紹)。
部署階段,DCGAN 的權重大小約 14 MB(fp32)、匯出 ONNX 約 14 MB。生成單張影像在 T4 上約 0.3 毫秒,在 CPU 上約 25 毫秒(生成器的轉置卷積在 CPU 上比判別器的卷積慢很多)。如果要部署到工業瑕疵檢測的「合成罕見瑕疵」任務,建議的設計是:先訓練一個 MVTec AD 子類別的 DCGAN(Day 39 會深入)、把生成的樣本用 MVTec AD 的分類器做品質過濾(信心 > 0.9 才保留),最後把過濾後的樣本混入訓練集。這個「GAN + 分類器篩選」的搭配在 2024 年的瑕疵資料增強文獻中仍是標準流程,比單純用 DCGAN 全部接收的資料有效。
小結
今天用 PyTorch 2.5 從零訓練了一個 DCGAN,從生成器與判別器的架構定義、到 MNIST 上的對抗式訓練迴圈、到「影像間距」的模式崩潰診斷、到三個對策(TTUR、β1 調整、標籤平滑)的實測比較。重點回顧:第一,DCGAN = 轉置卷積的生成器 + 卷積的判別器 + Adam(β1=0.5) + LeakyReLU(0.2),這四個元素缺一不可;第二,把影像正規化到 [-1, 1] 與 tanh 對齊是訓練初期的關鍵,忘記做會完全卡住;第三,GAN 的 loss 曲線會震盪是正常的,要看「影像間距」這個 proxy 來判斷模式崩潰;第四,當模式崩潰發生時,TTUR + 標籤平滑是最通用的兩個對策。明天我們會從 GAN 的「對抗式訓練」跳到 Diffusion 的「漸進式去噪」——它不用對抗賽局,而是用「加噪 → 學去噪」的對稱結構產生影像,數學基礎比 GAN 更穩定、訓練也比 GAN 更容易收斂。
結語
今天的重點是把 GAN 的對抗式訓練流程跑起來。我們先寫了 DCGAN 的標準架構(轉置卷積生成器 + 卷積判別器 + LeakyReLU + Adam(β1=0.5)),接著在 MNIST 上做了 5 epoch 的對抗式訓練、用「影像間距」診斷模式崩潰、最後實測三個對策(TTUR、β1=0.0、標籤平滑)的效果。讀完這篇你應該能回答:為什麼 DCGAN 用轉置卷積而不用 nearest-neighbor 上取樣?為什麼判別器要用 LeakyReLU(0.2) 而生成器用 ReLU?為什麼 GAN 的 loss 震盪是正常的、要看什麼指標?當模式崩潰發生時,三個對策各能解決哪一類問題?明天,我們會把 GAN 的「對抗式」換成 Diffusion 的「漸進去噪式」——DDPM 用一條馬可夫鏈把影像逐步加噪、再用一個神經網路學逆向去噪,整個訓練沒有對手、只有單純的 MSE 損失,是 2024 年生成模型的主流路線。
延伸資源
- Goodfellow 等人,2014,Generative Adversarial Networks(NeurIPS 2014):
https://arxiv.org/abs/1406.2661,GAN 原始論文,minimax 賽局與非對稱 loss 的設計動機。 - Radford 等人,2015,Unsupervised Representation Learning with Deep Convolutional GANs(ICLR 2016):
https://arxiv.org/abs/1511.06434,DCGAN 原始論文,轉置卷積、BatchNorm、LeakyReLU 的設計原則。 - Salimans 等人,2016,Improved Techniques for Training GANs(NeurIPS 2016):
https://arxiv.org/abs/1606.03498,標籤平滑、TTUR、feature matching、virtual batch normalization 等改良的原始論文。 - torch.nn.ConvTranspose2d 官方文件(PyTorch 2.5,2024):
https://pytorch.org/docs/stable/generated/torch.nn.ConvTranspose2d.html,轉置卷積的 kernel、stride、padding、output_padding 完整公式。 - PyTorch 官方 DCGAN 教學(2024):
https://pytorch.org/tutorials/beginner/dcgan_faces_tutorial.html,用 CelebA 臉部資料集訓練 DCGAN 的官方範例,從資料下載到視覺化都包好。 - MNIST 官方網站(CC BY-SA 3.0,Yann LeCun 維護):
http://yann.lecun.com/exdb/mnist/,60,000 張手寫數字影像的原始來源與授權。
留言
張貼留言