Conseguir que entrene y que generalice
Una red de seis capas cuya pérdida no se mueve de ln 2, arreglada medición a medición. Luego double descent: 5.000 parámetros en 40 puntos.
En esta página
La red del Capítulo 5 funciona. Tiene nueve parámetros, aprende XOR y sus gradientes coinciden con PyTorch hasta dieciséis decimales.
Hazla de seis capas de profundidad y deja de aprender por completo. No lentamente: por completo. Aquí tienes una red de seis capas en un problema de clasificación de dos espirales, entrenada durante 5000 pasos:
step 1: loss 0.693147
step 5000: loss 0.693147
accuracy: 50.0 %Ese número no es arbitrario. es la entropía cruzada binaria de un modelo que devuelve probabilidad para todo, y el 50 % es lanzar una moneda en un conjunto de datos equilibrado. Tras cinco mil pasos, la red no se ha movido ni una sola cifra. Nada ha fallado, nada ha avisado y los gradientes siguen siendo exactamente correctos.
Este capítulo trata sobre la brecha entre una red que se ejecuta y una red que funciona. Tiene dos mitades que parecen temas distintos y son el mismo trabajo: hacer que la pérdida vaya hacia abajo, y hacer que baje en datos que el modelo nunca ha visto.
Por qué la red de seis capas está muerta
Enlace a la sección: Por qué la red de seis capas está muertaEmpieza mirando, en vez de adivinar. Pasa un lote de entradas por la red e imprime la desviación estándar de las activaciones en cada capa, y después la desviación estándar de los gradientes de los pesos:
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 inicializaciones, misma arquitectura, seis capas de :
| inicialización | desv. estándar de activación, capas 1→6 |
|---|---|
| normal, std | 0.0145 · 0.0016 · 0.0002 · 0.0000 · 0.0000 · 0.0000 |
| normal, std | 0.6573 · 0.9296 · 0.9585 · 0.9634 · 0.9637 · 0.9625 |
| Xavier | 0.1579 · 0.1493 · 0.1353 · 0.1333 · 0.1325 · 0.1403 |
| inicialización | desv. estándar de gradiente, primera capa → última |
|---|---|
| normal, std | 3.20e-06 · 4.97e-07 · … · 6.40e-06 |
| normal, std | 1.94e+03 · 2.28e+02 · 1.22e+02 · 4.43e+01 · 1.85e+01 · 7.30e+00 |
| Xavier | 2.31e+00 · 4.50e-01 · 4.26e-01 · 3.89e-01 · 4.39e-01 · 4.73e-01 |
La primera fila es la red de arriba, y no está aprendiendo despacio: ya no le queda señal. En la capa cuatro, la desviación estándar de activación se ha subdesbordado hasta cero a cuatro decimales. Cada entrada produce la misma salida, la salida es una constante y el gradiente de una constante no es nada. Los pesos se inicializaron pequeños «para ir sobre seguro», y pequeño fue fatal.
La segunda fila es el fallo contrario y merece entenderse porque va contra la intuición. Las activaciones parecen sanas —alrededor de 0.96—, pero eso es saturado, clavado cerca de su límite, exactamente el régimen que el Capítulo 5 midió como una pérdida de un factor de casi diez mil en el gradiente. Y aun así los gradientes son enormes: 1940 en la primera capa. Ambas cosas son ciertas a la vez. Cada paso hacia atrás multiplica por , y con 128 entradas a varianza unitaria ese factor tiene una ganancia de alrededor de , que supera la contracción del saturado. Los gradientes crecen geométricamente al retroceder. Este es el gradiente explosivo, y produce valores de pérdida de nan en unos pocos pasos en cualquier entrenamiento real.
La tercera fila es lo que quieres: activaciones con una escala más o menos constante a través de la profundidad, gradientes con una escala más o menos constante a través de la profundidad. Nada muere, nada explota.
Normalización, y cuál sobrevivió
Enlace a la sección: Normalización, y cuál sobrevivióInicializar bien fija la escala en el paso cero. No la mantiene fija: los pesos se mueven, y para el paso cinco mil el cuidadoso argumento de varianza ya no se aplica.
Las capas de normalización imponen la escala de forma continua. Dado un vector de activaciones, resta una media, divide por una desviación estándar y después aplica una escala aprendida y un desplazamiento para que la capa pueda deshacer la normalización si eso resulta ser lo que quiere:
La única pregunta real es sobre qué promedias. La normalización por lote3 toma y a través de la dimensión del lote, una estadística por característica. La normalización por capa4 los toma a través de las características, una estadística por ejemplo.
Esa elección parece menor y decide casi todo lo que viene después:
BatchNorm hace que la salida de cada ejemplo dependa de los demás ejemplos que hayan acabado en su lote. Durante el entrenamiento eso es un regularizador suave. En inferencia no hay lote, así que tiene que mantener una media móvil de las estadísticas recogidas durante el entrenamiento, lo que significa que la capa se comporta de forma distinta en modo entrenamiento y en modo evaluación, y olvidarse de cambiar de modo es uno de los bugs más comunes del campo. También se degrada con lotes pequeños, y es incómoda con secuencias de longitud variable, porque «la media sobre el lote en la posición 40» se calcula con cuantas secuencias resulten ser tan largas.
LayerNorm normaliza cada ejemplo por sí solo. Sin dependencia del lote, sin estadísticas móviles, comportamiento idéntico en entrenamiento e inferencia, indiferente al tamaño del lote, indiferente a la longitud de la secuencia. Cada una de esas propiedades es un requisito, no un detalle agradable, cuando estás generando un token cada vez para un usuario, que es adonde acaba llegando el Capítulo 13.
Por eso LayerNorm es la que volverás a encontrar en el Capítulo 9 sin cambios: el bloque transformer la usa, y la usa por las razones de la columna derecha, no porque funcione mejor en abstracto.
Arreglar una cosa cada vez, que es la skill real
Enlace a la sección: Arreglar una cosa cada vez, que es la skill realCuatro posibles arreglos para la red muerta: inicialización Xavier, LayerNorm, conexiones residuales y Adam en vez de SGD. La tentación es aplicar los cuatro y seguir adelante. Hazlo y nunca sabrás cuál importaba, y la próxima vez que ocurra no tendrás un método: solo un ritual.
Así que aplícalos de uno en uno. Misma semilla, mismos datos, misma arquitectura, 800 pasos:
| qué se añadió | pérdida final | precisión |
|---|---|---|
| nada | 0.6931 | 50.0 % |
| inicialización Xavier | 0.5692 | 60.4 % |
| LayerNorm | 0.6230 | 61.5 % |
| conexiones residuales | 0.6651 | 56.6 % |
| Adam | 0.6787 | 58.7 % |
| los cuatro | 0.0000 | 100.0 % |
Lee esa tabla como la leerías a las 2 de la mañana y la conclusión es: nada funciona solo, todo funciona junto, por tanto el deep learning es alquimia. Esa conclusión es falsa, y descubrir por qué es lo más útil de este capítulo.
Dale a cada ejecución seis veces el presupuesto —5000 pasos en vez de 800— y cambia por completo:
| qué se añadió | pérdida final @ 5000 | precisión |
|---|---|---|
| nada | 0.6931 | 50.0 % |
| inicialización Xavier | 0.0007 | 100.0 % |
| LayerNorm | 0.0002 | 100.0 % |
| conexiones residuales | 0.6653 | 56.7 % |
| Adam | 0.6908 | 53.4 % |
| Xavier + Adam | 0.0000 | 100.0 % |
| Xavier + LayerNorm | 0.0001 | 100.0 % |
Ahora la imagen es nítida, y es un diagnóstico en vez de un ritual.
La inicialización por sí sola lo arregla. La normalización por sí sola lo arregla. Cada una ataca la enfermedad real —la señal hacia delante colapsando a cero— y cualquiera de las dos es suficiente. A los 800 pasos solo parecían mérito parcial, porque habían resuelto el problema y aún estaban saliendo del agujero.
Las conexiones residuales y Adam no lo arreglan, con ningún presupuesto. No porque sean malas, sino porque tratan otra enfermedad. Una conexión residual da al gradiente un camino alrededor de una capa bloqueante; eso vale muchísimo cuando el problema es el gradiente, y no vale nada cuando la señal hacia delante ya es cero, porque un atajo alrededor de una capa muerta sigue llevando un valor muerto. Adam reescala el paso de cada parámetro según su propio historial de gradientes; eso ayuda cuando los gradientes tienen magnitudes muy distintas, y no puede resucitar una red cuya salida no depende de su entrada.
Y «nada» sigue siendo exactamente 0.6931 después de cinco mil pasos. No 0.6929. No va lenta; está muerta, y esa distinción ahora es visible de una forma que antes no lo era, porque tienes la fila que dice que un arreglo funciona para comparar.
Ganarse PyTorch
Enlace a la sección: Ganarse PyTorchA partir de aquí este curso usa PyTorch. Eso debería ganarse, no anunciarse, así que aquí tienes exactamente qué hace que tú ya sabes hacer.
Un optimizador es una regla para convertir gradientes en actualizaciones de parámetros. El gradient descent puro usa el gradiente. Momentum usa una media móvil de él, lo que suaviza el ruido y gana velocidad en direcciones que se mantienen consistentes:
v = beta * v + p.grad
p -= lr * v Adam5 mantiene dos medias móviles —del gradiente y del gradiente al cuadrado— y divide una por la raíz cuadrada de la otra, de modo que cada parámetro recibe un paso escalado a la magnitud reciente de su propio gradiente:
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) Diez líneas. Ejecuta ambos contra torch.optim en el mismo problema durante 50 pasos:
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-07Idénticos a precisión float32. torch.optim.Adam son esas cinco líneas, más décadas de cuidado con casos límite y un kernel en C++. Ese es el intercambio que harás a partir de ahora: no magia a cambio de comprensión, sino velocidad a cambio de líneas que ya has escrito.
Por qué existe Adam: curvatura
Enlace a la sección: Por qué existe Adam: curvaturaLa explicación habitual de Adam es «learning rates adaptativos por parámetro», que es una descripción más que una razón. La razón es la geometría, y se puede medir.
Toma una pérdida cuya curvatura difiere entre direcciones: empinada en una, suave en otra. SGD tiene un learning rate global, así que debe escoger un valor lo bastante pequeño para ser estable en la dirección más empinada; y ese valor queda entonces demasiado pequeño para la dirección suave, donde el progreso se arrastra. Esto es lo que causa la imagen clásica de gradient descent zigzagueando por un valle estrecho.
Dos ratios de curvatura, tres optimizadores, 300 pasos, y a cada optimizador se le da el mejor learning rate de un barrido para que nadie tenga desventaja:
| ratio de curvatura | SGD | SGD + momentum | Adam |
|---|---|---|---|
| 10 : 1 | error 0.000002 | error 0.000000 | error 0.000000 |
| 1000 : 1 | error 1.925485 | error 0.001432 | error 0.000000 |
| divergió en (1000:1) | 4 de 8 tasas | 4 de 8 tasas | 0 de 6 tasas |
Con un ratio de diez, todo funciona y no hay nada que discutir. Con mil, SGD puro no puede llegar a la respuesta con ningún learning rate probado —su mejor resultado sigue siendo un error de 1.93— y diverge directamente con la mitad de las tasas. Adam aterriza exactamente en el objetivo y no diverge con ninguna.
Esa última columna es la razón práctica por la que Adam es el valor por defecto. No es que Adam encuentre soluciones mejores; en problemas bien condicionados, SGD bien ajustado a menudo lo iguala o lo supera. Es que Adam es mucho menos sensible al learning rate que elegiste, y las redes reales tienen ratios de curvatura mucho peores que mil a través de sus millones de parámetros.
Aquí van dos piezas más y ambas son de una línea. Gradient clipping reescala el vector de gradiente siempre que su norma supera un umbral, lo que convierte la fila «la pérdida salta de repente a un valor enorme» de la tabla de diagnóstico en un no suceso. Y learning rate schedules: un warmup corto desde casi cero durante los primeros cientos de pasos, porque las estimaciones de varianza de Adam son basura hasta que han visto algunos gradientes y un paso de tamaño completo tomado sobre basura puede destrozar una inicialización; después cosine decay hacia cero, porque terminar una ejecución con el mismo tamaño de paso con el que empezaste significa temblar alrededor del mínimo en vez de asentarte en él.
La segunda mitad: el modelo que ajusta perfecto y no predice nada
Enlace a la sección: La segunda mitad: el modelo que ajusta perfecto y no predice nadaTodo hasta ahora trataba de hacer bajar la pérdida. Ahora viene la mitad más difícil, porque que la pérdida baje no es el objetivo: es un proxy del objetivo, y el proxy falla de una forma concreta y famosa.
Doce puntos de una función suave con un poco de ruido. Ajusta polinomios de grado creciente:
| grado | RMSE entrenamiento | RMSE test |
|---|---|---|
| 1 | 0.764499 | 0.6985 |
| 3 | 0.252605 | 0.3031 |
| 5 | 0.164437 | 0.1568 |
| 9 | 0.088960 | 0.2347 |
| 11 | 0.000000 | 1.2094 |
El grado 11 con 12 puntos pasa por todos y cada uno exactamente —error de entrenamiento cero a seis decimales— y es ocho veces peor que el grado 5 en datos que no ha visto. Pide al grado 3 y al grado 11 que predigan en , justo fuera del rango de entrenamiento:
degree 3: predicts -1.053 (truth -0.012)
degree 11: predicts +61.224 (truth -0.012)Sesenta y uno, cuando la respuesta es aproximadamente cero. El modelo no aprendió la función; aprendió los doce puntos, y entre ellos hace lo que la aritmética exija.
Esto es overfitting, y su opuesto —grado 1, que no puede representar la curva en absoluto y es malo en todas partes— es underfitting. La explicación clásica divide el error esperado de un modelo en tres partes: sesgo, el error de que el modelo sea demasiado rígido para representar la verdad; varianza, el error de que el modelo sea tan flexible que persiga el ruido de esta muestra concreta; y ruido irreducible, que nada arregla. Los modelos simples están sesgados, los modelos flexibles tienen alta varianza, y la receta clásica es encontrar el punto dulce en medio: grado 5 en la tabla de arriba.
Las herramientas estándar atacan todas el término de varianza:
- Regularización L2 (weight decay) añade a la pérdida, tirando de los pesos hacia cero y haciendo la función más suave. En la tabla de arriba, el mayor coeficiente del grado 11 hace el daño; penalizar el tamaño lo desactiva.
- L1 añade en su lugar. La diferencia no es cosmética: el gradiente de L2 es proporcional al peso y por tanto se reduce a medida que lo hace el peso, acercándose a cero sin llegar, mientras que el gradiente de L1 es una constante que sigue empujando hasta el final. Por tanto L1 produce pesos que son exactamente cero: selecciona características. L2 produce pesos pequeños. Usa L2 cuando quieras suavidad, L1 cuando quieras escasez.
- Dropout7 pone a cero un subconjunto aleatorio de activaciones en cada paso de entrenamiento, de modo que ninguna unidad pueda depender de que otra unidad concreta esté presente.
- Parada temprana vigila la pérdida de validación y se detiene cuando empieza a subir.
- Aumento de datos fabrica más ejemplos de entrenamiento a partir de los que tienes, lo que ataca el problema en su origen: el overfitting es tanto una escasez de datos como un exceso de parámetros.
- Validación cruzada divide los datos de maneras y entrena veces, lo que compra una estimación fiable del error de test cuando tienes demasiado pocos datos para reservar un conjunto aparte.
Double descent, o por qué la sección anterior no es toda la historia
Enlace a la sección: Double descent, o por qué la sección anterior no es toda la historiaAhora el hecho que rompe la imagen.
La historia sesgo-varianza dice que, pasado el punto dulce, más parámetros significan peor generalización. Los modelos de lenguaje modernos tienen muchos más parámetros de los que las reglas clásicas permitirían para los datos que ven, y generalizan magníficamente. Ambas afirmaciones son ciertas, y reconciliarlas es lo más útil de este capítulo.
Cuarenta puntos de entrenamiento, entradas de veinte dimensiones, características ReLU aleatorias, y el número de características barrido de 2 a 5000, con la solución de norma mínima elegida siempre que hay muchas que ajustan:
| RMSE entrenamiento | RMSE test | |||
|---|---|---|---|---|
| 10 | 0.25 | 0.8822 | 1.2520 | 1.89 |
| 20 | 0.50 | 0.5962 | 1.1634 | 2.59 |
| 30 | 0.75 | 0.3896 | 1.5323 | 4.15 |
| 38 | 0.95 | 0.1769 | 3.7163 | 10.25 |
| 40 | 1.00 | 0.0000 | 5.8140 | 14.83 |
| 42 | 1.05 | 0.0000 | 3.1623 | 9.35 |
| 60 | 1.50 | 0.0000 | 1.1058 | 2.78 |
| 200 | 5.00 | 0.0000 | 0.6638 | 0.98 |
| 1500 | 37.50 | 0.0000 | 0.5859 | 0.33 |
| 5000 | 125.00 | 0.0000 | 0.5664 | 0.18 |
Léela en tres partes. Hasta la historia clásica se cumple exactamente: el error cae y luego empieza a subir. En —el umbral de interpolación, donde el modelo tiene exactamente suficientes parámetros para pasar por todos los puntos de entrenamiento— el error de test alcanza su pico, en 5.81, cinco veces peor que el modelo pequeño. Ese pico es la advertencia clásica, y es real.
Luego desciende de nuevo. Y sigue descendiendo, más allá de , más allá de , hasta , donde el error de test de 0.5664 es mejor que el mejor modelo infraparametrizado jamás logró. Un modelo con 5000 parámetros ajustado a 40 puntos es el mejor modelo de la tabla.
Esto es double descent,89 y el mecanismo se ve en la última columna. Una vez que hay infinitas configuraciones de parámetros que ajustan exactamente los datos de entrenamiento, y cuál obtienes depende de cómo elijas. La solución de norma mínima escoge la más pequeña, y muestra lo que eso significa: alcanza un pico de 14.83 justo en el umbral —donde hay exactamente una solución interpolante y te quedas con ella, por extrema que sea— y luego cae monótonamente a medida que crece, porque más parámetros significan más soluciones interpolantes entre las que elegir, lo que significa que la más pequeña disponible se hace más pequeña. En la norma es 0.18, ochenta veces menor que en el umbral.
Así que los parámetros extra no añaden complejidad. Añaden elección, y la regla de selección gasta esa elección en simplicidad. La regularización no está en la función de pérdida; está en el algoritmo. Gradient descent desde una inicialización pequeña tiene un sesgo documentado hacia soluciones de norma pequeña, que es por lo que este comportamiento aparece en redes reales entrenadas de la forma ordinaria y no solo en el álgebra lineal de arriba.
La consecuencia práctica, de la que depende el Capítulo 10: «el modelo tiene más parámetros que datos, así que hará overfit» no es un argumento válido. Era una buena regla cuando los modelos vivían a la izquierda del umbral. Ahora todo lo interesante vive muy a la derecha, donde la regla se invierte.
Adónde va esto ahora
Enlace a la sección: Adónde va esto ahoraLas herramientas de este capítulo bastan para entrenar una red que funcione con datos que puedes poner en una tabla: filas de números, una columna de etiquetas.
El lenguaje no es eso. Antes de que un modelo pueda predecir la siguiente palabra, algo tiene que decidir qué es siquiera una «palabra», y la respuesta no son ni letras ni palabras, sino un vocabulario que el modelo aprende a partir de los bytes brutos de los datos de entrenamiento. Esa decisión, tomada una vez antes de que empiece el entrenamiento, determina cuántas cosas puede decir el modelo, cuánto cuesta una petición y por qué modelos que pueden aprobar un examen de Derecho no pueden contar de forma fiable las letras de strawberry.
El Capítulo 7 construye un tokenizer.
Fuentes y método
Enlace a la sección: Fuentes y métodoPara las conexiones residuales usadas arriba, He et al., Deep Residual Learning for Image Recognition (arXiv:1512.03385). Building makemore Part 3: Activations & Gradients, BatchNorm, de Andrej Karpathy, recorre el diagnóstico de histogramas de activación en un modelo real y es el mejor tratamiento práctico de la primera mitad de este capítulo. Las clases 8 y 11–13 de Learning From Data, de Yaser Abu-Mostafa, presentan correctamente la teoría clásica de la generalización, incluidas las partes que este capítulo comprimió en un párrafo.
Referencias
Enlace a la sección: Referencias-
Glorot, X. and Bengio, Y. Understanding the difficulty of training deep feedforward neural networks. AISTATS (2010). El argumento de preservación de la varianza reproducido en el recuadro de arriba. ↩
-
He, K., Zhang, X., Ren, S. and Sun, J. Delving Deep into Rectifiers: Surpassing Human-Level Performance on ImageNet Classification. arXiv:1502.01852 (2015). ↩
-
Ioffe, S. and Szegedy, C. Batch Normalization: Accelerating Deep Network Training by Reducing Internal Covariate Shift. arXiv:1502.03167 (2015). Ten en cuenta que la explicación de «internal covariate shift» del título ha sido desde entonces bastante cuestionada; la capa funciona, pero la explicación original de por qué lo hace está discutida. ↩
-
Ba, J. L., Kiros, J. R. and Hinton, G. E. Layer Normalization. arXiv:1607.06450 (2016). ↩
-
Kingma, D. P. and Ba, J. Adam: A Method for Stochastic Optimization. arXiv:1412.6980 (2014). ↩
-
Loshchilov, I. and Hutter, F. Decoupled Weight Decay Regularization. arXiv:1711.05101 (2017). ↩
-
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). ↩
-
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). El artículo que dio nombre al fenómeno. ↩
-
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). Muestra el efecto en redes profundas reales, y también a lo largo del eje del tiempo de entrenamiento, además del eje del tamaño del modelo. ↩