「本物そっくりの偽札を作る偽造者と、それを見破ろうとする鑑定士がいたとしたら、両者が競い合ううちに何が起きるでしょうか?」
偽造者はより精巧な偽札を作ろうとし、鑑定士はより正確に偽札を見抜こうとします。この競争を繰り返すうちに、偽造者の技術は本物と見分けがつかないレベルにまで到達します。この「敵対的なゲーム」のアイデアを深層学習に持ち込んだのが、GAN(Generative Adversarial Network、敵対的生成ネットワーク) です。
2014年にイアン・グッドフェローらによって提案されたGANは、画像生成の分野に革命をもたらしました。人の顔、風景画、アニメキャラクターなど、実在しないものを「本物と見分けがつかない」品質で生成できるようになったのです。GANの応用は画像生成にとどまらず、以下のような幅広い分野で活用されています。
- 画像超解像: 低解像度の画像から高解像度画像を生成する(SRGAN)
- スタイル変換: 写真を絵画風に変換する、昼の画像を夜に変換する(CycleGAN)
- データ拡張: 学習データが不足する医療画像やレアケースの合成データ生成
- 異常検知: 正常データの分布を学習し、逸脱するデータを異常として検出する
本記事の内容
- GANの基本的な仕組みと直感的理解
- ミニマックスゲームとしての数学的定式化
- 最適な識別器と生成器の理論的導出
- 学習アルゴリズムの詳細
- PyTorchによる実装(2次元ガウス混合と画像生成)
前提知識
この記事を読む前に、以下の記事を読んでおくと理解が深まります。
- ニューラルネットワークの基礎 — 多層パーセプトロンの仕組み
- 確率分布の基礎 — 確率密度関数と期待値
- 最尤推定 — モデルパラメータの推定
- KLダイバージェンス — 分布間の距離尺度
GANとは — 偽造者と鑑定士のゲーム
2つのネットワークの競争
GANの核となるアイデアは、2つのニューラルネットワークを「競わせる」ことで、データの分布を学習するというものです。この2つのネットワークは以下の役割を持ちます。
生成器(Generator)$G$: ランダムなノイズベクトル $\bm{z}$ を入力として受け取り、本物のデータに似たサンプルを生成するネットワークです。偽造者に相当します。生成器の目標は、識別器を騙せるほどリアルなデータを作ることです。
識別器(Discriminator)$D$: データを入力として受け取り、それが本物のデータか生成器が作った偽物かを判別するネットワークです。鑑定士に相当します。識別器の目標は、本物と偽物を正確に見分けることです。
この関係を日常の例で考えてみましょう。美術の世界で、ある画家(生成器)がピカソの作品の模倣を試みているとします。美術鑑定士(識別器)は本物のピカソの作品と画家の模倣作品を見比べて、どちらが本物かを判断します。画家は鑑定士に見破られるたびに技術を改善し、鑑定士は画家の腕前が上がるたびにより細かな違いを見極めるようになります。
この競争が十分に進むと、画家の模倣作品は鑑定士でも見分けがつかないレベルに到達します。これがGANの学習が目指すゴールです。数学的に言えば、生成器が出力するデータの分布 $p_g$ が、本物のデータの分布 $p_{\text{data}}$ に一致する状態です。
データの流れ
GANのデータフローをもう少し具体的に見てみましょう。
- ノイズのサンプリング: まず、簡単な分布(通常は標準正規分布)からノイズベクトル $\bm{z} \sim p_z(\bm{z})$ をサンプリングします
- データの生成: 生成器 $G$ がノイズ $\bm{z}$ を変換して偽のデータ $G(\bm{z})$ を作り出します
- 識別: 識別器 $D$ が、本物のデータ $\bm{x} \sim p_{\text{data}}$ と偽のデータ $G(\bm{z})$ の両方を受け取り、各データが本物である確率を出力します
- フィードバック: 識別器の判断結果に基づいて、両方のネットワークのパラメータを更新します
重要なのは、生成器は本物のデータを一度も直接見ないということです。生成器が受け取るフィードバックは、識別器の判断だけです。それにもかかわらず、間接的に本物のデータの分布を学習できるのがGANの驚くべき特性です。
GANの直感的な仕組みを理解したところで、次にこのアイデアを数学的に定式化しましょう。
ミニマックスゲームとしての定式化
目的関数
GANの学習は、ゲーム理論における二人ゼロ和ゲーム(two-player minimax game)として定式化されます。識別器 $D$ と生成器 $G$ が最適化する目的関数は以下のとおりです。
$$ \begin{equation} \min_G \max_D V(D, G) = \mathbb{E}_{\bm{x} \sim p_{\text{data}}(\bm{x})}[\ln D(\bm{x})] + \mathbb{E}_{\bm{z} \sim p_z(\bm{z})}[\ln(1 – D(G(\bm{z})))] \end{equation} $$
この式が何を表しているかを、各項の意味から理解しましょう。
第1項 $\mathbb{E}_{\bm{x} \sim p_{\text{data}}}[\ln D(\bm{x})]$: 本物のデータ $\bm{x}$ に対して識別器が出力する確率 $D(\bm{x})$ の対数の期待値です。$D(\bm{x})$ は「$\bm{x}$ が本物である確率」を表すので、本物のデータに対してはこの値が1に近い($\ln D(\bm{x})$ が0に近い)ことが望ましいです。
第2項 $\mathbb{E}_{\bm{z} \sim p_z}[\ln(1 – D(G(\bm{z})))]$: 生成器が作った偽のデータ $G(\bm{z})$ に対して、識別器が「偽物である」と判断する確率 $1 – D(G(\bm{z}))$ の対数の期待値です。
識別器の視点($\max_D$): 識別器は $V(D, G)$ を最大化したい。つまり、本物には $D(\bm{x}) \to 1$(第1項が大きく)、偽物には $D(G(\bm{z})) \to 0$(第2項が大きく)と判断したい。
生成器の視点($\min_G$): 生成器は $V(D, G)$ を最小化したい。つまり、識別器が偽物を本物と間違える $D(G(\bm{z})) \to 1$(第2項の $\ln(1-D(G(\bm{z}))) \to -\infty$)ようにしたい。
このように、識別器と生成器が正反対の目的を持って $V(D, G)$ を最適化する構造が「ミニマックスゲーム」です。
二値交差エントロピーとの関係
目的関数 $V(D, G)$ は、実は二値交差エントロピー(binary cross-entropy)の負の値と密接に関連しています。
データ点 $\bm{x}$ に対して、ラベル $y = 1$(本物)または $y = 0$(偽物)が付与されているとすると、二値交差エントロピーは
$$ \mathcal{L}_{\text{BCE}} = -[y \ln D(\bm{x}) + (1 – y)\ln(1 – D(\bm{x}))] $$
これは標準的な二値分類の損失関数です。GANの識別器の学習は、本物データ($y=1$)と偽データ($y=0$)の混合データに対する二値分類問題そのものなのです。
この接続は実装面で重要です。PyTorchの BCELoss をそのまま使えることを意味しています。
ミニマックスゲームの定式化を理解したところで、次にこのゲームの均衡点、すなわち最適な識別器と生成器の形を理論的に導出しましょう。
最適な識別器の導出
任意の $G$ に対する最適な $D$
生成器 $G$ を固定したとき、目的関数 $V(D, G)$ を最大化する最適な識別器 $D^*$ を求めます。
目的関数を積分の形で書き直すと
$$ V(D, G) = \int_{\bm{x}} p_{\text{data}}(\bm{x}) \ln D(\bm{x}) \, d\bm{x} + \int_{\bm{x}} p_g(\bm{x}) \ln(1 – D(\bm{x})) \, d\bm{x} $$
ここで $p_g(\bm{x})$ は生成器が定義する分布です。つまり $\bm{z} \sim p_z$ に対して $G(\bm{z})$ がしたがう分布です。
2つの積分を1つにまとめると
$$ V(D, G) = \int_{\bm{x}} \left[ p_{\text{data}}(\bm{x}) \ln D(\bm{x}) + p_g(\bm{x}) \ln(1 – D(\bm{x})) \right] d\bm{x} $$
各点 $\bm{x}$ で被積分関数を $D(\bm{x})$ について最大化すればよいです。$a = p_{\text{data}}(\bm{x})$, $b = p_g(\bm{x})$, $t = D(\bm{x})$ とおくと、最大化する関数は
$$ h(t) = a \ln t + b \ln(1 – t), \quad t \in (0, 1) $$
$t$ で微分してゼロとおくと
$$ h'(t) = \frac{a}{t} – \frac{b}{1-t} = 0 $$
これを $t$ について解きます。$a(1-t) = bt$ より $a = t(a+b)$ なので
$$ t^* = \frac{a}{a + b} $$
二階微分 $h”(t) = -a/t^2 – b/(1-t)^2 < 0$ なので、これは確かに最大値を与えます。元の変数に戻すと、最適な識別器は
$$ \begin{equation} D^*_G(\bm{x}) = \frac{p_{\text{data}}(\bm{x})}{p_{\text{data}}(\bm{x}) + p_g(\bm{x})} \end{equation} $$
この結果は直感的にも納得できます。ある点 $\bm{x}$ において、本物のデータが頻繁に出現し($p_{\text{data}}(\bm{x})$ が大きい)、偽物のデータが稀であれば($p_g(\bm{x})$ が小さい)、識別器は高い確率でそれを本物と判断します。逆に、偽物のデータが頻繁に出現する領域では、識別器は低い確率を出力します。
特に注目すべきは、$p_g = p_{\text{data}}$ のとき(生成器が完璧なとき)、$D^*(\bm{x}) = 1/2$ となることです。本物と偽物の分布が完全に一致していれば、識別器はもはや区別できず、コイン投げ(50%)と同じ判断しかできません。
最適な識別器の形がわかったところで、次にこれを目的関数に代入して、生成器にとっての最適化問題を分析しましょう。
最適な生成器の導出とナッシュ均衡
$D^*$ を代入した目的関数
最適な識別器 $D^*_G$ を目的関数に代入すると、生成器に関する最適化問題が得られます。
$$ C(G) = V(D^*_G, G) = \mathbb{E}_{\bm{x} \sim p_{\text{data}}}\left[\ln \frac{p_{\text{data}}(\bm{x})}{p_{\text{data}}(\bm{x}) + p_g(\bm{x})}\right] + \mathbb{E}_{\bm{x} \sim p_g}\left[\ln \frac{p_g(\bm{x})}{p_{\text{data}}(\bm{x}) + p_g(\bm{x})}\right] $$
ここで、$p_m(\bm{x}) = (p_{\text{data}}(\bm{x}) + p_g(\bm{x})) / 2$ という混合分布を導入します。すると
$$ C(G) = \mathbb{E}_{\bm{x} \sim p_{\text{data}}}\left[\ln \frac{p_{\text{data}}(\bm{x})}{2 \, p_m(\bm{x})}\right] + \mathbb{E}_{\bm{x} \sim p_g}\left[\ln \frac{p_g(\bm{x})}{2 \, p_m(\bm{x})}\right] $$
対数の性質を使って $\ln 2$ を分離すると
$$ C(G) = -\ln 4 + \mathbb{E}_{\bm{x} \sim p_{\text{data}}}\left[\ln \frac{p_{\text{data}}(\bm{x})}{p_m(\bm{x})}\right] + \mathbb{E}_{\bm{x} \sim p_g}\left[\ln \frac{p_g(\bm{x})}{p_m(\bm{x})}\right] $$
ここで現れた2つの期待値は、それぞれKLダイバージェンスです。
$$ \begin{equation} C(G) = -\ln 4 + D_{\text{KL}}(p_{\text{data}} \| p_m) + D_{\text{KL}}(p_g \| p_m) \end{equation} $$
ジェンセン・シャノン・ダイバージェンス
右辺の2つのKLダイバージェンスの和は、ジェンセン・シャノン・ダイバージェンス(Jensen-Shannon divergence, JSD)と呼ばれる量の2倍です。
$$ \begin{equation} \text{JSD}(p_{\text{data}} \| p_g) = \frac{1}{2} D_{\text{KL}}(p_{\text{data}} \| p_m) + \frac{1}{2} D_{\text{KL}}(p_g \| p_m) \end{equation} $$
したがって
$$ C(G) = -\ln 4 + 2 \, \text{JSD}(p_{\text{data}} \| p_g) $$
JSDはKLダイバージェンスと異なり、対称性 $\text{JSD}(p \| q) = \text{JSD}(q \| p)$ を持ち、常に非負 $\text{JSD} \geq 0$ で、$\text{JSD}(p \| q) = 0$ となるのは $p = q$ のときに限ります。
したがって、$C(G)$ の最小値は $-\ln 4$ であり、これは $p_g = p_{\text{data}}$ のとき(かつそのときに限り)達成されます。つまり、GANのミニマックスゲームのナッシュ均衡は
$$ p_g^* = p_{\text{data}}, \quad D^*(\bm{x}) = \frac{1}{2} \quad \forall \bm{x} $$
です。理論的には、十分な容量を持つ生成器と識別器が与えられ、学習が均衡に達すれば、生成器は本物のデータ分布を完全に再現できるのです。
この美しい理論的結果は、GANの学習が何を最適化しているのかを明確にしています。実際の学習アルゴリズムでこの均衡に到達する方法を、次に見ていきましょう。
学習アルゴリズム
交互最適化
GANの学習は、識別器と生成器を交互に更新する手続きで行います。各イテレーションでは以下のステップを実行します。
ステップ1: 識別器の更新($k$ ステップ)
- ミニバッチの本物データ $\{\bm{x}^{(1)}, \ldots, \bm{x}^{(m)}\}$ をサンプリング
- ノイズ $\{\bm{z}^{(1)}, \ldots, \bm{z}^{(m)}\}$ をサンプリングし、偽データ $G(\bm{z}^{(i)})$ を生成
- 識別器のパラメータ $\theta_D$ を以下の勾配で更新
$$ \nabla_{\theta_D} \frac{1}{m} \sum_{i=1}^{m} \left[\ln D(\bm{x}^{(i)}) + \ln(1 – D(G(\bm{z}^{(i)})))\right] $$
ステップ2: 生成器の更新(1ステップ)
- 新しいノイズ $\{\bm{z}^{(1)}, \ldots, \bm{z}^{(m)}\}$ をサンプリング
- 生成器のパラメータ $\theta_G$ を以下の勾配で更新
$$ \nabla_{\theta_G} \frac{1}{m} \sum_{i=1}^{m} \ln(1 – D(G(\bm{z}^{(i)}))) $$
実用上は $k = 1$ とすることが多く、識別器と生成器を1ステップずつ交互に更新します。
学習初期の勾配問題と実用的な修正
理論的な目的関数では、生成器は $\ln(1 – D(G(\bm{z})))$ を最小化します。しかし、学習初期には生成器が生み出すデータの質が低いため、識別器は容易に見分けられ $D(G(\bm{z})) \approx 0$ となります。このとき $\ln(1 – D(G(\bm{z}))) \approx \ln 1 = 0$ となり、勾配がほぼ消失してしまいます。
この問題を回避するために、実際の実装では生成器の目的関数を以下のように変更します。
$$ \text{最小化: } -\ln D(G(\bm{z})) \quad \text{(理論: 最小化 } \ln(1 – D(G(\bm{z}))) \text{)} $$
すなわち、「識別器を騙せない確率を最小化する」代わりに、「識別器を騙せる確率を最大化する」と読み替えます。$D(G(\bm{z})) \approx 0$ のとき $-\ln D(G(\bm{z})) \to \infty$ なので、大きな勾配が得られ、学習が進みやすくなります。
この修正は同じ均衡点 $p_g = p_{\text{data}}$ を持ちますが、学習のダイナミクスが改善されます。
学習の不安定性
GANの学習には、以下のような不安定性が知られています。
モード崩壊(Mode collapse): 生成器がデータ分布の一部のモード(山)しか学習せず、多様性が失われる現象です。たとえば、手書き数字の生成で特定の数字しか生成しなくなるような状態です。識別器が見抜けない「楽な」パターンに生成器が固執することで起こります。
振動(Oscillation): 識別器と生成器のパラメータが収束せず、振動し続ける現象です。一方が強くなると他方が弱くなり、均衡に到達しないことがあります。
勾配消失: 識別器が完璧に近い性能を持つと、生成器への勾配信号が消失します。前述の修正された目的関数はこの問題を軽減しますが、完全には解決しません。
これらの問題に対処するために、WGANやSpectral Normalizationなどの様々な改良手法が提案されています。
学習アルゴリズムの全体像を理解したところで、次にPyTorchを使って実際にGANを実装し、理論が実際にどのように動作するのかを確認しましょう。
Pythonでの実装 — 2次元ガウス混合データ
2次元のGAN
まず、GANの基本的な動作を確認するために、2次元のガウス混合分布を対象に実装します。2次元であれば生成結果を直接可視化でき、学習の進行を目で追えるという利点があります。
import numpy as np
import matplotlib.pyplot as plt
import torch
import torch.nn as nn
import torch.optim as optim
# 乱数シード
torch.manual_seed(42)
np.random.seed(42)
# --- 2次元ガウス混合分布からのサンプル生成 ---
def sample_real_data(n_samples):
"""8つのガウス分布を円状に配置した混合分布"""
n_clusters = 8
radius = 2.0
std = 0.05
angles = np.linspace(0, 2 * np.pi, n_clusters, endpoint=False)
centers = np.stack([radius * np.cos(angles), radius * np.sin(angles)], axis=1)
samples = []
for _ in range(n_samples):
idx = np.random.randint(n_clusters)
sample = np.random.normal(centers[idx], std)
samples.append(sample)
return torch.tensor(np.array(samples), dtype=torch.float32)
# --- 生成器の定義 ---
class Generator(nn.Module):
def __init__(self, noise_dim=2, hidden_dim=128, output_dim=2):
super().__init__()
self.net = nn.Sequential(
nn.Linear(noise_dim, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, output_dim),
)
def forward(self, z):
return self.net(z)
# --- 識別器の定義 ---
class Discriminator(nn.Module):
def __init__(self, input_dim=2, hidden_dim=128):
super().__init__()
self.net = nn.Sequential(
nn.Linear(input_dim, hidden_dim),
nn.LeakyReLU(0.2),
nn.Linear(hidden_dim, hidden_dim),
nn.LeakyReLU(0.2),
nn.Linear(hidden_dim, 1),
nn.Sigmoid(),
)
def forward(self, x):
return self.net(x)
上のコードでは、GANの基本要素を3つ定義しました。まず sample_real_data 関数は8つのガウス分布を円状に配置した混合分布からサンプルを生成します。これが生成器が再現すべき「本物のデータ」です。生成器 Generator は2次元のノイズを受け取り、2次元のデータ点を出力する3層のネットワークです。識別器 Discriminator は2次元のデータ点を受け取り、それが本物である確率(0~1)を出力します。識別器にはLeakyReLUを使用しています。これはGANの学習安定化のためによく用いられるテクニックで、負の入力に対しても小さな勾配を伝播させることで、勾配消失を防ぎます。
次に、学習ループを実装します。
# --- 学習パラメータ ---
noise_dim = 2
n_epochs = 5000
batch_size = 512
lr = 1e-3
# モデルの初期化
G = Generator(noise_dim=noise_dim)
D = Discriminator()
optimizer_G = optim.Adam(G.parameters(), lr=lr, betas=(0.5, 0.999))
optimizer_D = optim.Adam(D.parameters(), lr=lr, betas=(0.5, 0.999))
criterion = nn.BCELoss()
# --- 学習ループ ---
d_losses, g_losses = [], []
for epoch in range(n_epochs):
# 本物データのサンプリング
real_data = sample_real_data(batch_size)
real_labels = torch.ones(batch_size, 1)
fake_labels = torch.zeros(batch_size, 1)
# === 識別器の更新 ===
z = torch.randn(batch_size, noise_dim)
fake_data = G(z).detach() # 生成器の勾配は不要
d_real = D(real_data)
d_fake = D(fake_data)
d_loss = criterion(d_real, real_labels) + criterion(d_fake, fake_labels)
optimizer_D.zero_grad()
d_loss.backward()
optimizer_D.step()
# === 生成器の更新 ===
z = torch.randn(batch_size, noise_dim)
fake_data = G(z)
d_fake = D(fake_data)
# 生成器は識別器を騙したい -> fake を real と判断させたい
g_loss = criterion(d_fake, real_labels)
optimizer_G.zero_grad()
g_loss.backward()
optimizer_G.step()
d_losses.append(d_loss.item())
g_losses.append(g_loss.item())
学習ループの各イテレーションで、まず識別器を更新し、次に生成器を更新しています。識別器の更新では、本物データに対して「本物」ラベル(1)、偽データに対して「偽物」ラベル(0)を使ってBCELossを計算します。fake_data = G(z).detach() の .detach() は、識別器の更新時に生成器の勾配が計算されないようにするための重要な操作です。生成器の更新では、生成した偽データを識別器に通し、「本物」ラベルとの損失を計算します。これが前述の「$-\ln D(G(\bm{z}))$ を最小化する」修正された目的関数に対応しています。
# --- 学習結果の可視化 ---
fig, axes = plt.subplots(1, 3, figsize=(16, 5))
# (a) 損失の推移
ax = axes[0]
window = 100
d_smooth = np.convolve(d_losses, np.ones(window)/window, mode='valid')
g_smooth = np.convolve(g_losses, np.ones(window)/window, mode='valid')
ax.plot(d_smooth, label='Discriminator', alpha=0.8)
ax.plot(g_smooth, label='Generator', alpha=0.8)
ax.set_xlabel('Epoch', fontsize=12)
ax.set_ylabel('Loss', fontsize=12)
ax.set_title('Training Loss', fontsize=13)
ax.legend(fontsize=10)
ax.grid(True, alpha=0.3)
# (b) 本物データ vs 生成データ
ax = axes[1]
real_samples = sample_real_data(1000).numpy()
with torch.no_grad():
z = torch.randn(1000, noise_dim)
fake_samples = G(z).numpy()
ax.scatter(real_samples[:, 0], real_samples[:, 1], s=5, alpha=0.5,
label='Real', color='blue')
ax.scatter(fake_samples[:, 0], fake_samples[:, 1], s=5, alpha=0.5,
label='Generated', color='red')
ax.set_xlabel('$x_1$', fontsize=12)
ax.set_ylabel('$x_2$', fontsize=12)
ax.set_title('Real vs Generated Data', fontsize=13)
ax.legend(fontsize=10)
ax.set_aspect('equal')
ax.grid(True, alpha=0.3)
# (c) 識別器の決定境界
ax = axes[2]
xx, yy = np.meshgrid(np.linspace(-3.5, 3.5, 200),
np.linspace(-3.5, 3.5, 200))
grid = torch.tensor(np.c_[xx.ravel(), yy.ravel()], dtype=torch.float32)
with torch.no_grad():
d_values = D(grid).numpy().reshape(xx.shape)
im = ax.contourf(xx, yy, d_values, levels=20, cmap='RdBu_r', alpha=0.8)
plt.colorbar(im, ax=ax, label='D(x)')
ax.scatter(real_samples[:, 0], real_samples[:, 1], s=3, alpha=0.3,
color='blue', label='Real')
ax.scatter(fake_samples[:, 0], fake_samples[:, 1], s=3, alpha=0.3,
color='red', label='Generated')
ax.set_xlabel('$x_1$', fontsize=12)
ax.set_ylabel('$x_2$', fontsize=12)
ax.set_title('Discriminator Decision Surface', fontsize=13)
ax.legend(fontsize=9)
ax.set_aspect('equal')
plt.tight_layout()
plt.savefig('gan_2d_result.png', dpi=150, bbox_inches='tight')
plt.show()
この可視化結果から、GANの学習が正しく進んでいることを確認できます。
-
左図(損失の推移): 識別器の損失と生成器の損失が学習を通じて拮抗する値に収束しています。これはミニマックスゲームが均衡に近づいていることを示唆しています。学習初期には識別器の損失が急速に下がり(容易に見分けられる)、その後生成器が改善するにつれて識別器の損失が上昇して均衡に至っています
-
中央図(本物vs生成データ): 青い点(本物データ)が形成する8つのクラスタに対して、赤い点(生成データ)が同様のクラスタ構造を形成しています。生成器が8モードのガウス混合分布の構造を捉えていることがわかります。全てのモードが再現されているかどうかは、モード崩壊が起きていないかの重要な確認点です
-
右図(識別器の決定面): 識別器の出力値(D(x))を色で示しています。理想的なナッシュ均衡では $D(\bm{x}) = 0.5$ となるべきですが、実際にはデータが存在する領域(クラスタ付近)で値が0.5に近くなっている様子が見て取れます。データの存在しない領域では識別器の出力が極端な値をとりますが、これは学習に影響しません
MNIST画像の生成
画像生成への拡張
2次元データで基本的な動作を確認したので、次にMNIST手書き数字データセットを使って画像を生成してみましょう。画像データはピクセル値の高次元ベクトルであり、生成器と識別器のネットワーク構造を適切に設計する必要があります。
import torch
import torch.nn as nn
import torch.optim as optim
from torchvision import datasets, transforms
from torch.utils.data import DataLoader
import numpy as np
import matplotlib.pyplot as plt
torch.manual_seed(42)
# --- データの準備 ---
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,)) # [0,1] -> [-1,1]
])
dataset = datasets.MNIST(root='./data', train=True, download=True,
transform=transform)
dataloader = DataLoader(dataset, batch_size=128, shuffle=True)
# --- ネットワーク定義 ---
noise_dim = 64
img_dim = 28 * 28 # 784
class GeneratorMNIST(nn.Module):
def __init__(self):
super().__init__()
self.net = nn.Sequential(
nn.Linear(noise_dim, 256),
nn.LeakyReLU(0.2),
nn.Linear(256, 512),
nn.LeakyReLU(0.2),
nn.Linear(512, 1024),
nn.LeakyReLU(0.2),
nn.Linear(1024, img_dim),
nn.Tanh(), # 出力を[-1, 1]に
)
def forward(self, z):
return self.net(z)
class DiscriminatorMNIST(nn.Module):
def __init__(self):
super().__init__()
self.net = nn.Sequential(
nn.Linear(img_dim, 1024),
nn.LeakyReLU(0.2),
nn.Dropout(0.3),
nn.Linear(1024, 512),
nn.LeakyReLU(0.2),
nn.Dropout(0.3),
nn.Linear(512, 256),
nn.LeakyReLU(0.2),
nn.Dropout(0.3),
nn.Linear(256, 1),
nn.Sigmoid(),
)
def forward(self, x):
return self.net(x.view(-1, img_dim))
G = GeneratorMNIST()
D = DiscriminatorMNIST()
optimizer_G = optim.Adam(G.parameters(), lr=2e-4, betas=(0.5, 0.999))
optimizer_D = optim.Adam(D.parameters(), lr=2e-4, betas=(0.5, 0.999))
criterion = nn.BCELoss()
MNIST用のネットワークでは、いくつかの工夫を加えています。生成器の出力活性化関数に Tanh を使い、出力を $[-1, 1]$ の範囲にしています。これは入力画像を同じ範囲に正規化しているためです。識別器には Dropout(確率0.3)を追加しています。これは識別器が強くなりすぎるのを防ぎ、生成器に十分な勾配が伝わるようにするためのテクニックです。学習率は $2 \times 10^{-4}$、Adamオプティマイザの $\beta_1 = 0.5$ はGAN学習のベストプラクティスとして広く使われている設定です。
# --- 学習ループ ---
n_epochs = 50
d_losses, g_losses = [], []
for epoch in range(n_epochs):
for batch_idx, (real_imgs, _) in enumerate(dataloader):
batch_size_curr = real_imgs.size(0)
real_labels = torch.ones(batch_size_curr, 1)
fake_labels = torch.zeros(batch_size_curr, 1)
# 識別器の更新
z = torch.randn(batch_size_curr, noise_dim)
fake_imgs = G(z).detach()
d_real = D(real_imgs)
d_fake = D(fake_imgs)
d_loss = criterion(d_real, real_labels) + criterion(d_fake, fake_labels)
optimizer_D.zero_grad()
d_loss.backward()
optimizer_D.step()
# 生成器の更新
z = torch.randn(batch_size_curr, noise_dim)
fake_imgs = G(z)
d_fake = D(fake_imgs)
g_loss = criterion(d_fake, real_labels)
optimizer_G.zero_grad()
g_loss.backward()
optimizer_G.step()
d_losses.append(d_loss.item())
g_losses.append(g_loss.item())
if (epoch + 1) % 10 == 0:
print(f"Epoch [{epoch+1}/{n_epochs}] D_loss: {d_loss.item():.4f} "
f"G_loss: {g_loss.item():.4f}")
学習の構造は2次元の場合と全く同じです。各ミニバッチについて識別器を1回更新し、続いて生成器を1回更新します。実データの画像は自動的に $[-1, 1]$ に正規化されており、生成器の出力も Tanh で同じ範囲になっているため、識別器が公平に比較できます。
# --- 生成画像の可視化 ---
fig, axes = plt.subplots(2, 2, figsize=(14, 12))
# (a) 損失の推移
ax = axes[0, 0]
ax.plot(d_losses, label='Discriminator', alpha=0.8)
ax.plot(g_losses, label='Generator', alpha=0.8)
ax.set_xlabel('Epoch', fontsize=12)
ax.set_ylabel('Loss', fontsize=12)
ax.set_title('Training Loss (MNIST GAN)', fontsize=13)
ax.legend(fontsize=10)
ax.grid(True, alpha=0.3)
# (b) 生成画像のグリッド
ax = axes[0, 1]
with torch.no_grad():
z = torch.randn(64, noise_dim)
generated = G(z).view(-1, 28, 28).numpy()
grid_img = np.zeros((8 * 28, 8 * 28))
for i in range(8):
for j in range(8):
grid_img[i*28:(i+1)*28, j*28:(j+1)*28] = generated[i*8+j]
ax.imshow(grid_img, cmap='gray')
ax.set_title('Generated Digits (64 samples)', fontsize=13)
ax.axis('off')
# (c) 潜在空間の補間
ax = axes[1, 0]
z1 = torch.randn(1, noise_dim)
z2 = torch.randn(1, noise_dim)
n_interp = 10
interp_imgs = []
with torch.no_grad():
for alpha in np.linspace(0, 1, n_interp):
z_interp = (1 - alpha) * z1 + alpha * z2
img = G(z_interp).view(28, 28).numpy()
interp_imgs.append(img)
interp_grid = np.concatenate(interp_imgs, axis=1)
ax.imshow(interp_grid, cmap='gray')
ax.set_title('Latent Space Interpolation', fontsize=13)
ax.axis('off')
# (d) 本物vs生成の比較
ax = axes[1, 1]
real_batch, _ = next(iter(dataloader))
real_grid = np.zeros((2 * 28, 8 * 28))
with torch.no_grad():
z = torch.randn(8, noise_dim)
fake_batch = G(z).view(-1, 28, 28).numpy()
for j in range(8):
real_grid[0:28, j*28:(j+1)*28] = real_batch[j, 0].numpy()
real_grid[28:56, j*28:(j+1)*28] = fake_batch[j]
ax.imshow(real_grid, cmap='gray')
ax.set_title('Real (top) vs Generated (bottom)', fontsize=13)
ax.axis('off')
plt.tight_layout()
plt.savefig('gan_mnist_result.png', dpi=150, bbox_inches='tight')
plt.show()
この可視化から、MNIST画像生成GANの学習結果を確認できます。
-
左上(損失の推移): 識別器と生成器の損失がエポックごとに推移しています。GANの学習では損失が単調に減少するのではなく、両者が拮抗しながら推移するのが正常な振る舞いです。損失がゼロに近づく場合は、どちらか一方が圧倒していることを示し、学習が崩壊している可能性があります
-
右上(生成画像): 64枚の生成画像をグリッド状に並べています。各画像は手書き数字に類似した構造を持っており、数字として認識可能なものが多数含まれています。ただし、全結合層のみのGANでは画像の解像度や鮮明さに限界があり、ぼやけた画像やノイジーな画像も混在しています。これはDCGANなどの畳み込み層ベースのアーキテクチャで改善されます
-
左下(潜在空間の補間): 2つのランダムなノイズベクトル間を線形補間して生成した画像の系列です。画像がスムーズに遷移している場合、生成器が潜在空間の構造を適切に学習していることを示しています。これはGANが単にデータを暗記しているのではなく、データの内部表現を獲得していることの証拠です
-
右下(本物vs生成の比較): 上段が本物のMNIST画像、下段が同数の生成画像です。全結合GANの限界として、生成画像は本物ほど鮮明ではありませんが、大まかな構造(ストロークの位置、数字の形)は捉えています
GANの理論的課題と改良の方向性
モード崩壊の数学的理解
モード崩壊はGANの最も深刻な問題の一つです。これを数学的に理解するために、ミニマックスゲームの非凸性に注目します。
実際の学習では、$G$ と $D$ はニューラルネットワークでパラメータ化されるため、目的関数は $\theta_G$ と $\theta_D$ について非凸です。ミニマックス定理の条件(凸凹性)を満たさないため
$$ \min_{\theta_G} \max_{\theta_D} V(\theta_D, \theta_G) \neq \max_{\theta_D} \min_{\theta_G} V(\theta_D, \theta_G) $$
が一般に成り立ちます。交互最適化では、生成器が $\max_D V$ の $G$ に対する最小化を行おうとしますが、識別器も同時に変化するため、生成器がデータ分布の一部のモードに「ロック」されてしまうことがあります。
JSDの限界
ジェンセン・シャノン・ダイバージェンスは、$p_{\text{data}}$ と $p_g$ の台(support)が重ならないとき(高次元空間ではほぼ確実に起こる)、飽和して意味のある勾配を提供しません。
$$ \text{JSD}(p_{\text{data}} \| p_g) = \ln 2 \quad \text{when } \text{supp}(p_{\text{data}}) \cap \text{supp}(p_g) = \emptyset $$
この問題は、ワッサースタイン距離を使うWGANで解決されます。ワッサースタイン距離は、分布の台が重ならない場合でも滑らかな勾配を提供します。
主な改良手法
GANの理論的・実用的な課題に対して、以下の改良が提案されています。
| 手法 | 解決する問題 | 主なアイデア |
|---|---|---|
| DCGAN | 画像の品質 | 畳み込み層の導入、アーキテクチャガイドライン |
| WGAN | 学習の不安定性 | ワッサースタイン距離の使用 |
| WGAN-GP | WGAN の制約処理 | 勾配ペナルティによるリプシッツ制約 |
| Spectral Normalization | 識別器の爆発 | 識別器の重み行列のスペクトルノルム正規化 |
| Progressive GAN | 高解像度画像 | 段階的な解像度の向上 |
これらの手法は、GAN の基本理論を土台として、より安定で高品質な生成を目指しています。
まとめ
本記事では、GAN(敵対的生成ネットワーク)の理論と実装について解説しました。
- GANの基本原理: 生成器(偽造者)と識別器(鑑定士)の敵対的な学習により、データの分布を学習する
- ミニマックスゲーム: GANの学習はゲーム理論の枠組みで定式化され、目的関数は $\min_G \max_D V(D, G)$ で表される
- 最適な識別器: $D^*(\bm{x}) = p_{\text{data}}(\bm{x}) / (p_{\text{data}}(\bm{x}) + p_g(\bm{x}))$ であり、生成が完璧なときは $D^* = 1/2$
- JSDとの関係: GANの目的関数はジェンセン・シャノン・ダイバージェンスの最小化と等価であり、均衡点は $p_g = p_{\text{data}}$
- 学習の実際: 交互最適化による学習、勾配問題の修正、モード崩壊などの課題がある
- 実装: 2次元ガウス混合データとMNIST画像で、GANが多峰分布や手書き数字の構造を学習できることを確認した
GANの基本理論を理解した次のステップとして、以下の記事も参考にしてください。
- DCGANの理論と実装 — 畳み込みを導入した画像生成の改善
- WGAN — ワッサースタイン距離による安定学習 — JSDの限界を克服する手法
- VAEの実装チュートリアル — 生成モデルの別のアプローチ