Ves al contingut
6/30Capítol 6 de 30

Fer que entreni i que generalitzi

Una xarxa de sis capes amb pèrdua clavada a ln 2, arreglada mesura a mesura. Després, doble descens: 5.000 paràmetres en 40 punts.

En aquesta pàgina

La xarxa del capítol 5 funciona. Té nou paràmetres, aprèn XOR i els seus gradients coincideixen amb PyTorch fins a setze decimals.

Fes-la de sis capes de profunditat i deixa d’aprendre del tot. No lentament: del tot. Aquí tens una xarxa de sis capes en un problema de classificació de dues espirals, entrenada durant 5000 passos:

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

Aquest nombre no és arbitrari. ln2=0.693147\ln 2 = 0.693147 és l’entropia creuada binària d’un model que dona una probabilitat de 0.50.5 per a tot, i el 50 % és llançar una moneda en un dataset equilibrat. Després de cinc mil passos, la xarxa no ha mogut ni un sol dígit. Res no ha fallat, no ha aparegut cap avís, i els gradients continuen sent exactament correctes.

Aquest capítol tracta de l’escletxa entre una xarxa que s’executa i una xarxa que funciona. Té dues meitats que semblen temes diferents però són la mateixa feina: aconseguir que la pèrdua vagi avall, i aconseguir que baixi en dades que el model no ha vist mai.

Comença mirant, en lloc d’endevinar. Fes passar un batch d’entrades i imprimeix la desviació estàndard de les activacions a cada capa, i després la desviació estàndard dels gradients dels pesos:

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}")

Tres inicialitzacions, mateixa arquitectura, sis capes de tanh\tanh:

inicialitzaciódesviació estàndard de l’activació, capes 1→6
normal, std 0.010.010.0145 · 0.0016 · 0.0002 · 0.0000 · 0.0000 · 0.0000
normal, 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
inicialitzaciódesviació estàndard del gradient, primera capa → última
normal, std 0.010.013.20e-06 · 4.97e-07 · … · 6.40e-06
normal, 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

La primera fila és la xarxa d’abans, i no està aprenent lentament: ja no li queda cap senyal. A la quarta capa, la desviació estàndard de l’activació ha fet underflow fins a zero amb quatre decimals. Cada entrada produeix la mateixa sortida, la sortida és una constant, i el gradient d’una constant no és res. Els pesos es van inicialitzar petits «per anar sobre segur», i petit va ser fatal.

La segona fila és el fracàs oposat i val la pena entendre’l perquè és contraintuïtiu. Les activacions semblen sanes —al voltant de 0.96—, però això és tanh\tanh saturat, fixat prop del seu límit, exactament el règim que el capítol 5 va mesurar com una pèrdua d’un factor de gairebé deu mil en el gradient. I, tanmateix, els gradients són enormes: 1940 a la primera capa. Totes dues coses són certes alhora. Cada pas enrere multiplica per WW^\top, i amb 128 entrades a variància unitària aquest factor té un guany d’aproximadament 12811\sqrt{128} \approx 11, que ofega la reducció del tanh\tanh saturat. Els gradients creixen geomètricament en tornar enrere. Això és l’exploding gradient, i produeix valors de pèrdua de nan en pocs passos en qualsevol entrenament real.

La tercera fila és el que vols: activacions amb una escala aproximadament constant al llarg de la profunditat, gradients amb una escala aproximadament constant al llarg de la profunditat. Res no mor, res no explota.

Inicialitzar bé arregla l’escala al pas zero. No la manté fixa: els pesos es mouen, i al pas cinc mil l’argument acurat de la variància ja no s’aplica.

Les capes de normalització imposen l’escala contínuament. Donat un vector d’activacions, resta una mitjana, divideix per una desviació estàndard, i després aplica una escala apresa γ\gamma i un desplaçament β\beta perquè la capa pugui desfer la normalització si resulta que això és el que vol:

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

L’única pregunta real és sobre què fas la mitjana. Batch normalisation3 pren μ\mu i σ\sigma al llarg de la dimensió del batch, una estadística per feature. Layer normalisation4 les pren al llarg de les features, una estadística per exemple.

Aquesta tria sembla menor i decideix gairebé tot el que ve després:

BatchNorm fa que la sortida de cada exemple depengui dels altres exemples que han coincidit al seu batch. Durant l’entrenament això és una regularització suau. En inferència no hi ha batch, així que ha de mantenir una mitjana mòbil de les estadístiques recollides durant l’entrenament; això vol dir que la capa es comporta diferent en mode entrenament i en mode avaluació, i oblidar-se de canviar de mode és un dels bugs més habituals del camp. També es degrada amb batchs petits, i és incòmode amb seqüències de longitud variable, perquè «la mitjana del batch a la posició 40» es calcula amb les seqüències que, per casualitat, són prou llargues.

LayerNorm normalitza cada exemple per si sol. Sense dependència del batch, sense estadístiques mòbils, comportament idèntic en entrenament i inferència, indiferent a la mida del batch, indiferent a la longitud de la seqüència. Cadascuna d’aquestes propietats és un requisit més que no pas un detall agradable quan generes un token cada vegada per a un usuari, que és on acaba el capítol 13.

Per això LayerNorm és la que tornaràs a veure al capítol 9 sense canvis: el bloc transformer la fa servir, i la fa servir per les raons de la columna de la dreta, no perquè funcioni millor en abstracte.

Arreglar una cosa cada vegada, que és el veritable skill

Enllaç a la secció: Arreglar una cosa cada vegada, que és el veritable skill

Quatre possibles arreglos per a la xarxa morta: inicialització Xavier, LayerNorm, connexions residuals i Adam en lloc de SGD. La temptació és aplicar-los tots quatre i continuar. Fes-ho i mai sabràs quin importava, i la pròxima vegada que passi no tindràs cap mètode: només un ritual.

Així que aplica’ls d’un en un. Mateixa seed, mateixes dades, mateixa arquitectura, 800 passos:

què s’ha afegitpèrdua finalexactitud
res0.693150.0 %
inicialització Xavier0.569260.4 %
LayerNorm0.623061.5 %
connexions residuals0.665156.6 %
Adam0.678758.7 %
tots quatre0.0000100.0 %

Llegeix aquesta taula com la llegiries a les 2 de la matinada i la conclusió és: res no funciona sol, tot funciona conjuntament, per tant el deep learning és alquímia. Aquesta conclusió és equivocada, i descobrir per què és el més útil d’aquest capítol.

Dona a cada execució sis vegades més pressupost —5000 passos en lloc de 800— i canvia completament:

què s’ha afegitpèrdua final @ 5000exactitud
res0.693150.0 %
inicialització Xavier0.0007100.0 %
LayerNorm0.0002100.0 %
connexions residuals0.665356.7 %
Adam0.690853.4 %
Xavier + Adam0.0000100.0 %
Xavier + LayerNorm0.0001100.0 %

Ara la imatge és nítida, i és un diagnòstic més que no pas un ritual.

La inicialització sola ho arregla. La normalització sola ho arregla. Cadascuna aborda la malaltia real —el senyal cap endavant que col·lapsa fins a zero— i qualsevol de les dues és suficient. Als 800 passos només semblaven mèrit parcial, perquè havien resolt el problema i encara estaven sortint del sot.

Les connexions residuals i Adam no ho arreglen, amb cap pressupost. No perquè siguin dolents, sinó perquè tracten una altra malaltia. Una connexió residual dona al gradient un camí al voltant d’una capa bloquejant; això val molt quan el problema és el gradient, i no val res quan el senyal cap endavant ja és zero, perquè una drecera al voltant d’una capa morta continua portant un valor mort. Adam reescala el pas de cada paràmetre segons el seu propi historial de gradients; això ajuda quan els gradients tenen magnituds molt diferents, i no pot ressuscitar una xarxa la sortida de la qual no depèn de l’entrada.

I «res» continua sent exactament 0.6931 després de cinc mil passos. No 0.6929. No és lenta; és morta, i aquesta distinció ara és visible d’una manera que abans no ho era, perquè tens la fila que diu que un arreglament funciona per comparar-hi.

A partir d’aquí aquest curs fa servir PyTorch. Això s’ha de guanyar més que anunciar, així que aquí tens exactament què fa que tu ja saps fer.

Un optimitzador és una regla per convertir gradients en actualitzacions de paràmetres. El gradient descent simple fa servir el gradient. Momentum en fa servir una mitjana mòbil, cosa que suavitza el soroll i acumula velocitat en direccions que es mantenen consistents:

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

Adam5 manté dues mitjanes mòbils —del gradient i del gradient al quadrat— i divideix una per l’arrel quadrada de l’altra, de manera que cada paràmetre rep un pas escalat segons la seva magnitud recent de gradient:

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)   

Deu línies. Executa tots dos contra torch.optim en el mateix problema durant 50 passos:

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

Idèntic fins a la precisió float32. torch.optim.Adam són aquestes cinc línies, més dècades de cura amb casos límit i un kernel C++. Aquest és l’intercanvi que fas a partir d’ara: no màgia a canvi d’entendre, sinó velocitat a canvi de línies que ja has escrit.

L’explicació habitual d’Adam és «learning rates adaptatius per paràmetre», que és una descripció més que una raó. La raó és geometria, i es pot mesurar.

Agafa una pèrdua amb curvatura diferent segons la direcció: pronunciada en una, suau en una altra. SGD té un únic learning rate global, així que ha de triar un valor prou petit per ser estable en la direcció més pronunciada; i aquest valor és llavors molt massa petit per a la suau, on el progrés s’arrossega. Això és el que causa la imatge clàssica del gradient descent fent ziga-zaga avall per una vall estreta.

Dues ràtios de curvatura, tres optimitzadors, 300 passos, i cada optimitzador amb el millor learning rate d’un sweep perquè ningú tingui desavantatge:

ràtio de curvaturaSGDSGD + momentumAdam
10 : 1error 0.000002error 0.000000error 0.000000
1000 : 1error 1.925485error 0.001432error 0.000000
divergit a (1000:1)4 de 8 rates4 de 8 rates0 de 6 rates

Amb una ràtio de deu, tot funciona i no hi ha res a discutir. Amb una de mil, l’SGD simple no pot arribar a la resposta amb cap learning rate provat —el seu millor resultat continua sent un error d’1.93— i divergeix directament amb la meitat dels rates. Adam aterra exactament al target i no divergeix amb cap.

Aquesta última columna és la raó pràctica per la qual Adam és el valor per defecte. No és que Adam trobi solucions millors; en problemes ben condicionats, SGD ajustat sovint l’iguala o el supera. És que Adam és molt menys sensible al learning rate que has triat, i les xarxes reals tenen ràtios de curvatura molt pitjors que mil al llarg dels seus milions de paràmetres.

Aquí hi van dues peces més, i totes dues són d’una línia. Gradient clipping reescala el vector de gradient sempre que la seva norma supera un llindar, cosa que converteix la fila «la pèrdua salta de sobte a un valor enorme» de la taula de diagnòstic en un no-esdeveniment. I learning rate schedules: un warmup curt des de gairebé zero durant els primers centenars de passos, perquè les estimacions de variància d’Adam són escombraries fins que han vist alguns gradients i un pas de mida completa pres sobre escombraries pot destrossar una inicialització; després, cosine decay cap a zero, perquè acabar una execució amb la mateixa mida de pas amb què has començat vol dir tremolar al voltant del mínim en lloc d’assentar-s’hi.

La segona meitat: el model que encaixa perfectament i no prediu res

Enllaç a la secció: La segona meitat: el model que encaixa perfectament i no prediu res

Tot fins ara anava d’aconseguir que la pèrdua baixés. Ara ve la meitat més difícil, perquè que la pèrdua baixi no és l’objectiu: és un proxy de l’objectiu, i el proxy falla d’una manera específica i famosa.

Dotze punts d’una funció suau amb una mica de soroll. Ajusta polinomis de grau creixent:

grauRMSE entrenamentRMSE prova
10.7644990.6985
30.2526050.3031
50.1644370.1568
90.0889600.2347
110.0000001.2094

El grau 11 a través de 12 punts passa per tots i cadascun exactament —error d’entrenament zero fins a sis decimals— i és vuit vegades pitjor que el grau 5 en dades que no ha vist. Demana al grau 3 i al grau 11 que prediguin a x=3.25x = 3.25, just fora del rang d’entrenament:

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

Seixanta-u, quan la resposta és aproximadament zero. El model no ha après la funció; ha après els dotze punts, i entre ells fa el que l’aritmètica exigeixi.

Això és overfitting, i el seu contrari —grau 1, que no pot representar la corba de cap manera i és dolent a tot arreu— és underfitting. L’explicació clàssica divideix l’error esperat d’un model en tres parts: biaix, l’error perquè el model és massa rígid per representar la veritat; variància, l’error perquè el model és tan flexible que persegueix el soroll d’aquesta mostra concreta; i soroll irreductible, que res no arregla. Els models simples són esbiaixats, els models flexibles tenen variància alta, i la prescripció clàssica és trobar el punt dolç al mig: el grau 5 de la taula anterior.

Les eines estàndard ataquen totes el terme de variància:

  • Regularització L2 (weight decay) afegeix λw2\lambda \lVert w \rVert^2 a la pèrdua, estirant els pesos cap a zero i fent la funció més suau. A la taula anterior, el coeficient més gran del grau 11 fa el mal; penalitzar la mida el desactiva.
  • L1 afegeix λwi\lambda \sum |w_i| en lloc d’això. La diferència no és cosmètica: el gradient de L2 és proporcional al pes i, per tant, es redueix a mesura que ho fa el pes, acostant-se a zero sense arribar-hi, mentre que el gradient de L1 és una constant ±λ\pm\lambda que continua empenyent fins al final. Per això L1 produeix pesos que són exactament zero: selecciona features. L2 produeix pesos petits. Fes servir L2 quan vulguis suavitat, L1 quan vulguis esparsitat.
  • Dropout7 posa a zero un subconjunt aleatori d’activacions a cada pas d’entrenament, de manera que cap unitat pot dependre que cap altra unitat concreta hi sigui present.
  • Early stopping observa la pèrdua de validació i s’atura quan gira cap amunt.
  • Data augmentation fabrica més exemples d’entrenament a partir dels que tens, cosa que ataca el problema a l’arrel: l’overfitting és tant una manca de dades com un excés de paràmetres.
  • Cross-validation divideix les dades de kk maneres i entrena kk vegades, cosa que compra una estimació fiable de l’error de prova quan tens massa poques dades per reservar un conjunt separat.

Doble descens, o per què la secció anterior no és tota la història

Enllaç a la secció: Doble descens, o per què la secció anterior no és tota la història

Ara ve el fet que trenca la imatge.

La història biaix-variància diu que, passat el punt dolç, més paràmetres volen dir pitjor generalització. Els models de llenguatge moderns tenen molts més paràmetres dels que les regles clàssiques permetrien per a les dades que veuen, i generalitzen magníficament. Totes dues afirmacions són certes, i reconciliar-les és el més útil d’aquest capítol.

Quaranta punts d’entrenament, entrades de vint dimensions, features ReLU aleatòries, i el nombre de features PP recorregut de 2 a 5000, amb la solució de norma mínima triada sempre que n’hi ha moltes que encaixen:

PPP/nP/nRMSE entrenamentRMSE provaw\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

Llegeix-ho en tres parts. Fins a P/n=0.5P/n = 0.5, la història clàssica es compleix exactament: l’error cau i després comença a pujar. A P=n=40P = n = 40 —el llindar d’interpolació, on el model té exactament prou paràmetres per passar per cada punt d’entrenament— l’error de prova arriba al pic, a 5.81, cinc vegades pitjor que el model petit. Aquest pic és l’advertiment clàssic, i és real.

Després torna a baixar. I continua baixant, més enllà de P=5nP = 5n, més enllà de P=37nP = 37n, fins a P=125nP = 125n, on l’error de prova de 0.5664 és millor que el millor model subparametritzat que s’hagi aconseguit. Un model amb 5000 paràmetres ajustat a 40 punts és el millor model de la taula.

Això és doble descens,89 i el mecanisme és visible a l’última columna. Un cop P>nP > n hi ha infinites configuracions de paràmetres que encaixen exactament amb les dades d’entrenament, i quina obtens depèn de com tries. La solució de norma mínima tria la més petita, i w\lVert w \rVert mostra què vol dir això: arriba al pic de 14.83 just al llindar —on hi ha exactament una solució interpoladora i t’hi quedes atrapat, per extrema que sigui— i després cau monòtonament a mesura que PP creix, perquè més paràmetres vol dir més solucions interpoladores per triar, cosa que vol dir que la més petita disponible es fa més petita. A P=5000P = 5000 la norma és 0.18, vuitanta vegades més petita que al llindar.

Així que els paràmetres extra no afegeixen complexitat. Afegeixen tria, i la regla de selecció gasta aquesta tria en simplicitat. La regularització no és a la funció de pèrdua; és a l’algoritme. Gradient descent des d’una inicialització petita té un biaix documentat cap a solucions de norma petita, i per això aquest comportament apareix en xarxes reals entrenades de la manera ordinària i no només en l’àlgebra lineal anterior.

La conseqüència pràctica, de la qual depèn el capítol 10: «el model té més paràmetres que dades, així que farà overfitting» no és un argument vàlid. Era una bona regla quan els models vivien a l’esquerra del llindar. Ara tot el que és interessant viu molt a la dreta, on la regla s’inverteix.

Les eines d’aquest capítol són suficients per entrenar una xarxa que funcioni amb dades que pots posar en una taula: files de nombres, una columna d’etiquetes.

El llenguatge no és això. Abans que un model pugui predir la paraula següent, alguna cosa ha de decidir què és exactament una «paraula», i la resposta no són ni lletres ni paraules, sinó un vocabulari que el model aprèn dels bytes crus de les dades d’entrenament. Aquesta decisió, presa una vegada abans que comenci l’entrenament, determina quantes coses pot dir el model, quant costa una petició, i per què models que poden aprovar un examen de dret no poden comptar de manera fiable les lletres de strawberry.

El capítol 7 construeix un tokenizer.


Per a les connexions residuals utilitzades més amunt, He et al., Deep Residual Learning for Image Recognition (arXiv:1512.03385). Building makemore Part 3: Activations & Gradients, BatchNorm, d’Andrej Karpathy, recorre el diagnòstic amb histogrames d’activació en un model real i és el millor tractament pràctic de la primera meitat d’aquest capítol. Les lliçons 8 i 11–13 de Learning From Data, de Yaser Abu-Mostafa, expliquen correctament la teoria clàssica de la generalització, incloses les parts que aquest capítol ha comprimit en un paràgraf.

  1. Glorot, X. and Bengio, Y. Understanding the difficulty of training deep feedforward neural networks. AISTATS (2010). L’argument de preservació de la variància reproduït al quadre anterior.

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

  3. Ioffe, S. and Szegedy, C. Batch Normalization: Accelerating Deep Network Training by Reducing Internal Covariate Shift. arXiv:1502.03167 (2015). Tingues en compte que l’explicació de l’«internal covariate shift» del títol ha estat substancialment qüestionada des de llavors; la capa funciona, però l’explicació original del perquè és discutida.

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

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

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

  7. Srivastava, N., Hinton, G., Krizhevsky, A., Sutskever, I. and 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. and Mandal, S. Reconciling modern machine-learning practice and the classical bias–variance trade-off. PNAS 116(32), pp. 15849–15854 (2019). L’article que va donar nom al fenomen.

  9. Nakkiran, P., Kaplun, G., Bansal, Y., Yang, T., Barak, B. and Sutskever, I. Deep Double Descent: Where Bigger Models and More Data Hurt. arXiv:1912.02292 (2019). Mostra l’efecte en xarxes profundes reals, i també al llarg de l’eix del temps d’entrenament, a més de l’eix de la mida del model.

A punt per deixar que triï LIA?

Crea amb tots els models d'IA en un sol lloc — comença gratis avui mateix.