ML Atlas

06 · Sieci · 4 min czytania · Interaktywne · aktualizacja

Czym różnią się SGD, momentum i Adam i który optymalizator wybrać?

W skrócie

Optymalizator zamienia gradient na krok. SGD idzie wzdłuż gradientu, momentum dodaje rozpęd, a Adam skaluje krok każdej wagi jej typowym gradientem.

Co to jest

Optymalizator to algorytm, który na podstawie gradientu (i historii gradientów) decyduje, jak zmienić wagi w jednym kroku. SGD robi krok proporcjonalny do bieżącego gradientu; momentum dokłada wykładniczo wygaszaną średnią poprzednich kroków; Adam (Kingma i Ba 2014) łączy momentum z osobnym skalowaniem kroku dla każdej wagi.

Optymalizatory są częścią treningu każdej sieci neuronowej; w treningu dużych modeli językowych standardem jest AdamW.

Intuicja: SGD to piechur, który w każdym kroku patrzy tylko pod nogi. Momentum to kula tocząca się po zboczu — nabiera rozpędu na długich spadkach i nie odbija się od każdej nierówności. Adam to piechur, który dodatkowo dostosowuje długość kroku do terenu: na stromym drobi, na płaskim wydłuża krok.

Mechanizm — dlaczego tak działa

SGD: w ← w − η g, gdzie g to gradient z mini-batcha. Problem: krajobraz straty ma różną krzywiznę w różnych kierunkach. W kierunkach stromych ten sam η jest za duży (zygzaki), w płaskich za mały (pełzanie).

Momentum: v ← β v + g, w ← w − η v, typowo β = 0,9. Prędkość v uśrednia gradienty z około 1/(1 − β) = 10 ostatnich kroków. W kierunkach, gdzie gradient zmienia znak (zygzak w wąskiej dolinie), składowe się znoszą; w kierunkach stałych sumują się aż do η g/(1 − β), czyli nawet 10 razy dłuższego kroku. Rozpęd tłumi zygzaki i przyspiesza w wąwozach.

Adam utrzymuje dwie średnie — m (średnia gradientów, jak momentum, β₁ = 0,9) i v (średnia kwadratów gradientów, β₂ = 0,999) — i robi krok w ← w − η m/(√v + ε). Dzielenie przez pierwiastek średniego kwadratu gradientu oznacza, że każda waga dostaje krok o podobnej wielkości niezależnie od skali swojego gradientu: wagi cech o dużych wartościach (duże gradienty) dostają mniejszy efektywny krok, wagi cech o małych wartościach — większy. Adam częściowo wyrównuje więc skale cech sam z siebie, dlatego jego przewaga nad SGD rośnie, gdy wejścia nie są standaryzowane. Nie zastępuje to standaryzacji: wyrównuje skalę kroku, nie kształt doliny. Korekta obciążenia (bias correction) poprawia m i v na początku treningu, gdy średnie są jeszcze bliskie zera.

Zastrzeżenie: Adam bywa gorszy w generalizacji niż dobrze dostrojony SGD z momentum na obrazach, a kara L2 dodana do gradientu działa w nim inaczej niż prawdziwy weight decay — stąd AdamW (Loshchilov i Hutter 2019), który odejmuje weight decay od wag bezpośrednio.

Na przykładzie

Na Digits 8×8 (1347 obrazów treningowych, 450 testowych, random_state=0) trenowałem MLPClassifier z 64 neuronami ukrytymi przez 20 epok z batchem 64, zmieniając tylko optymalizator. Zwykły SGD (learning rate 0,01) kończy ze stratą treningową 1,41 i trafnością testową 84,4% — idzie w dobrą stronę, ale powoli. Ten sam SGD z momentum 0,9 schodzi do straty 0,16 i trafności 95,8%. Adam z domyślnym lr = 0,001 daje 95,6%, a z lr = 0,01 — 98,2%.

Potem pomnożyłem piksele przez 160, żeby wejścia miały wartości do 160 zamiast do 1. SGD — z momentum i bez — przestał się uczyć: trafność 10,0–10,4%, czyli zgadywanie wśród 10 cyfr. Adam z lr = 0,001 nadal osiągnął 91,3%, a z lr = 0,01 — 94,4%: adaptacyjne skalowanie kroku uratowało trening, choć wynik i tak jest gorszy niż na dobrze przeskalowanych danych.

Ta ilustracja działa w przeglądarce z włączonym JavaScriptem: SGD, momentum i Adam na tej samej dolinie: kto dochodzi do dna pierwszy i dlaczego.

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

W praktyce

  • Domyślne PyTorch: Adam(lr=1e-3, betas=(0.9, 0.999), eps=1e-8); SGD(lr, momentum=0) — momentum trzeba włączyć jawnie (zwykle 0,9). W scikit-learn MLPClassifier(solver='adam') to domyślny wybór.
  • LLM: AdamW z β₂ = 0,95 zamiast 0,999 dla stabilności, weight decay około 0,1, learning rate rzędu 10⁻⁴ z rozgrzewką i harmonogramem kosinusowym.
  • Pierwszy wybór dla nowego problemu: Adam lub AdamW z lr = 0,001. SGD z momentum i harmonogramem, gdy liczy się ostatnia setna procenta na obrazach.
  • Adam potrzebuje dwóch dodatkowych buforów wielkości modelu — przy dużych modelach to istotny koszt pamięci.
  • Typowy błąd: przeniesienie learning rate z SGD (0,1) do Adama — dla Adama to zwykle o dwa rzędy wielkości za dużo.

Najczęstsze pytania

Adam czy SGD — który wybrać?
Na start Adam (lub AdamW): jest mniej wrażliwy na learning rate i skalę cech i działa dobrze od razu. SGD z momentum i dobrym harmonogramem potrafi dać lepszą generalizację w widzeniu komputerowym, ale wymaga więcej strojenia. W LLM standardem jest AdamW.
Co robi momentum w optymalizatorze?
Utrzymuje „prędkość” — wygaszaną średnią poprzednich gradientów — i dodaje ją do kroku. W kierunkach, gdzie gradient oscyluje, oscylacje się znoszą; w kierunkach stałych krok rośnie nawet 10-krotnie (przy β = 0,9). Trening idzie gładziej i szybciej przez wąskie doliny.
Czym różni się Adam od AdamW?
W Adamie kara L2 jest dodawana do gradientu i potem dzielona przez √v, więc wagi o dużych gradientach są słabiej regularyzowane. AdamW odejmuje weight decay od wag osobno, poza adaptacyjnym skalowaniem, co daje spójną regularyzację i lepsze wyniki — dlatego to on jest standardem.

Źródła

  • Kingma, D., Ba, J. (2014). "Adam: a method for stochastic optimization". ICLR 2015. arXiv:1412.6980
  • Loshchilov, I., Hutter, F. (2019). "Decoupled weight decay regularization". ICLR. arXiv:1711.05101
  • Sutskever, Martens, Dahl, Hinton (2013). "On the importance of initialization and momentum in deep learning". ICML.
  • Goodfellow, Bengio, Courville (2016). Deep Learning, rozdz. 8.3 "Basic algorithms", 8.5 "Algorithms with adaptive learning rates". https://www.deeplearningbook.org/contents/optimization.html
  • Zhang i in. Dive into Deep Learning, rozdz. 12.6 "Momentum", 12.10 "Adam". https://d2l.ai/chapter_optimization/adam.html

Zobacz też