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 metodzieforward; gotowe sieci są wtorchvision.models.resnet18,resnet50itd. - 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”).