Gradient říká kterým směrem parametr posunout. Optimizer rozhoduje o tom jak daleko, a to je překvapivě celá věda. Špatně nastavený krok učení je nejčastější důvod, proč trénink buď nekonverguje, nebo vyloženě exploduje. Tady je, co dělá AdamW, proč se learning rate v čase mění a odkud se berou ty desítky gigabajtů navíc.
Od nejjednoduššího k tomu, co se používá
SGD: posuň parametr proti gradientu o pevný krok.
w ← w − lr × gradient
Funguje, ale je citlivý: pro každý parametr se hodí jiná velikost kroku a v úzkých údolích ztrátové krajiny to poskakuje ze strany na stranu.
Moment: přidej setrvačnost. Místo aktuálního gradientu se použije jeho klouzavý průměr, takže se šum vyruší a v konzistentním směru to naopak nabere rychlost.
Adam: ke setrvačnosti přidá adaptivní krok pro každý parametr zvlášť. Sleduje dvě věci:
m = klouzavý průměr gradientu („kterým směrem to dlouhodobě táhne“)
v = klouzavý průměr druhých mocnin („jak moc tenhle parametr kolísá“)
w ← w − lr × m / (√v + ε)
Parametr s velkými, rozkolísanými gradienty tak dostane menší krok, klidný parametr větší. Díky tomu Adam funguje „z krabice“ i na modely, kde by SGD vyžadoval týdny ladění.
AdamW: dnešní standard. Jediný rozdíl proti Adamu je v tom, že weight decay (mírné stahování vah k nule, které brání přeučení) se aplikuje odděleně od gradientu, ne skrz něj. Ukázalo se, že to funguje výrazně líp, a od té doby to používá prakticky každý.
Odtud těch 56 GB navíc
Adam si pro každý parametr drží m a v, obvykle ve fp32:
model 7 B:
m: 7 mld × 4 B = 28 GB
v: 7 mld × 4 B = 28 GB
= 56 GB jen na stavy optimizeru
Proto se u velkých modelů používají optimizery s menší paměťovou stopou (8bitový Adam, faktorizované varianty) nebo se stavy rozdělí mezi GPU přes ZeRO/FSDP, viz distribuovaný trénink.
A takhle vypadá learning rate v čase u skoro každého velkého tréninku:
lr
| ,-*-.
| ,' `-. cosine decay
| ,' `--.
| ,' `---.
| / `-----.____
|/ warmup `---
+-------------------------------------------> krok
0 2 % 100 %
Warmup existuje proto, že na začátku jsou váhy náhodné, gradienty obrovské a plný krok by model rozhodil. Decay proto, že ke konci chceš doladit, ne skákat.
Learning rate: nejdůležitější číslo tréninku
Nemění se v čase náhodou, má standardní tvar:
lr
│ ╭──────╮
│ ╱ ╲___
│ ╱ ╲____
│ ╱ ╲______
└──┴──────────────────────────────→ kroky
warmup cosine decay
Warmup (první procenta tréninku). Začne se u nuly a lineárně roste. Důvod: na začátku jsou
váhy náhodné, gradienty obrovské a plný krok by model rozhodil hned v prvních stovkách kroků.
Adam navíc potřebuje pár set kroků, aby měl smysluplné odhady m a v.
Cosine decay (zbytek). Krok postupně klesá skoro k nule. Na začátku chceš dělat velké skoky a najít správnou oblast, na konci jemně doladit. Model natrénovaný s doběhnutým decayem je znatelně lepší než ten samý model zastavený uprostřed.
⚠️ Proto se nedá jen tak „trénovat dál“. Když je rozvrh naplánovaný na milion kroků a ty ho po pěti stech tisících zastavíš, dostaneš model uprostřed sestupu a jeho dotrénování vyžaduje rozvrh znovu promyslet.
Další nastavení, která rozhodují
- Batch size. Kolik tokenů se zpracuje před jednou úpravou vah. Větší dávka = méně šumu v gradientu = dá se použít větší learning rate. Velké modely trénují na milionech tokenů v dávce.
- Gradient accumulation. Když se velká dávka nevejde do paměti, spočítá se po částech a gradienty se sčítají. Efekt je stejný jako u velké dávky, jen to trvá déle.
- Gradient clipping. Když norma gradientu překročí práh (typicky 1,0), zkrátí se. Je to pojistka proti tomu, aby jedna divná dávka rozhodila celý model.
- Weight decay. Typicky 0,1; drží váhy malé a zlepšuje zobecňování.
Jak vypadá, když se to pokazí
| Příznak | Nejčastější příčina |
|---|---|
| Loss stoupá už v prvních krocích | Chybí warmup nebo je lr příliš velký |
| Loss náhle vystřelí a nevrátí se | Loss spike; řeší se návratem k checkpointu a přeskočením dávky |
| Loss klesá a pak stagnuje vysoko | Příliš malý lr nebo příliš malá kapacita modelu |
| Trénovací loss klesá, validační roste | Přeučení, u fine-tuningu na malé sadě klasika |
Cvičení
- Trénuješ model 1,5 B v bf16 s AdamW. Spočítej paměť na váhy, gradienty a stavy optimizeru. Kolik z toho ušetří, když přejdeš na SGD s momentem?
- Learning rate zvýšíš desetkrát a loss začne po pár stech krocích růst. Co se děje a co uděláš dřív, než sáhneš po jiném optimizeru?
Náčrt řešení: rozbal, až si cvičení zkusíš sám
- Dohromady 18 GB, z toho 12 GB jsou stavy optimizeru. SGD s momentem by ušetřil 6 GB.
Váhy
1,5 × 2 = 3 GB, gradienty3 GB, AdamW drží dva momenty ve fp32, tedy1,5 × 8 = 12 GB. SGD s momentem drží jen jeden, tedy 6 GB. Jenže ta úspora se nevyplatí: bez adaptivního kroku pro každý parametr se velké modely trénují výrazně hůř, takže se šetří jinde, například ZeRO nebo gradient checkpointingem. - Krok je tak velký, že model přeskakuje minimum a rozhoupává se. První dvě věci k vyzkoušení jsou vrátit learning rate zpět a prodloužit warmup, teprve pak ořezávání normy gradientu. Optimizer je skoro nikdy ta správná první proměnná: AdamW s rozumným lr a plánovačem funguje napříč modely, zatímco špatně zvolený lr rozbije cokoli. Viz když se trénink pokazí.
Shrnutí
- AdamW = setrvačnost + adaptivní krok pro každý parametr + oddělený weight decay.
- Stavy optimizeru zaberou zhruba čtyřnásobek velikosti modelu v bf16.
- Learning rate má warmup (aby se to na začátku nerozsypalo) a decay (aby se to na konci doladilo).
- Gradient clipping a rozumná velikost dávky jsou hlavní pojistky proti divergenci.
K čemu slouží warmup na začátku tréninku?
Na začátku jsou váhy náhodné a gradienty velké, takže plný krok učení by model hned rozhodil. Learning rate proto postupně roste od nuly. Adam navíc potřebuje pár stovek kroků, než má použitelné odhady svých klouzavých průměrů.
Proč AdamW potřebuje zhruba čtyřnásobek velikosti modelu jen na své stavy?
Pro každý parametr drží dva klouzavé průměry (gradientu a jeho druhé mocniny), obvykle ve fp32. U modelu 7 B to je 2 × 7 mld × 4 B ≈ 56 GB, zatímco samotné váhy v bf16 zabírají 14 GB.
Můžeš vzít model, jehož trénink byl naplánovaný na milion kroků, zastavit ho v polovině a prostě pokračovat později?
Ne bez rozmyslu. Learning rate se řídí rozvrhem, který na konci klesá skoro k nule, model zastavený uprostřed je v jiném režimu než dotrénovaný. Pokračování vyžaduje rozvrh přeplánovat, jinak dostaneš horší výsledek než při doběhnutí původního plánu.
