はじめに
VAE (variational Autoencoder) は、深層学習とベイズ推論を組み合わせた生成モデルです。 2 つのネットワークに対する GAN とは異なり、VAE は構造化された潜在空間を学習し、内挿、制御された生成、および正確な尤度推定を可能にします。
1. オートエンコーダーの要約
import torch.nn as nn
class Autoencoder(nn.Module):
def __init__(self):
super().__init__()
self.encoder = nn.Sequential(
nn.Linear(784, 256),
nn.ReLU(),
nn.Linear(256, 64), # bottleneck
)
self.decoder = nn.Sequential(
nn.Linear(64, 256),
nn.ReLU(),
nn.Linear(256, 784),
nn.Sigmoid(),
)
def forward(self, x):
z = self.encoder(x) # compress
x_hat = self.decoder(z) # reconstruct
return x_hat
# Loss: ||x - x_hat||² (reconstruction error)
問題: オートエンコーダーの潜在空間が断続的である → ランダムな z をサンプリングして生成できない。
2. VAE — 中心となるアイデア
┌─────────────────────────────────────────────────────────┐
│ VAE Architecture │
│ │
│ x ──→ Encoder ──→ μ, σ ──→ z = μ + σ·ε ──→ Decoder → x̂│
│ ↑ │
│ ε ~ N(0, 1) │
│ (reparameterization trick) │
│ │
│ Loss = Reconstruction + KL Divergence │
│ = ||x - x̂||² + KL(q(z|x) || p(z)) │
└─────────────────────────────────────────────────────────┘
ELBO — 証拠の下限
$$\mathcal{L} = \mathbb{E}{q(z|x)}[\log p(x|z)] - D{KL}(q(z|x) || p(z))$$
- 再構成損失: $\mathbb{E}_{q(z|x)}[\log p(x|z)]$ — デコードは入力と同じである必要があります
- KL 発散: $D_{KL}(q(z|x) || p(z))$ — 潜在分布を $\mathcal{N}(0, I)$ に近づける
3. VAE を実装する
class VAE(nn.Module):
def __init__(self, input_dim=784, latent_dim=20):
super().__init__()
# Encoder
self.fc1 = nn.Linear(input_dim, 400)
self.fc_mu = nn.Linear(400, latent_dim) # mean
self.fc_logvar = nn.Linear(400, latent_dim) # log variance
# Decoder
self.fc3 = nn.Linear(latent_dim, 400)
self.fc4 = nn.Linear(400, input_dim)
def encode(self, x):
h = torch.relu(self.fc1(x))
mu = self.fc_mu(h)
logvar = self.fc_logvar(h)
return mu, logvar
def reparameterize(self, mu, logvar):
"""Reparameterization trick: z = μ + σ · ε"""
std = torch.exp(0.5 * logvar)
eps = torch.randn_like(std) # ε ~ N(0, 1)
return mu + eps * std
def decode(self, z):
h = torch.relu(self.fc3(z))
return torch.sigmoid(self.fc4(h))
def forward(self, x):
mu, logvar = self.encode(x.view(-1, 784))
z = self.reparameterize(mu, logvar)
x_hat = self.decode(z)
return x_hat, mu, logvar
def vae_loss(x_hat, x, mu, logvar):
# Reconstruction loss (BCE)
recon = nn.functional.binary_cross_entropy(
x_hat, x.view(-1, 784), reduction='sum'
)
# KL divergence: -0.5 * Σ(1 + log(σ²) - μ² - σ²)
kl = -0.5 * torch.sum(1 + logvar - mu.pow(2) - logvar.exp())
return recon + kl
4. 再パラメータ化のトリック
Vấn đề: z ~ q(z|x) → sampling không differentiable → không thể backprop
Giải pháp: z = μ + σ · ε, với ε ~ N(0, 1)
- μ, σ là output của encoder (differentiable)
- ε là random noise (không phụ thuộc parameters)
→ Gradient có thể flow qua μ, σ
# ❌ Không thể backprop qua sampling
z = torch.distributions.Normal(mu, std).sample()
# ✅ Reparameterization trick
eps = torch.randn_like(std)
z = mu + eps * std # gradient flows through mu and std
5. 潜在宇宙探査
補間
def interpolate(model, x1, x2, steps=10):
"""Interpolate between 2 images in latent space"""
model.eval()
with torch.no_grad():
mu1, _ = model.encode(x1.view(-1, 784))
mu2, _ = model.encode(x2.view(-1, 784))
images = []
for alpha in torch.linspace(0, 1, steps):
z = (1 - alpha) * mu1 + alpha * mu2
img = model.decode(z)
images.append(img.view(28, 28))
return images # smooth transition từ x1 → x2
ランダム生成
def generate(model, num_images=16):
"""Generate new images by sampling from latent space"""
model.eval()
with torch.no_grad():
z = torch.randn(num_images, 20) # sample từ N(0, I)
images = model.decode(z)
return images.view(num_images, 28, 28)
6. 条件付き VAE (CVAE)
class ConditionalVAE(nn.Module):
"""VAE conditioned on label → kiểm soát generation"""
def __init__(self, input_dim=784, latent_dim=20, num_classes=10):
super().__init__()
# Encoder nhận cả x và label
self.fc1 = nn.Linear(input_dim + num_classes, 400)
self.fc_mu = nn.Linear(400, latent_dim)
self.fc_logvar = nn.Linear(400, latent_dim)
# Decoder nhận z và label
self.fc3 = nn.Linear(latent_dim + num_classes, 400)
self.fc4 = nn.Linear(400, input_dim)
def encode(self, x, y_onehot):
h = torch.relu(self.fc1(torch.cat([x, y_onehot], dim=1)))
return self.fc_mu(h), self.fc_logvar(h)
def decode(self, z, y_onehot):
h = torch.relu(self.fc3(torch.cat([z, y_onehot], dim=1)))
return torch.sigmoid(self.fc4(h))
# Generate digit "7":
y = torch.zeros(1, 10)
y[0, 7] = 1 # one-hot cho số 7
z = torch.randn(1, 20)
img = model.decode(z, y) # → ảnh số 7
7. VQ-VAE — ベクトル量子化 VAE
Ý tưởng: Thay continuous latent → discrete codebook
- Encoder output → tìm nearest codebook vector
- Codebook: tập các learnable vectors
- Decoder nhận discrete code → reconstruct
Ưu điểm:
- Avoid posterior collapse
- Codebook = "vocabulary" của visual concepts
- Nền tảng cho DALL-E 1 (VQ-VAE + Transformer)
class VectorQuantizer(nn.Module):
def __init__(self, num_embeddings=512, embedding_dim=64):
super().__init__()
self.codebook = nn.Embedding(num_embeddings, embedding_dim)
def forward(self, z):
# Tìm nearest codebook vector
distances = torch.cdist(z, self.codebook.weight)
indices = distances.argmin(dim=-1)
z_q = self.codebook(indices)
# Straight-through estimator
z_q = z + (z_q - z).detach()
return z_q, indices
8. VAE 対 GAN
| 特長 | VAE | ガン |
|---|---|---|
| トレーニング | 安定した単一の目標 | 不安定、ミニマックス |
| 生み出される品質 | ぼやけた | シャープ |
| 潜在空間 | 構造的でスムーズ | 非構造化 |
| 可能性 | トラクタブル (ELBO) | 難治性 |
| 多様性 | 良い | モード崩壊のリスク |
| 補間 | スムーズ | 予測不可能 |
| 使用例 | 潜在的な操作 | 高品質な合成 |
概要
| コンセプト | 説明 |
|---|---|
| VAE | 確率的潜在空間を備えたエンコーダ-デコーダ |
| エルボ | 証拠の下限 = 偵察 + KL |
| 再パラメータ化 | z = μ + σ·ε — サンプリングによる逆伝播を有効にする |
| KLダイバージェンス | N(0, I) |
| CVA | 条件付き VAE — 制御された生成 |
| VQ-VAE | 離散コードブック — DALL-E 1 プラットフォーム |
📌 次の投稿: 拡散モデル — ゼロからの数学、直感、および DDPM。