GAN
Kunna träna en liten GAN och förklara instabiliteten.
Förkunskaper
Intuition
En GAN är två nät i ett spel:
- Generatorn G tar brus och producerar en bild. Mål: lura diskriminatorn.
- Diskriminatorn D får en bild och avgör om den är äkta eller genererad. Mål: avslöja generatorn.
De tränas växelvis. Om det fungerar blir G så bra att D inte kan skilja — men jämvikten är svår att nå.
Tre klassiska problem:
- Mode collapse: G hittar ett fåtal bilder som lurar D och producerar bara dem. Variationen dör.
- Instabilitet: om D blir för bra får G inga användbara gradienter och slutar lära sig.
- Ingen likelihood: det finns inget mått som säger «hur bra är modellen?» — därför är det svårt att veta när man ska sluta träna.
Formellt
Målfunktionen (Goodfellow m.fl. 2014):
Vid optimal minimerar detta Jensen–Shannon-divergensen mellan datafördelningen och generatorns fördelning. Problemet: när fördelningarna knappt överlappar (vilket de inte gör i början) är JS-divergensen konstant och gradienten noll — därför används i praktiken det «non-saturating» målet .
Varför WGAN hjälper: Wasserstein-avståndet ger användbara gradienter även utan överlapp, och kritikerns värde korrelerar med bildkvalitet — ett mått att följa, vilket vanlig GAN saknar. Kräver Lipschitz-begränsning (gradient penalty).
Var GAN fortfarande används: superupplösning, bild-till-bild, och överallt där en framåtpassning krävs (realtid). Diffusion vann på kvalitet och stabilitet men behöver många steg; GAN-sampling är ett enda anrop.
Utvärdering: FID (jämför statistik i ett förtränat nätverks featurerum) är standard, men fångar inte mode collapse väl — komplettera med precision/recall för generativa modeller och mänsklig bedömning.
Kod
import torch, torch.nn as nn
G = nn.Sequential(nn.Linear(64, 256), nn.ReLU(), nn.Linear(256, 784), nn.Tanh())
D = nn.Sequential(nn.Linear(784, 256), nn.LeakyReLU(0.2), nn.Linear(256, 1))
optG = torch.optim.Adam(G.parameters(), lr=2e-4, betas=(0.5, 0.999))
optD = torch.optim.Adam(D.parameters(), lr=2e-4, betas=(0.5, 0.999))
bce = nn.BCEWithLogitsLoss()
for xb in loader:
# 1. Diskriminatorn: äkta → 1, falska → 0
z = torch.randn(xb.size(0), 64)
falska = G(z).detach()
lossD = bce(D(xb), torch.ones(len(xb), 1)) + bce(D(falska), torch.zeros(len(xb), 1))
lossD.backward(); optD.step(); optD.zero_grad()
# 2. Generatorn: non-saturating — maximera log D(G(z))
z = torch.randn(xb.size(0), 64)
lossG = bce(D(G(z)), torch.ones(len(xb), 1))
lossG.backward(); optG.step(); optG.zero_grad()
# Diagnostik: om lossD → 0 har D vunnit och G får inga gradienter.
# Om de genererade bilderna blir allt mer lika varandra: mode collapse.
Behärskning innebär
- Förklarar generator, diskriminator och minimax-målet
- Känner igen mode collapse och instabilitet
- Vet varför GAN saknar likelihood
Logga in för att göra övningarna och bygga upp din behärskning.
Källor
- arXiv — Generative Adversarial Networks — arXiv (öppen åtkomst; licens per artikel)
- arXiv — Wasserstein GAN — arXiv (öppen åtkomst; licens per artikel)