生成對抗網路(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()

在上述範例中:

與 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 在勢均力敵的情況下,互相競爭並持續進步;而為了讓兩個網路維持勢均力敵,不會有一方變強的太快,是需要不少訓練技巧配合的。其中,前面的幾個範例已經用到的有:

前述範例中,尚未展示的則有:

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()

在上述範例中:

GAN 還有其他不少的著名變體,如: