Naar inhoud springen
6/30Hoofdstuk 6 van 30

Het netwerk laten trainen én laten generaliseren

Een netwerk met zes lagen waarvan de loss op ln 2 blijft hangen, stap voor stap gefixt. Daarna double descent: 5.000 parameters op 40 punten.

Op deze pagina

Het netwerk uit hoofdstuk 5 werkt. Het heeft negen parameters, het leert XOR en zijn gradients komen tot zestien decimalen overeen met PyTorch.

Maak het zes lagen diep en het stopt volledig met leren. Niet langzaam — volledig. Hier is een netwerk met zes lagen op een classificatieprobleem met twee spiralen, getraind voor 5000 stappen:

TEXT
step    1: loss 0.693147
step 5000: loss 0.693147
accuracy: 50.0 %

Dat getal is niet willekeurig. ln2=0.693147\ln 2 = 0.693147 is de binary cross-entropy van een model dat voor alles kans 0.50.5 uitvoert, en 50 % is een muntworp op een gebalanceerde dataset. Na vijfduizend stappen is het netwerk geen enkel cijfer opgeschoven. Niets crashte, niets gaf een waarschuwing, en de gradients zijn nog steeds exact goed.

Dit hoofdstuk gaat over de kloof tussen een netwerk dat draait en een netwerk dat werkt. Het heeft twee helften die op verschillende onderwerpen lijken, maar dezelfde taak zijn: de loss omlaag krijgen, en hem omlaag krijgen op data die het model nog nooit heeft gezien.

Begin met kijken, niet met raden. Duw een batch inputs erdoorheen en print de standaarddeviatie van de activaties bij elke laag, en daarna de standaarddeviatie van de weight gradients:

profile.pyPYTHON
def profile(model, x):
    h = x
    for layer in model:
        h = layer(h)
        if isinstance(layer, (nn.Tanh, nn.ReLU)):
            print(f"activation std: {h.std().item():.4f}")
    model(x).sum().backward()
    for p in model.parameters():
        if p.dim() == 2:
            print(f"gradient std: {p.grad.std().item():.2e}")

Drie initialisaties, dezelfde architectuur, zes lagen tanh\tanh:

initialisatieactivatie-std, lagen 1→6
normaal, std 0.010.010.0145 · 0.0016 · 0.0002 · 0.0000 · 0.0000 · 0.0000
normaal, std 110.6573 · 0.9296 · 0.9585 · 0.9634 · 0.9637 · 0.9625
Xavier0.1579 · 0.1493 · 0.1353 · 0.1333 · 0.1325 · 0.1403
initialisatiegradient-std, eerste laag → laatste
normaal, std 0.010.013.20e-06 · 4.97e-07 · … · 6.40e-06
normaal, std 111.94e+03 · 2.28e+02 · 1.22e+02 · 4.43e+01 · 1.85e+01 · 7.30e+00
Xavier2.31e+00 · 4.50e-01 · 4.26e-01 · 3.89e-01 · 4.39e-01 · 4.73e-01

De eerste rij is het netwerk hierboven, en het leert niet langzaam — er is geen signaal meer over. Bij laag vier is de standaarddeviatie van de activatie tot nul ondergelopen op vier decimalen. Elke input produceert dezelfde output, de output is een constante, en de gradient van een constante is niets. De weights waren klein geïnitialiseerd “voor de zekerheid”, en klein was fataal.

De tweede rij is de tegenovergestelde fout en die is het waard om te begrijpen, omdat hij tegenintuïtief is. De activaties zien er gezond uit — rond 0.96 — maar dat is tanh\tanh verzadigd, vastgepind dicht bij zijn limiet, precies het regime dat hoofdstuk 5 mat als een verlies van bijna een factor tienduizend in gradient. En toch zijn de gradients enorm: 1940 bij de eerste laag. Beide dingen zijn tegelijk waar. Elke backward-stap vermenigvuldigt met WW^\top, en met 128 inputs op unit variance heeft die factor een gain van ongeveer 12811\sqrt{128} \approx 11, wat de krimp door de verzadigde tanh\tanh overweldigt. De gradients groeien geometrisch op de weg terug. Dit is de exploding gradient, en die produceert loss-waarden van nan binnen een paar stappen in elke echte trainingsrun.

De derde rij is wat je wilt: activaties blijven ongeveer constant in schaal door de diepte heen, gradients blijven ongeveer constant in schaal door de diepte heen. Niets sterft, niets explodeert.

Goed initialiseren fixt de schaal op stap nul. Het houdt die schaal niet vast: de weights bewegen, en bij stap vijfduizend geldt het zorgvuldige variantie-argument niet meer.

Normalisatielagen dwingen de schaal continu af. Gegeven een vector activaties: trek een gemiddelde af, deel door een standaarddeviatie, en pas daarna een geleerde schaal γ\gamma en verschuiving β\beta toe, zodat de laag de normalisatie kan terugdraaien als dat blijkt te zijn wat hij wil:

h^=hμσ2+ϵ,y=γh^+β\hat{h} = \frac{h - \mu}{\sqrt{\sigma^2 + \epsilon}}, \qquad y = \gamma\hat{h} + \beta

De enige echte vraag is waarover je middelt. Batch normalisation3 neemt μ\mu en σ\sigma over de batch-dimensie, één statistiek per feature. Layer normalisation4 neemt ze over de features, één statistiek per voorbeeld.

Die keuze lijkt klein en bepaalt bijna alles stroomafwaarts:

BatchNorm laat de output van elk voorbeeld afhangen van de andere voorbeelden die toevallig in zijn batch zaten. Tijdens training is dat een milde regularizer. Tijdens inference is er geen batch, dus moet BatchNorm een lopend gemiddelde bijhouden van de statistieken die tijdens training zijn verzameld — wat betekent dat de laag zich anders gedraagt in trainings- en evaluatiemodus, en vergeten van modus te wisselen is een van de meest voorkomende bugs in het veld. Het degradeert ook bij kleine batches, en het is onhandig met sequenties van variabele lengte, omdat “het gemiddelde over de batch op positie 40” wordt berekend uit hoeveel sequenties toevallig zo lang zijn.

LayerNorm normaliseert elk voorbeeld op zichzelf. Geen batch-afhankelijkheid, geen lopende statistieken, identiek gedrag tijdens training en inference, ongevoelig voor batchgrootte, ongevoelig voor sequentielengte. Elk van die eigenschappen is eerder een vereiste dan een aardigheid zodra je één token tegelijk genereert voor één gebruiker, en dat is waar hoofdstuk 13 uitkomt.

Daarom is LayerNorm degene die je in hoofdstuk 9 ongewijzigd opnieuw tegenkomt: het transformer-blok gebruikt hem, en gebruikt hem om de redenen in de rechterkolom, niet omdat hij in abstracte zin beter werkt.

Eén ding tegelijk fixen, wat de echte skill is

Link naar de sectie: Eén ding tegelijk fixen, wat de echte skill is

Vier kandidaat-fixes voor het dode netwerk: Xavier-initialisatie, LayerNorm, residual connections, en Adam in plaats van SGD. De verleiding is om ze alle vier toe te passen en door te gaan. Doe dat en je zult nooit weten welke ertoe deed, en de volgende keer dat het gebeurt heb je geen methode — alleen een ritueel.

Pas ze dus één voor één toe. Dezelfde seed, dezelfde data, dezelfde architectuur, 800 stappen:

wat is toegevoegduiteindelijke lossaccuracy
niets0.693150.0 %
Xavier-initialisatie0.569260.4 %
LayerNorm0.623061.5 %
residual connections0.665156.6 %
Adam0.678758.7 %
alle vier0.0000100.0 %

Lees die tabel zoals je hem om 2 uur ’s nachts zou lezen en de conclusie is: niets werkt alleen, alles werkt samen, dus deep learning is alchemie. Die conclusie is fout, en ontdekken waarom is het nuttigste in dit hoofdstuk.

Geef elke run zes keer zoveel budget — 5000 stappen in plaats van 800 — en het verandert volledig:

wat is toegevoegduiteindelijke loss @ 5000accuracy
niets0.693150.0 %
Xavier-initialisatie0.0007100.0 %
LayerNorm0.0002100.0 %
residual connections0.665356.7 %
Adam0.690853.4 %
Xavier + Adam0.0000100.0 %
Xavier + LayerNorm0.0001100.0 %

Nu is het beeld scherp, en het is een diagnose in plaats van een ritueel.

Initialisatie alleen fixt het. Normalisatie alleen fixt het. Elk pakt de echte ziekte aan — het forward-signaal dat naar nul instort — en elk van beide is voldoende. Bij 800 stappen leken ze alleen maar gedeeltelijke punten te krijgen, omdat ze het probleem hadden opgelost en nog bezig waren eruit te klimmen.

Residual connections en Adam fixen het niet, bij geen enkel budget. Niet omdat ze slecht zijn, maar omdat ze een andere ziekte behandelen. Een residual connection geeft de gradient een pad om een blokkerende laag heen; dat is veel waard wanneer de gradient het probleem is, en niets waard wanneer het forward-signaal al nul is, omdat een snelweg om een dode laag heen nog steeds een dode waarde draagt. Adam herschaalt de stap van elke parameter op basis van zijn eigen gradient-geschiedenis; dat helpt wanneer gradients enorm verschillende groottes hebben, en kan geen netwerk tot leven wekken waarvan de output niet afhangt van de input.

En “niets” is na vijfduizend stappen nog steeds exact 0.6931. Niet 0.6929. Het is niet traag; het is dood, en dat onderscheid is nu zichtbaar op een manier waarop het eerder niet zichtbaar was, omdat je de rij hebt die zegt dat een fix werkt om mee te vergelijken.

Vanaf hier gebruikt deze cursus PyTorch. Dat moet verdiend worden in plaats van aangekondigd, dus hier is precies wat het doet dat je zelf al weet te doen.

Een optimizer is een regel om gradients om te zetten in parameterupdates. Gewone gradient descent gebruikt de gradient. Momentum gebruikt er een lopend gemiddelde van, wat de ruis gladstrijkt en snelheid opbouwt in richtingen die consistent blijven:

optim_by_hand.pyPYTHON
v = beta * v + p.grad          
p -= lr * v                    

Adam5 houdt twee lopende gemiddelden bij — van de gradient en van de gradient in het kwadraat — en deelt het ene door de wortel van het andere, zodat elke parameter een stap krijgt die is geschaald naar zijn eigen recente gradient-grootte:

optim_by_hand.pyPYTHON
m = b1 * m + (1 - b1) * g          # mean of the gradient          
v = b2 * v + (1 - b2) * g * g      # mean of the squared gradient  
m_hat = m / (1 - b1 ** t)          # bias correction: both averages start at zero
v_hat = v / (1 - b2 ** t)
p -= lr * m_hat / (v_hat.sqrt() + eps)   

Tien regels. Draai beide tegen torch.optim op hetzelfde probleem voor 50 stappen:

TEXT
SGD+momentum   by hand [2.7781870365142822, -1.0304985046386719]
               torch   [2.7781870365142822, -1.0304983854293823]   max |diff| = 1.19e-07
Adam           by hand [0.4893140196800232, -0.46317872405052185]
               torch   [0.48931416869163513, -0.46317875385284424]   max |diff| = 1.49e-07

Identiek tot float32-precisie. torch.optim.Adam is die vijf regels, plus decennia aan zorg voor randgevallen en een C++-kernel. Dat is de ruil die je vanaf nu maakt: geen magie voor begrip, maar snelheid voor regels die je al hebt geschreven.

De gebruikelijke uitleg van Adam is “adaptieve learning rates per parameter”, wat een beschrijving is en geen reden. De reden is geometrie, en die kun je meten.

Neem een loss waarvan de kromming per richting verschilt: steil in de ene, ondiep in de andere. SGD heeft één globale learning rate, dus het moet een waarde kiezen die klein genoeg is om stabiel te blijven in de steilste richting — en die waarde is vervolgens veel te klein voor de ondiepe richting, waar vooruitgang kruipt. Dit veroorzaakt het klassieke beeld van gradient descent die zigzaggend door een smalle vallei omlaaggaat.

Twee krommingsverhoudingen, drie optimizers, 300 stappen, en elke optimizer krijgt de beste learning rate uit een sweep zodat niemand wordt benadeeld:

krommingsverhoudingSGDSGD + momentumAdam
10 : 1fout 0.000002fout 0.000000fout 0.000000
1000 : 1fout 1.925485fout 0.001432fout 0.000000
gedivergeerd bij (1000:1)4 van 8 rates4 van 8 rates0 van 6 rates

Bij een verhouding van tien werkt alles en valt er niets te bespreken. Bij duizend kan gewone SGD het antwoord bij geen enkele geprobeerde learning rate bereiken — zijn beste resultaat is nog steeds een fout van 1.93 — en het divergeert ronduit bij de helft van de rates. Adam landt exact op het target en divergeert bij geen enkele.

Die laatste kolom is de praktische reden waarom Adam de default is. Het is niet dat Adam betere oplossingen vindt; op goed geconditioneerde problemen evenaart of verslaat getunede SGD hem vaak. Het is dat Adam veel minder gevoelig is voor de learning rate die je hebt gekozen, en echte netwerken hebben krommingsverhoudingen die veel erger zijn dan duizend over hun miljoenen parameters.

Twee extra stukken horen hierbij en beide zijn één regel. Gradient clipping herschaalt de gradient-vector wanneer zijn norm een drempel overschrijdt, waardoor de rij “loss springt plotseling naar een enorme waarde” uit de diagnostische tabel een non-event wordt. En learning rate schedules: een korte warmup vanaf bijna nul over de eerste paar honderd stappen, omdat Adams variantieschattingen waardeloos zijn totdat ze wat gradients hebben gezien en een stap op volle grootte op waardeloze schattingen een initialisatie kan slopen; daarna cosine decay richting nul, omdat een run eindigen met dezelfde stapgrootte waarmee je begon betekent dat je rond het minimum blijft trillen in plaats van erin te landen.

De tweede helft: het model dat perfect fit en niets voorspelt

Link naar de sectie: De tweede helft: het model dat perfect fit en niets voorspelt

Alles tot nu toe ging over de loss omlaag krijgen. Nu de moeilijkere helft, omdat de loss omlaag krijgen niet het doel is — het is een proxy voor het doel, en die proxy faalt op een specifieke en beroemde manier.

Twaalf punten uit een gladde functie met een beetje ruis. Fit polynomen van oplopende graad:

graadtrain RMSEtest RMSE
10.7644990.6985
30.2526050.3031
50.1644370.1568
90.0889600.2347
110.0000001.2094

Graad 11 door 12 punten gaat door elk afzonderlijk punt exact heen — train error nul tot zes decimalen — en is acht keer slechter dan graad 5 op data die het niet heeft gezien. Vraag graad 3 en graad 11 om te voorspellen bij x=3.25x = 3.25, net buiten het trainingsbereik:

TEXT
degree  3: predicts   -1.053   (truth -0.012)
degree 11: predicts  +61.224   (truth -0.012)

Eenenzestig, waar het antwoord ongeveer nul is. Het model heeft de functie niet geleerd; het heeft de twaalf punten geleerd, en ertussen doet het wat de rekenkunde eist.

Dit is overfitting, en het tegenovergestelde — graad 1, die de curve helemaal niet kan representeren en overal slecht is — is underfitting. De klassieke uitleg splitst de verwachte fout van een model in drie delen: bias, de fout doordat het model te rigide is om de waarheid te representeren; variance, de fout doordat het model zo flexibel is dat het de ruis in deze specifieke sample najaagt; en irreducible noise, die door niets wordt gefixt. Simpele modellen hebben bias, flexibele modellen hebben hoge variance, en het klassieke voorschrift is om de sweet spot in het midden te vinden — graad 5 in de tabel hierboven.

De standaardtools vallen allemaal de variance-term aan:

  • L2-regularisatie (weight decay) voegt λw2\lambda \lVert w \rVert^2 toe aan de loss, trekt weights richting nul en maakt de functie gladder. In de tabel hierboven doet de grootste coëfficiënt van graad 11 de schade; grootte bestraffen maakt hem onschadelijk.
  • L1 voegt in plaats daarvan λwi\lambda \sum |w_i| toe. Het verschil is niet cosmetisch: de gradient van L2 is proportioneel aan de weight en krimpt dus mee met de weight, nadert nul zonder er te komen, terwijl de gradient van L1 een constante ±λ\pm\lambda is die helemaal blijft duwen. L1 produceert daarom weights die exact nul zijn — het selecteert features. L2 produceert kleine weights. Gebruik L2 wanneer je gladheid wilt, L1 wanneer je sparsity wilt.
  • Dropout7 zet bij elke trainingsstap een willekeurige subset van activaties op nul, zodat geen enkele unit erop kan vertrouwen dat een bepaalde andere unit aanwezig is.
  • Early stopping kijkt naar de validation loss en stopt wanneer die omhoog draait.
  • Data augmentation maakt meer trainingsvoorbeelden uit de voorbeelden die je hebt, en pakt daarmee het probleem bij de bron aan: overfitting is net zo goed een tekort aan data als een teveel aan parameters.
  • Cross-validation splitst de data kk manieren en traint kk keer, wat een betrouwbare schatting van de test error oplevert wanneer je te weinig data hebt om een held-out set apart te houden.

Double descent, of waarom de vorige sectie niet het hele verhaal is

Link naar de sectie: Double descent, of waarom de vorige sectie niet het hele verhaal is

Nu het feit dat het beeld breekt.

Het bias-variance-verhaal zegt dat voorbij de sweet spot meer parameters slechtere generalisatie betekenen. Moderne taalmodellen hebben veel meer parameters dan de klassieke regels toestaan voor de data die ze zien, en generaliseren uitstekend. Beide uitspraken zijn waar, en ze met elkaar verzoenen is het nuttigste in dit hoofdstuk.

Veertig trainingspunten, twintigdimensionale inputs, willekeurige ReLU-features, en het aantal features PP gesweept van 2 tot 5000 — met de minimum-norm-oplossing gekozen wanneer er veel oplossingen zijn die fitten:

PPP/nP/ntrain RMSEtest RMSEw\lVert w \rVert
100.250.88221.25201.89
200.500.59621.16342.59
300.750.38961.53234.15
380.950.17693.716310.25
401.000.00005.814014.83
421.050.00003.16239.35
601.500.00001.10582.78
2005.000.00000.66380.98
150037.500.00000.58590.33
5000125.000.00000.56640.18

Lees het in drie delen. Tot P/n=0.5P/n = 0.5 klopt het klassieke verhaal precies: de fout daalt en begint dan te stijgen. Bij P=n=40P = n = 40 — de interpolatiedrempel, waar het model exact genoeg parameters heeft om door elk trainingspunt heen te gaan — piekt de test error op 5.81, vijf keer slechter dan het kleine model. Die piek is de klassieke waarschuwing, en die is echt.

Daarna daalt hij opnieuw. En hij blijft dalen, voorbij P=5nP = 5n, voorbij P=37nP = 37n, helemaal tot P=125nP = 125n, waar de test error van 0.5664 beter is dan het beste ondergeparametriseerde model ooit haalde. Een model met 5000 parameters gefit op 40 punten is het beste model in de tabel.

Dit is double descent,89 en het mechanisme is zichtbaar in de laatste kolom. Zodra P>nP > n zijn er oneindig veel parameterinstellingen die de trainingsdata exact fitten, en welke je krijgt hangt af van hoe je kiest. De minimum-norm-oplossing kiest de kleinste, en w\lVert w \rVert laat zien wat dat betekent: hij piekt op 14.83 precies bij de drempel — waar er exact één interpolerende oplossing is en je eraan vastzit, hoe extreem ook — en daalt daarna monotoon terwijl PP groeit, omdat meer parameters meer interpolerende oplossingen betekenen om uit te kiezen, wat betekent dat de kleinste beschikbare kleiner wordt. Bij P=5000P = 5000 is de norm 0.18, tachtig keer kleiner dan bij de drempel.

De extra parameters voegen dus geen complexiteit toe. Ze voegen keuze toe, en de selectieregel besteedt die keuze aan eenvoud. De regularisatie zit niet in de loss function; die zit in het algoritme. Gradient descent vanuit een kleine initialisatie heeft een gedocumenteerde bias richting small-norm-oplossingen, en daarom verschijnt dit gedrag in echte netwerken die op de gewone manier worden getraind, niet alleen in de lineaire algebra hierboven.

Het praktische gevolg, waar hoofdstuk 10 van afhangt: “het model heeft meer parameters dan data, dus het zal overfitten” is geen geldig argument. Het was een goede regel toen modellen links van de drempel leefden. Alles wat nu interessant is leeft er ver rechts van, waar de regel omkeert.

De tools in dit hoofdstuk zijn genoeg om een netwerk te trainen dat werkt op data die je in een tabel kunt zetten: rijen getallen, een kolom labels.

Taal is dat niet. Voordat een model het volgende woord kan voorspellen, moet iets bepalen wat een “woord” überhaupt is — en het antwoord is noch letters noch woorden, maar een vocabulaire dat het model leert uit de ruwe bytes van de trainingsdata. Die beslissing, eenmaal genomen voordat training begint, bepaalt hoeveel dingen het model kan zeggen, hoeveel een request kost, en waarom modellen die een rechtenexamen kunnen halen niet betrouwbaar de letters in strawberry kunnen tellen.

Hoofdstuk 7 bouwt een tokenizer.


Voor de residual connections hierboven: He et al., Deep Residual Learning for Image Recognition (arXiv:1512.03385). Andrej Karpathy’s Building makemore Part 3: Activations & Gradients, BatchNorm loopt door de activatie-histogramdiagnose op een echt model en is de beste hands-on behandeling van de eerste helft van dit hoofdstuk. Yaser Abu-Mostafa’s colleges Learning From Data, 8 en 11–13, geven de klassieke generalisatietheorie zoals het hoort, inclusief de delen die dit hoofdstuk tot één alinea heeft samengeperst.

  1. Glorot, X. en Bengio, Y. Understanding the difficulty of training deep feedforward neural networks. AISTATS (2010). Het variantiebehoud-argument dat in het kader hierboven is gereproduceerd.

  2. He, K., Zhang, X., Ren, S. en Sun, J. Delving Deep into Rectifiers: Surpassing Human-Level Performance on ImageNet Classification. arXiv:1502.01852 (2015).

  3. Ioffe, S. en Szegedy, C. Batch Normalization: Accelerating Deep Network Training by Reducing Internal Covariate Shift. arXiv:1502.03167 (2015). Merk op dat de verklaring “internal covariate shift” in de titel sindsdien stevig is betwist; de laag werkt, maar de oorspronkelijke uitleg waarom is omstreden.

  4. Ba, J. L., Kiros, J. R. en Hinton, G. E. Layer Normalization. arXiv:1607.06450 (2016).

  5. Kingma, D. P. en Ba, J. Adam: A Method for Stochastic Optimization. arXiv:1412.6980 (2014).

  6. Loshchilov, I. en Hutter, F. Decoupled Weight Decay Regularization. arXiv:1711.05101 (2017).

  7. Srivastava, N., Hinton, G., Krizhevsky, A., Sutskever, I. en Salakhutdinov, R. Dropout: A Simple Way to Prevent Neural Networks from Overfitting. JMLR 15, pp. 1929–1958 (2014).

  8. Belkin, M., Hsu, D., Ma, S. en Mandal, S. Reconciling modern machine-learning practice and the classical bias–variance trade-off. PNAS 116(32), pp. 15849–15854 (2019). Het paper dat het fenomeen een naam gaf.

  9. Nakkiran, P., Kaplun, G., Bansal, Y., Yang, T., Barak, B. en Sutskever, I. Deep Double Descent: Where Bigger Models and More Data Hurt. arXiv:1912.02292 (2019). Toont het effect in echte diepe netwerken, en langs de as van training time én de as van modelgrootte.

Klaar om LIA te laten kiezen?

Bouw met elk AI-model op één plek — begin vandaag nog gratis.