生成對抗網路(Generative Adversarial Network,GAN)是一種以對抗訓練為核心的生成模型,最初由 Goodfellow 等人於 2014 年提出,並被廣泛應用於圖像生成、風格轉換、影像修復、超解析度等多媒體處理任務。GAN 由兩個網路組成,一個是負責從隨機雜訊生成資料的生成器(generator),另一個是負責判斷輸入資料為真實或生成的判別器(discriminator);兩者在訓練過程中相互博弈,generator 試圖欺騙 discriminator,discriminator 則試圖正確區分真假;而我們期待透過這樣的對抗過程,來讓 generator 逐漸學會生成與真實資料難以區分的樣本。
與先前介紹的 Auto Encoder 相比,GAN 的 generator 扮演的角色類似於 AE 的 decoder,同樣負責從低維度表示重建出資料;然而 AE 的訓練目標是最小化重建誤差,GAN 則以對抗損失取而代之,不依賴逐像素的比較,因此往往能生成更清晰、更具視覺真實感的結果。本篇章將介紹 GAN 的基礎概念以及相關應用。
與 Auto Encoder 相似,generator 和 discriminator 的網路架構可以是 MLP、CNN,或其他任何架構。最早的 GAN 實作採用 MLP,概念簡單但生成品質有限;因此,Radford 等人於 2015 年提出了 DCGAN (Deep Convolutional GAN) ,將 generator 和 discriminator 兩者均改為卷積網路,有效提升了訓練穩定性與生成品質,並成為後續許多 GAN 變形的架構基礎。以下是用 DCGAN 並基於 Fashion-MNIST 資料集生成新圖案的範例:
import matplotlib.pyplot as plt import torch import torch.nn as nn from torch.utils.data import DataLoader from torchvision import datasets, transforms LATENT_DIM = 64 EPOCHS = 10 BATCH_SIZE = 64 LR = 0.0002 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,)) ]) train_dataset = datasets.FashionMNIST(root='./data', train=True, download=True, transform=transform) train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True) class Generator(nn.Module): def __init__(self, latent_dim): super().__init__() self.fc = nn.Linear(latent_dim, 128 * 7 * 7) self.net = nn.Sequential( nn.ConvTranspose2d(128, 64, kernel_size=4, stride=2, padding=1), nn.BatchNorm2d(64), nn.ReLU(), nn.ConvTranspose2d(64, 1, kernel_size=4, stride=2, padding=1), nn.Tanh() ) def forward(self, z): x = self.fc(z) x = x.view(-1, 128, 7, 7) return self.net(x) class Discriminator(nn.Module): def __init__(self): super().__init__() self.net = nn.Sequential( nn.Conv2d(1, 32, kernel_size=4, stride=2, padding=1), nn.LeakyReLU(0.2), nn.Conv2d(32, 64, kernel_size=4, stride=2, padding=1), nn.BatchNorm2d(64), nn.LeakyReLU(0.2) ) self.fc = nn.Linear(64 * 7 * 7, 1) def forward(self, x): x = self.net(x) x = x.view(-1, 64 * 7 * 7) return torch.sigmoid(self.fc(x)) G = Generator(LATENT_DIM).to(device) D = Discriminator().to(device) criterion = nn.BCELoss() optimizer_G = torch.optim.Adam(G.parameters(), lr=LR, betas=(0.5, 0.999)) optimizer_D = torch.optim.Adam(D.parameters(), lr=LR, betas=(0.5, 0.999)) G.train() D.train() for epoch in range(EPOCHS): print(f'Epoch {epoch+1}/{EPOCHS}') for real_imgs, _ in train_loader: real_imgs = real_imgs.to(device) batch_size = real_imgs.size(0) real_labels = torch.ones(batch_size, 1).to(device) fake_labels = torch.zeros(batch_size, 1).to(device) # Update D z = torch.randn(batch_size, LATENT_DIM).to(device) fake_imgs = G(z) loss_D = criterion(D(real_imgs), real_labels) + criterion(D(fake_imgs.detach()), fake_labels) optimizer_D.zero_grad() loss_D.backward() optimizer_D.step() # Update G loss_G = criterion(D(fake_imgs), real_labels) optimizer_G.zero_grad() loss_G.backward() optimizer_G.step() print(f'\tLoss D: {loss_D.item():.4f} Loss G: {loss_G.item():.4f}') G.eval() fixed_z = torch.randn(64, LATENT_DIM).to(device) with torch.no_grad(): generated = G(fixed_z).cpu() generated = (generated + 1) / 2 grid = generated.view(8, 8, 28, 28).permute(0, 2, 1, 3).reshape(8 * 28, 8 * 28) plt.imshow(grid, cmap='gray') plt.axis('off') plt.show()在上述範例中:
- 進行 transform 的目的是把影像的取值範圍,從 [0, 1] 變為 [-1, 1],以方便訓練。如果你使用了不同的影像,或者某些特定的預訓練模型,都可以或者可能需要更換 transform 的設定。
- G 的輸入是雜點,輸出是產生的影像,訓練的目標是盡量騙過 D;D 的輸入是一張影像,輸出是該影像是否為真實資料,訓練的目標是盡量分辨真偽。
- 我們會讓 G 和 D 交替進行訓練,以讓兩者在勢均力敵的情況下,互相競爭並持續進步。如果 D 一開始就變得太強,則會讓 G 拿到的梯度接近消失;而如果 G 一開始就太強,則可能僅產生少量樣本就騙過 D,而不會有動力進步。
與 Auto Encoder 相仿,我們當然也可以在 GAN 當中加上 condition 的輸入。以下是用 Conditional DCGAN 並基於 Fashion-MNIST 資料集生成新圖案的範例:
import matplotlib.pyplot as plt import torch import torch.nn as nn from torch.utils.data import DataLoader from torchvision import datasets, transforms NUM_CLASSES = 10 LATENT_DIM = 64 EPOCHS = 10 BATCH_SIZE = 64 LR = 0.0002 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,)) ]) train_dataset = datasets.FashionMNIST(root='./data', train=True, download=True, transform=transform) train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True) def to_onehot(labels, num_classes=NUM_CLASSES): out = torch.zeros(labels.size(0), num_classes) out[torch.arange(labels.size(0)), labels] = 1 return out class Generator(nn.Module): def __init__(self, latent_dim, num_classes): super().__init__() self.fc = nn.Linear(latent_dim + num_classes, 128 * 7 * 7) self.net = nn.Sequential( nn.ConvTranspose2d(128, 64, kernel_size=4, stride=2, padding=1), nn.BatchNorm2d(64), nn.ReLU(), nn.ConvTranspose2d(64, 1, kernel_size=4, stride=2, padding=1), nn.Tanh() ) def forward(self, z, c): x = self.fc(torch.cat([z, c], dim=1)) x = x.view(-1, 128, 7, 7) return self.net(x) class Discriminator(nn.Module): def __init__(self, num_classes): super().__init__() self.net = nn.Sequential( nn.Conv2d(1 + num_classes, 32, kernel_size=4, stride=2, padding=1), nn.LeakyReLU(0.2), nn.Conv2d(32, 64, kernel_size=4, stride=2, padding=1), nn.BatchNorm2d(64), nn.LeakyReLU(0.2) ) self.fc = nn.Linear(64 * 7 * 7, 1) def forward(self, x, c): c_map = c.view(-1, NUM_CLASSES, 1, 1).expand(-1, -1, 28, 28) x = self.net(torch.cat([x, c_map], dim=1)) x = x.view(-1, 64 * 7 * 7) return torch.sigmoid(self.fc(x)) G = Generator(LATENT_DIM, NUM_CLASSES).to(device) D = Discriminator(NUM_CLASSES).to(device) criterion = nn.BCELoss() optimizer_G = torch.optim.Adam(G.parameters(), lr=LR, betas=(0.5, 0.999)) optimizer_D = torch.optim.Adam(D.parameters(), lr=LR, betas=(0.5, 0.999)) G.train() D.train() for epoch in range(EPOCHS): print(f'Epoch {epoch + 1}/{EPOCHS}') for real_imgs, labels in train_loader: real_imgs = real_imgs.to(device) labels = labels.to(device) batch_size = real_imgs.size(0) c = to_onehot(labels).to(device) real_labels = torch.ones(batch_size, 1).to(device) fake_labels = torch.zeros(batch_size, 1).to(device) # Update D z = torch.randn(batch_size, LATENT_DIM).to(device) fake_imgs = G(z, c) loss_D = criterion(D(real_imgs, c), real_labels) + criterion(D(fake_imgs.detach(), c), fake_labels) optimizer_D.zero_grad() loss_D.backward() optimizer_D.step() # Update G loss_G = criterion(D(fake_imgs, c), real_labels) optimizer_G.zero_grad() loss_G.backward() optimizer_G.step() print(f'\tLoss D: {loss_D.item():.4f} Loss G: {loss_G.item():.4f}') G.eval() n = 10 fixed_z = torch.randn(n, LATENT_DIM).to(device) fixed_z = fixed_z.repeat_interleave(n, dim=0) fixed_c = to_onehot(torch.arange(n).repeat(n)).to(device) with torch.no_grad(): generated = G(fixed_z, fixed_c).cpu() generated = (generated + 1) / 2 grid = generated.view(n, n, 28, 28).permute(0, 2, 1, 3).reshape(n * 28, n * 28) plt.imshow(grid, cmap='gray') plt.axis('off') plt.show()在上述範例中,我們讓 G 和 D 都接收 condition 的輸入,以讓 D 不只判斷輸入影像是否真實,還要判斷其是否與給定的 condition 一致;你也可以只讓 G 接收 condition,不過如此一來,G 只要生成任何看起來真實的影像就能騙過 D,不需要真的遵守 condition,有可能因此讓訓練比較不穩定。
前面提到,我們希望 generator 和 discriminator 在勢均力敵的情況下,互相競爭並持續進步;而為了讓兩個網路維持勢均力敵,不會有一方變強的太快,是需要不少訓練技巧配合的。其中,前面的幾個範例已經用到的有:
- 讓 generator 的訓練目標,從將「1 - discriminator 把假影像判為假影像的機率」最大化,改為將「discriminator 把假影像判為真影像的機率」最大化。這在 Binary Cross Entropy loss 的公式中,分別相當於 (1 - y) * log(1 - ŷ) 和 y * log(ŷ),因此它們在數學上等價,但因為前者在訓練早期 discriminator 先變強的時候,梯度會比較小,因此讓 generator 不易訓練。
- Discriminator 使用 LeakyReLU 而非 ReLU,避免負值區梯度為 0。Discriminator 使用 LeakyReLU 的原因是,梯度不為 0 就容易在訓練過程中持續地給予 generator 回饋;而 generator 沒有一起使用的原因是用 ReLU 造成幾個神經元死亡,對生成任務的影響有限,反倒是能讓負值通過的 LeakyReLU 不一定有幫助。
- BatchNorm 的位置在 generator 除最後一層外都加,discriminator 除第一層外都加。兩者的原因都是避免讓單一影像的特性被同一 batch 內的其他資料影響,generator 可以保留每一張生成影像自己的特性,discriminator 則可以更好的看到每一張輸入影像自己的特性。
前述範例中,尚未展示的則有:
- Real label 從 1 改為 0.9,避免 discriminator 過於自信。
- 對 discriminator 輸入的影像加入少量雜訊,讓其不要太容易分辨真假。
- 對 discriminator 的每一層做權重正規化,使得網路不會對輸入的微小變化反應過度,此方法在實作上可使用 nn.utils.spectral_norm 包住每一層網路,例如 nn.utils.spectral_norm(nn.Conv2d(1, 32)),可以取代或搭配 BatchNorm 來使用。
GAN 的另外一種變體是 Wasserstein GAN,簡稱 WGAN,主要是在網路架構和訓練技巧上面有些不同。以下範例,是由先前 DCGAN 的範例修改而來:
import matplotlib.pyplot as plt import torch import torch.nn as nn from torch.utils.data import DataLoader from torchvision import datasets, transforms LATENT_DIM = 64 EPOCHS = 10 BATCH_SIZE = 64 LR = 0.0002 N_CRITIC = 5 LAMBDA_GP = 10 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,)) ]) train_dataset = datasets.FashionMNIST(root='./data', train=True, download=True, transform=transform) train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True) class Generator(nn.Module): def __init__(self, latent_dim): super().__init__() self.fc = nn.Linear(latent_dim, 128 * 7 * 7) self.net = nn.Sequential( nn.ConvTranspose2d(128, 64, kernel_size=4, stride=2, padding=1), nn.BatchNorm2d(64), nn.ReLU(), nn.ConvTranspose2d(64, 1, kernel_size=4, stride=2, padding=1), nn.Tanh() ) def forward(self, z): x = self.fc(z) x = x.view(-1, 128, 7, 7) return self.net(x) class Critic(nn.Module): def __init__(self): super().__init__() self.net = nn.Sequential( nn.Conv2d(1, 32, kernel_size=4, stride=2, padding=1), nn.LeakyReLU(0.2), nn.Conv2d(32, 64, kernel_size=4, stride=2, padding=1), nn.LeakyReLU(0.2) ) self.fc = nn.Linear(64 * 7 * 7, 1) def forward(self, x): x = self.net(x) x = x.view(-1, 64 * 7 * 7) return self.fc(x) def gradient_penalty(critic, real_imgs, fake_imgs): alpha = torch.rand(real_imgs.size(0), 1, 1, 1).to(device) interpolates = (alpha * real_imgs + (1 - alpha) * fake_imgs).requires_grad_(True) d_interpolates = critic(interpolates) gradients = torch.autograd.grad( outputs=d_interpolates, inputs=interpolates, grad_outputs=torch.ones_like(d_interpolates), create_graph=True )[0] gradients = gradients.view(gradients.size(0), -1) return ((gradients.norm(2, dim=1) - 1) ** 2).mean() G = Generator(LATENT_DIM).to(device) C = Critic().to(device) optimizer_G = torch.optim.Adam(G.parameters(), lr=LR, betas=(0.5, 0.999)) optimizer_C = torch.optim.Adam(C.parameters(), lr=LR, betas=(0.5, 0.999)) G.train() C.train() for epoch in range(EPOCHS): print(f'Epoch {epoch + 1}/{EPOCHS}') for i, (real_imgs, _) in enumerate(train_loader): real_imgs = real_imgs.to(device) batch_size = real_imgs.size(0) # Update Critic z = torch.randn(batch_size, LATENT_DIM).to(device) fake_imgs = G(z) gp = gradient_penalty(C, real_imgs, fake_imgs.detach()) loss_C = C(fake_imgs.detach()).mean() - C(real_imgs).mean() + LAMBDA_GP * gp optimizer_C.zero_grad() loss_C.backward() optimizer_C.step() # Update Generator if i % N_CRITIC == 0: z = torch.randn(batch_size, LATENT_DIM).to(device) loss_G = -C(G(z)).mean() optimizer_G.zero_grad() loss_G.backward() optimizer_G.step() print(f'\tLoss C: {loss_C.item():.4f} Loss G: {loss_G.item():.4f}') G.eval() fixed_z = torch.randn(64, LATENT_DIM).to(device) with torch.no_grad(): generated = G(fixed_z).cpu() generated = (generated + 1) / 2 grid = generated.view(8, 8, 28, 28).permute(0, 2, 1, 3).reshape(8 * 28, 8 * 28) plt.imshow(grid, cmap='gray') plt.axis('off') plt.show()在上述範例中:
- Discriminator 在 WGAN 的情境中,通常會改稱 Critic;其網路結構稍有變化(請自行對照),且訓練目標之一為將真假影像的平均分數差異最大化。
- 承上,另一個訓練目標,是讓網路在處理「中間(藉由在真假影像之間做插值來產生)」影像時的梯度不要太大或太小,若是偏離 1 的話就要受到懲罰。
- Generator 每隔數個 batches 才會更新一次,這是為了讓 Critic 在 generator 要被更新前已經足夠準確。此處沒有「如果 D 一開始就變得太強,則會讓 G 拿到的梯度接近消失」的顧慮,是因為上一條提到的梯度懲罰的設計,可以讓梯度在 Critic 很強的時候也不會消失。
GAN 還有其他不少的著名變體,如:
- Pix2Pix:由 Isola 等人於 2017 年提出,將 Conditional GAN 的條件從類別標籤推廣為整張圖像,實現成對圖像之間的轉換,例如將語意分割圖轉換為真實場景、將素描轉換為彩色圖像等。
- CycleGAN:由 Zhu 等人於 2017 年提出,在 Pix2Pix 的基礎上去除了成對資料的需求,改以兩組 generator 和 discriminator 互相配合,並引入 cycle consistency loss 約束轉換的可逆性,使模型能在沒有對應標注的情況下學習兩個圖像域之間的風格轉換,例如馬與斑馬之間的互換。
- SRGAN:由 Ledig 等人於 2017 年提出,將 GAN 應用於單張圖像超解析度任務,generator 負責將低解析度圖像重建為高解析度版本,並引入感知損失(perceptual loss)取代逐像素的重建誤差,使生成結果在視覺上更清晰自然,而非僅在數值上接近目標。
- StyleGAN:由 NVIDIA 於 2019 年提出,以漸進式生成架構為基礎,透過將 latent vector 映射為風格向量並注入各層,實現對生成圖像不同層次特徵(如整體姿態、臉部特徵、膚色等)的細粒度控制,在高解析度人臉生成任務上達到當時最先進的效果。