ML Atlas

07 · Architektury · 4 min czytania · aktualizacja

Czym są generatywne sieci przeciwstawne (GAN) i jak generator uczy się od dyskryminatora?

W skrócie

GAN to dwie rywalizujące sieci: generator tworzy fałszywe próbki, dyskryminator odróżnia je od prawdziwych. Ten pojedynek uczy generator realizmu.

Co to jest

Generatywna sieć przeciwstawna (GAN, generative adversarial network) to układ dwóch sieci neuronowych trenowanych przeciwko sobie. Generator zamienia losowy szum na próbkę (np. obraz twarzy). Dyskryminator dostaje na przemian próbki prawdziwe i wygenerowane i ocenia, czy dana próbka jest prawdziwa. Generator uczy się tak, by oszukać dyskryminator, a dyskryminator — by nie dać się oszukać. Pomysł przedstawili Goodfellow i współpracownicy w 2014 roku.

Intuicja: fałszerz banknotów i ekspert w banku. Na początku fałszerz robi kiepskie podróbki, a ekspert łatwo je wyłapuje. Każda wpadka mówi fałszerzowi, co poprawić; każda udana podróbka uczy eksperta, na co jeszcze patrzeć. Jeśli obaj się rozwijają, podróbki stają się nie do odróżnienia od oryginału. Ważne: generator nigdy nie widzi prawdziwych danych bezpośrednio — uczy się wyłącznie z sygnału dyskryminatora.

GAN-y przez kilka lat wyznaczały granicę realizmu generowanych obrazów (np. StyleGAN tworzący fotorealistyczne twarze). Dziś w generowaniu obrazów dominują modele dyfuzyjne, ale GAN-y wciąż są używane tam, gdzie liczy się szybkość — generują obraz jednym przejściem sieci.

Mechanizm — dlaczego tak działa

Formalnie trening to gra minimaksowa z funkcją wartości V(D, G) = E_x[log D(x)] + E_z[log(1 − D(G(z)))], gdzie D(x) to prawdopodobieństwo, że x jest prawdziwe. Dyskryminator maksymalizuje V (dobrze rozpoznaje obie klasy), generator ją minimalizuje (chce, by D(G(z)) było bliskie 1). W praktyce oba modele aktualizuje się na zmianę, po jednym lub kilku krokach gradientu.

Dlaczego to w ogóle prowadzi do realistycznych danych? Dla ustalonego generatora optymalny dyskryminator to D*(x) = p_data(x) / (p_data(x) + p_g(x)) — stosunek gęstości prawdziwych i wygenerowanych danych. Podstawiając go do V, Goodfellow i in. pokazali, że generator minimalizuje wtedy dywergencję Jensena–Shannona między rozkładem danych a rozkładem generatora. Minimum osiąga się, gdy p_g = p_data: dyskryminator zwraca wszędzie 0,5, a V = −log 4. Dyskryminator działa więc jak uczona, elastyczna miara „jak bardzo się różnimy”, której gradient wskazuje generatorowi, w którą stronę przesunąć próbki.

W praktyce generator minimalizuje −log D(G(z)) zamiast log(1 − D(G(z))). Na początku treningu dyskryminator łatwo odrzuca podróbki, a wtedy oryginalna funkcja jest płaska i gradient prawie znika; zmodyfikowana wersja daje silny sygnał właśnie wtedy, gdy generator jest słaby. To tzw. strata nienasycająca (non-saturating).

GAN-y są znane z trudnego treningu. Załamanie trybów (mode collapse): generator odkrywa kilka próbek, które oszukują dyskryminator, i produkuje tylko je — np. wyłącznie jedną cyfrę zamiast dziesięciu. Niestabilność: gra dwóch graczy nie musi zbiegać, parametry mogą krążyć. Brak wiarygodnej miary postępu: strata generatora nie mówi, czy obrazy są lepsze. Rozwiązania to m.in. Wasserstein GAN (inna miara odległości rozkładów), normalizacja spektralna, kara gradientowa i starannie dobrane architektury (DCGAN, StyleGAN).

Na przykładzie

Niech prawdziwe dane mają rozkład N(0, 1), a generator na razie produkuje N(2, 1). Optymalny dyskryminator w punkcie x = 0 zwraca 0,88 (tu dane prawdziwe są dużo gęstsze), w x = 1 dokładnie 0,5 (obie gęstości równe), w x = 2 — 0,12. Gradient tej funkcji „pcha” próbki generatora w lewo, w stronę prawdziwych danych. Gdy generator dojdzie do N(0, 1), dyskryminator wszędzie zwraca 0,5 i nie ma już żadnej wskazówki — to punkt równowagi, w którym V = −log 4 ≈ −1,386.

Dlaczego strata nienasycająca pomaga? Załóżmy, że na początku treningu D(G(z)) = 0,01. Pochodna log(1 − D) po logicie dyskryminatora wynosi −D = −0,01, a pochodna −log D wynosi −(1 − D) = −0,99. Ta sama sytuacja, a sygnał dla generatora jest 99 razy silniejszy.

W praktyce

  • W PyTorch obie sieci trenuje się dwoma optymalizatorami; stratę liczy nn.BCEWithLogitsLoss z etykietą 1 dla prawdziwych i 0 dla fałszywych (a przy kroku generatora — z etykietą 1 dla fałszywych).
  • Sprawdzone ustawienia z DCGAN: Adam z krokiem 0,0002 i β1 = 0,5, normalizacja wsadowa, LeakyReLU w dyskryminatorze.
  • Jakość ocenia się miarą FID (odległość statystyk cech prawdziwych i wygenerowanych obrazów), a nie stratą.
  • Regularnie oglądaj próbki z tego samego, ustalonego szumu — szybko zobaczysz załamanie trybów.
  • Dla stabilności warto zacząć od WGAN-GP lub normalizacji spektralnej (torch.nn.utils.spectral_norm).

Najczęstsze pytania

Czym GAN różni się od autoenkodera wariacyjnego (VAE)?
VAE maksymalizuje dolne ograniczenie wiarygodności danych i ma koder, więc daje stabilny trening i sensowną przestrzeń ukrytą, ale często rozmyte obrazy. GAN nie modeluje wiarygodności wprost, za to generuje ostrzejsze próbki kosztem trudniejszego treningu.
Co to jest załamanie trybów?
Sytuacja, w której generator produkuje tylko niewielki wycinek możliwych danych, bo te próbki wystarczają do oszukania dyskryminatora. Rozkład wygenerowanych danych jest wtedy dużo uboższy niż prawdziwy.
Czy GAN-y są jeszcze używane?
Tak, choć rzadziej jako główny generator obrazów. Są cenione za szybkość (jedno przejście sieci), stosuje się je w superrozdzielczości, edycji obrazów i jako dodatkowa strata przeciwstawna w innych modelach, np. w koderach obrazu dla modeli dyfuzyjnych.

Źródła

  • Goodfellow i in. „Generative Adversarial Nets”, NeurIPS 2014, arXiv:1406.2661.
  • Radford, Metz, Chintala „Unsupervised Representation Learning with Deep Convolutional Generative Adversarial Networks”, ICLR 2016.
  • Arjovsky, Chintala, Bottou „Wasserstein Generative Adversarial Networks”, ICML 2017.
  • Karras, Laine, Aila „A Style-Based Generator Architecture for Generative Adversarial Networks”, CVPR 2019.
  • Goodfellow, Bengio, Courville „Deep Learning”, MIT Press, 2016, rozdz. 20.10.4.

Zobacz też