ML Atlas

07 · Architektury · 4 min czytania · aktualizacja

Czym są połączenia rezydualne (skip connections) i dlaczego pozwalają trenować bardzo głębokie sieci?

W skrócie

Połączenie rezydualne dodaje wejście bloku do jego wyjścia: y = x + F(x). Gradient ma wtedy skrót przez sieć, więc da się trenować setki warstw.

Co to jest

Połączenie rezydualne (skip connection, połączenie skrótowe) to ścieżka, która omija blok warstw i dodaje jego wejście do wyjścia: y = x + F(x). Blok F nie musi więc wytwarzać całej nowej reprezentacji, tylko poprawkę (resztę, ang. residual) do tego, co już jest. Architekturę opartą na takich blokach zaproponowali He, Zhang, Ren i Sun w 2015 roku jako ResNet; w 2016 roku praca ukazała się na CVPR.

Intuicja: zamiast przepisywać cały tekst od nowa na każdym etapie redakcji, każdy redaktor nanosi tylko poprawki na marginesie. Jeśli redaktor nie ma nic do dodania, tekst przechodzi dalej bez zmian. W sieci bez skrótów każda warstwa musi „przepisać” sygnał i każda może go trochę zepsuć.

Dziś połączenia rezydualne są wszędzie: w głębokich CNN, w każdym bloku transformera (wokół uwagi i wokół warstwy MLP), w sieciach dyfuzyjnych typu U-Net. Bez nich współczesne modele językowe z dziesiątkami warstw praktycznie by się nie uczyły.

Mechanizm — dlaczego tak działa

Przed ResNetem obserwowano paradoks degradacji: sieć głębsza miała większy błąd nie tylko na teście, ale i na zbiorze treningowym. To nie było przeuczenie — głębsza sieć po prostu gorzej się optymalizowała. A przecież teoretycznie powinna być co najmniej tak dobra: wystarczy, by dodatkowe warstwy liczyły tożsamość. Problem w tym, że stos nieliniowych warstw z losowymi wagami bardzo trudno nauczyć dokładnej tożsamości.

Połączenie rezydualne odwraca sytuację. Tożsamość jest teraz domyślna: wystarczy F(x) ≈ 0, czyli małe wagi, a takie sieć ma niemal od startu. Każdy blok zaczyna jako „prawie nic nie zmieniam” i uczy się tylko tego, co faktycznie poprawia wynik. Głębsza sieć przestaje być trudniejsza do wytrenowania od płytszej.

Drugie spojrzenie to przepływ gradientu. Pochodna bloku to ∂y/∂x = 1 + ∂F/∂x. W zwykłej sieci gradient z wyjścia do wejścia jest iloczynem pochodnych wszystkich warstw; jeśli każda ma wzmocnienie 0,5, po 20 warstwach zostaje 0,5^20 ≈ 0,000001. W sieci rezydualnej iloczyn czynników (1 + ∂F/∂x) zawiera składnik równy dokładnie 1 — gradient ma „autostradę” prosto do wczesnych warstw, niezależnie od tego, co robią bloki. To łagodzi problem zanikającego gradientu.

Trzecie spojrzenie: sieć rezydualna zachowuje się jak zespół wielu płytszych ścieżek o różnej długości, bo sygnał może każdy blok przejść albo ominąć. Badania krajobrazu funkcji straty pokazały też, że skróty wyraźnie go wygładzają.

Zastrzeżenie: same skróty nie wystarczą. Suma x + F(x) przy wielu blokach może rosnąć bez kontroli, dlatego łączy się je z normalizacją (BatchNorm w CNN, LayerNorm w transformerach) i ostrożną inicjalizacją. Gdy wymiary x i F(x) się różnią (np. inna liczba kanałów), skrót zawiera projekcję — splot 1×1.

Na przykładzie

Na zbiorze Digits 8×8 zbudowaliśmy sieć gęstą o szerokości 32 neuronów, z LayerNorm w każdym bloku, trenowaną 30 epok spadkiem gradientu (krok 0,05), w dwóch wersjach: zwykłej i z połączeniami rezydualnymi. Trzy ziarna losowe na każdą konfigurację. Przy 5 blokach obie wersje osiągają 96–98% na teście. Przy 20 blokach zwykła sieć utknęła na 10% dla wszystkich trzech ziaren — to poziom zgadywania — a wersja rezydualna osiągnęła 90–97%. Przy 50 blokach rezydualna wciąż uczy się do 85–96%. To miniatura zjawiska degradacji opisanego przez He i współpracowników.

Skróty pozwalają też oszczędzać. Blok „z wąskim gardłem” w ResNet-50 na 256 kanałach składa się ze splotów 1×1 (256→64), 3×3 (64→64) i 1×1 (64→256): 16 384 + 36 864 + 16 384 = 69 632 wag. Dwa zwykłe sploty 3×3 na 256 kanałach to 2·256·256·9 = 1 179 648 wag, czyli 17 razy więcej. Dlatego ResNet-50 ma tylko 25,6 mln parametrów, a 152-warstwowy ResNet-152 — 60,2 mln, mniej niż 16-warstwowy VGG-16 (138 mln).

Dane: Digits (ręcznie pisane cyfry 8×8)

W praktyce

  • W PyTorch blok to po prostu return x + self.f(x) w metodzie forward; gotowe sieci są w torchvision.models.resnet18, resnet50 itd.
  • Gdy blok zmienia liczbę kanałów lub rozdzielczość, skrót potrzebuje projekcji: nn.Conv2d(c_in, c_out, kernel_size=1, stride=2).
  • W transformerach dominuje wariant pre-norm: x = x + attn(norm(x)), który stabilniej trenuje głębokie modele niż post-norm z oryginalnej pracy.
  • Sztuczka inicjalizacyjna: zerowanie ostatniej warstwy normalizacji w każdym bloku sprawia, że na starcie każdy blok to dokładna tożsamość.
  • Częsty błąd: dodanie ReLU po sumie w transformerze lub zapomnienie o normalizacji — sygnał rośnie z każdym blokiem i trening się rozjeżdża.

Najczęstsze pytania

Czy połączenia rezydualne rozwiązują problem zanikającego gradientu?
W dużej mierze tak: składnik tożsamości daje gradientowi bezpośrednią ścieżkę do wczesnych warstw. Pełną stabilność daje dopiero połączenie z normalizacją i rozsądną inicjalizacją.
Czym różni się ResNet od DenseNet?
ResNet dodaje wejście do wyjścia bloku. DenseNet skleja (konkatenuje) wyjścia wszystkich wcześniejszych warstw, więc każda warstwa widzi wszystkie poprzednie mapy cech. Obie idee służą łatwiejszemu przepływowi informacji.
Czy każda głęboka sieć powinna mieć skróty?
Przy więcej niż kilkunastu warstwach praktycznie tak. Dla płytkich sieci (kilka warstw) nie są konieczne, choć zwykle nie szkodzą.

Źródła

  • He, Zhang, Ren, Sun „Deep Residual Learning for Image Recognition”, CVPR 2016, arXiv:1512.03385.
  • He, Zhang, Ren, Sun „Identity Mappings in Deep Residual Networks”, ECCV 2016.
  • Veit, Wilber, Belongie „Residual Networks Behave Like Ensembles of Relatively Shallow Networks”, NeurIPS 2016.
  • Li, Xu, Taylor, Studer, Goldstein „Visualizing the Loss Landscape of Neural Nets”, NeurIPS 2018.
  • Zhang i in. „Dive into Deep Learning”, d2l.ai, rozdz. 8.6 („Residual Networks”).

Zobacz też