: Sobreparametrización, generalización y SAM
El espectacular éxito del aprendizaje profundo moderno, especialmente en los dominios de la visión por computadora y el procesamiento del lenguaje natural, se basa en modelos "sobreparametrizados": modelos con parámetros más que suficientes para memorizar perfectamente los datos de entrenamiento. Funcionalmente, se puede diagnosticar que un modelo está sobreparametrizado cuando puede alcanzar fácilmente una precisión de entrenamiento casi perfecta (cerca del 100%) con una pérdida de entrenamiento cercana a cero para una tarea determinada.
Sin embargo, la utilidad de dicho modelo depende de si funciona bien con los datos de prueba extraídos de la misma distribución que el conjunto de entrenamiento, pero que no se ven durante el entrenamiento. Esta propiedad se llama "generalizabilidad" (la capacidad de un modelo para mantener el rendimiento en nuevos ejemplos) y es esencial para que cualquier modelo de aprendizaje profundo sea útil en la práctica.
La teoría clásica del aprendizaje automático nos dice que los modelos sobreparametrizados deberían sobreajustarse catastróficamente y, por lo tanto, generalizarse mal. Sin embargo, uno de los descubrimientos más sorprendentes de la última década es que los modelos de esta clase suelen generalizar notablemente bien.
Este fenómeno altamente contradictorio ha sido investigado en una serie de artículos, comenzando con los trabajos fundamentales de Belkin et al. (2018) y Nakkiran et al. (2019), que demostró que existe una curva de “doble descenso” para la generalización: a medida que aumenta el tamaño del modelo, la generalización primero empeora (como predice la teoría clásica) y luego mejora nuevamente más allá de un umbral crítico, siempre que el modelo se entrene con los métodos de optimización adecuados.
La Figura 1 muestra una caricatura de una curva de doble descenso. El eje y representa el error de prueba (una medida de generalización, donde un error más bajo indica una mejor generalización), mientras que el eje x muestra el número de parámetros del modelo. A medida que aumenta el tamaño del modelo, el error de entrenamiento (línea azul discontinua) se acerca rápidamente a cero, como se esperaba.
El error de prueba (línea azul sólida) muestra un comportamiento más interesante: inicialmente disminuye con el tamaño del modelo (el primer descenso, resaltado por el círculo rojo izquierdo) y luego aumenta hasta un pico en el umbral de interpolación marcado por la línea discontinua vertical, donde el modelo tiene la peor generalización. Sin embargo, más allá de este umbral, en el régimen sobreparametrizado, el error de prueba vuelve a disminuir (el segundo descenso, resaltado por el círculo rojo de la derecha) y continúa disminuyendo a medida que se agregan más parámetros. Este es el régimen de interés para los modelos modernos de aprendizaje profundo.
En Machine Learning, se encuentran los parámetros de un modelo minimizando una función de pérdida en el conjunto de datos de entrenamiento. Pero, ¿el simple hecho de minimizar nuestra función de pérdida favorita, como la entropía cruzada, en el conjunto de datos de entrenamiento garantiza propiedades de generalización satisfactorias para la clase de modelos sobreparametrizados? La respuesta es, en términos generales, ¡no! Ya sea que uno esté interesado en ajustar un modelo previamente entrenado o entrenar un modelo desde cero, es importante optimizar su algoritmo de entrenamiento para asegurarse de tener un modelo suficientemente generalizable. Esto es lo que hace que la elección del optimizador sea una elección de diseño crucial.
Nitidez-Aware-Minimización (SAM): presentado en un artículo de Foret et al. (2019): es un optimizador diseñado para mejorar la generalización de un modelo sobreparametrizado. En este artículo presento una revisión pedagógica de SAM que incluye:
Una comprensión intuitiva de cómo funciona SAM y por qué mejora la generalización. Una inmersión profunda en el algoritmo, que explica los pasos matemáticos clave involucrados. Una implementación de PyTorch de la clase optimizadora en un bucle de entrenamiento, que incluye una advertencia importante para los modelos con capas BatchNorm. Una demostración rápida de la eficacia del optimizador para mejorar la generalización en una tarea de clasificación de imágenes con un modelo ResNet-18.
El código completo utilizado en este artículo se puede encontrar en este repositorio de Github. ¡Siéntete libre de jugar con él!
La noción de nitidez
Para empezar, intentemos tener una idea intuitiva de por qué simplemente minimizar la función de pérdida puede no ser suficiente para una generalización óptima.
Un panorama útil a tener en cuenta es el del panorama de pérdidas. Para un modelo grande sobreparametrizado, el panorama de pérdidas tiene múltiples mínimos locales y globales. Las geometrías locales alrededor de dichos mínimos pueden variar significativamente a lo largo del paisaje. Por ejemplo, dos mínimos pueden tener valores de pérdida casi idénticos, pero diferir dramáticamente en su geometría local: uno puede ser agudo (valle estrecho) mientras que el otro es plano (valle ancho).
Una medida formal para comparar estas geometrías locales es la "nitidez". En cualquier punto w del panorama de pérdidas con función de pérdida L(w), la nitidez S(w) se define como:
Permítanme desentrañar la definición. Imagine que está en un punto w en el paisaje de pérdidas y perturba los parámetros de tal manera que el nuevo parámetro siempre se encuentra dentro de una bola de radio ρ con centro w. Luego, la nitidez se define como el cambio máximo en la función de pérdida dentro de esta familia de perturbaciones. En la literatura, también se le conoce como la peor nitidez de dirección por razones obvias.
Se puede ver fácilmente que para un mínimo pronunciado (un valle estrecho y empinado) el valor de la función de pérdida cambiará dramáticamente con pequeñas perturbaciones en ciertas direcciones y conducirá a un valor alto de nitidez. Por otro lado, para un mínimo plano (un valle ancho), el valor de la función de pérdida cambiará relativamente lentamente con pequeñas perturbaciones y conducirá a un valor más bajo de nitidez. Por lo tanto, la nitidez da una medida de planitud para un mínimo dado en el paisaje de pérdidas.
Existe una conexión profunda entre la geometría local de un mínimo (especialmente la medida de nitidez) y la propiedad de generalización del modelo resultante. Durante la última década, se ha realizado una cantidad significativa de investigaciones teóricas y empíricas para aclarar esta conexión. Por ejemplo, como señala el artículo de Keskar et al. (2016) señala: los mínimos globales con valores similares de la función de pérdida pueden tener propiedades de generalización significativamente diferentes dependiendo de sus medidas de nitidez.
La lección básica que parece surgir de estos estudios es: los mínimos más planos (menos agudos) se correlacionan positivamente con una mejor generalización de los modelos. En particular, el modelo debe evitar quedarse atrapado en mínimos pronunciados durante el entrenamiento si tiene que generalizarse bien. Por lo tanto, para entrenar un modelo con buena generalización, es necesario asegurarse de que el procedimiento de optimización no solo minimice la función de pérdida sino que también busque maximizar la planitud (o equivalentemente minimizar la nitidez) de los mínimos.
Este es precisamente el problema para el que está diseñado el optimizador SAM, y esto es lo que abordaremos en la siguiente sección.
Un breve comentario: tenga en cuenta que la imagen de arriba ofrece una explicación conceptual de por qué un modelo sobreparametrizado puede evitar potencialmente el problema del sobreajuste. Esto se debe a que un modelo grande tiene un rico panorama de pérdidas que proporciona una multiplicidad de mínimos globales planos con excelentes propiedades de generalización.
El algoritmo de minimización consciente de la nitidez (SAM)
Recordemos la optimización estándar de un modelo. Implica encontrar parámetros del modelo que minimicen una función de pérdida determinada calculada en un mini lote B. En cada paso de tiempo, se calcula el gradiente de pérdida con respecto a los parámetros y se actualizan los parámetros de acuerdo con la regla:
A diferencia de SGD o Adam, SAM no minimiza L directamente. En cambio, en un punto dado del panorama de pérdidas, primero escanea su vecindad de un tamaño dado ρ y encuentra la perturbación que maximiza la función de pérdida. En el segundo paso, minimiza esta función de pérdida máxima. Esto permite al optimizador encontrar parámetros que se encuentran en vecindarios con un valor de pérdida uniformemente bajo, lo que da como resultado valores de nitidez más pequeños y mínimos más planos.
Analicemos el procedimiento con un poco más de detalle. La función de pérdida del optimizador SAM es:
donde ρ denota el límite superior del tamaño de las perturbaciones. La perturbación que maximiza la función L (a menudo llamada perturbación adversaria ya que maximiza la pérdida convencional) se puede encontrar observando que:
donde la segunda igualdad es una aproximación obtenida mediante la expansión de Taylor de la función perturbada en el primer paso, y la última igualdad se deriva de la independencia ϵ del primer término entre corchetes en el paso anterior. Esta última igualdad se puede resolver para la perturbación adversaria de la siguiente manera:
Volviendo a incluir esto en la ecuación de la pérdida SAM, se pueden calcular los gradientes de la pérdida SAM al orden principal en derivadas de ϵ:
Esta es la ecuación más crucial para el procedimiento de optimización. Para el orden principal en derivadas de ϵ, los gradientes de la función de pérdida SAM pueden aproximarse mediante los gradientes de la función de pérdida convencional evaluados en el punto adversamente perturbado. Usando la fórmula anterior para los gradientes, ahora se puede ejecutar el paso estándar del optimizador:
Esto completa una iteración SAM completa. A continuación, traduzcamos el algoritmo del inglés a PyTorch.
Implementación de PyTorch en un ciclo de entrenamiento
En el bloque de código sam_training_loop.py se proporciona un ejemplo ilustrativo de un bucle de entrenamiento con un optimizador SAM. Para ser más concretos, hemos elegido un problema de clasificación de imágenes genérico, pero la misma estructura se aplica en términos generales a una amplia gama de tareas de visión por computadora y PNL. La clase de optimizador SAM se muestra en el bloque de código sam_optimizer_class.py.
Tenga en cuenta que definir un optimizador SAM requiere especificar dos datos:
Un optimizador básico (como SGD o Adam), ya que SAM implica al final un paso de optimizador estándar. Un hiperparámetro ρ, que pone un límite superior al tamaño de las perturbaciones admisibles.
Una única iteración del optimizador implica dos pases hacia adelante y dos pases hacia atrás. Rastreemos los pasos clave del código en sam_training_loop.py:
La línea 5 calcula la función de pérdida L(w, B) para el minilote B actual: el primer pase directo. La línea 6 calcula los gradientes de la función de pérdida L(w, B), el primer paso hacia atrás. La línea 7 llama a la función sam_optimizer.first_step de la clase de optimizador SAM (ver más abajo) que calcula la perturbación adversaria usando la fórmula analizada anteriormente y perturba los pesos del modelo como se analizó anteriormente. La línea 10 calcula la función de pérdida para el modelo perturbado: el segundo pase hacia adelante. La línea 11 calcula los gradientes de la función de pérdida para el modelo perturbado: el segundo paso hacia atrás. La línea 12 llama a la función sam_optimizer.segundo_paso de la clase de optimizador (ver más abajo) que restaura los pesos a w_t y luego usa el optimizador base para actualizar los pesos w_t usando los gradientes calculados en el punto perturbado.
Una advertencia: SAM con BatchNorm
Hay un punto importante que se debe tener en cuenta al implementar SAM en un ciclo de entrenamiento si el modelo tiene algún módulo que incluya capas de normalización por lotes. Durante el entrenamiento, BatchNorm implementa la normalización utilizando las estadísticas del lote actual y actualiza las estadísticas en ejecución en cada paso hacia adelante. Para la evaluación, utiliza las estadísticas en ejecución.
Ahora, como vimos anteriormente, SAM implica dos pases hacia adelante por iteración. Para la primera pasada, BatchNorm funciona de forma estándar. Sin embargo, durante la segunda pasada, utilizamos pesos perturbados para calcular la pérdida, y la función de entrenamiento ingenua en el bloque de código sam_training_loop.py permitirá que las capas BatchNorm actualicen las estadísticas de ejecución también durante la segunda pasada. Esto no es deseable porque las estadísticas en ejecución solo deben reflejar el comportamiento del modelo original, no el modelo perturbado, que es solo un paso intermedio para calcular los gradientes. Por lo tanto, es necesario deshabilitar explícitamente la actualización de estadísticas en ejecución durante el segundo paso y habilitarla antes de la siguiente iteración.
Para este propósito, usaremos dos funciones explícitas enable_bn_stats y enable_bn_stats en el ciclo de entrenamiento (se muestran ejemplos simples de tales funciones en el bloque de código running_stat.py) que alternan el parámetro track_running_stats (línea 4 y línea 9) de la función BatchNorm en PyTorch. El bucle de entrenamiento modificado se proporciona en el bloque de código mod_train.py.
Demostración: clasificación de imágenes con ResNet-18
Finalmente, demostremos cómo la optimización SAM mejora la generalización de un modelo en un ejemplo concreto. Consideraremos un problema de clasificación de imágenes utilizando el conjunto de datos Fashion-MNIST (licencia MIT): consta de 60.000 imágenes de entrenamiento y 10.000 imágenes de prueba en 10 clases distintas y mutuamente excluyentes, donde cada imagen está en escala de grises con 28*28 píxeles.
Como modelo clasificador, elegiremos un PreAct ResNet-18 sin ningún entrenamiento previo. Si bien una discusión sobre la arquitectura precisa de ResNet-18 no es muy relevante para nuestro propósito, recordemos que el modelo consta de una secuencia de bloques de construcción, cada uno de los cuales se compone de capas convolucionales, capas BatchNorm y activación ReLU con conexiones omitidas. PreAct (preactivación) indica que la función de activación (ReLU) viene antes de la capa convolucional en cada bloque. Para un ResNet-18 estándar, es al revés. Remitiría al lector al artículo: He et al. (2015) – para más detalles sobre la arquitectura.
Sin embargo, lo que es importante tener en cuenta es que este modelo tiene alrededor de 11,2 millones de parámetros y, por lo tanto, desde la perspectiva del aprendizaje automático clásico, es un modelo sobreparametrizado con una relación parámetro-muestra de aproximadamente 186:1. Además, dado que el modelo incluye capas BatchNorm, debemos tener cuidado al deshabilitar las estadísticas de ejecución para la segunda pasada mientras usamos SAM.
Ahora estamos listos para realizar el siguiente experimento. Primero entrenamos el modelo en el conjunto de datos Fashion-MNIST con el optimizador SGD estándar y luego con el optimizador SAM usando el mismo SGD como optimizador base. Consideraremos una configuración simple con una tasa de aprendizaje fija lr = 0,05 y con el impulso y la caída del peso establecidos en cero. El hiperparámetro ρ en SAM se establece en 0,05. Todas las ejecuciones se realizan en una única GPU A100.
Dado que cada actualización de peso de SAM requiere dos pasos de retropropagación (uno para calcular las perturbaciones y otro para calcular los gradientes finales), para una comparación justa, cada ejecución de entrenamiento que no sea SAM debe ejecutar el doble de épocas que cada ejecución de entrenamiento SAM. Por lo tanto, tendremos que comparar una métrica de una época de ejecución de entrenamiento SAM con una métrica de dos épocas de ejecución de entrenamiento no SAM. A esto lo llamaremos "época estandarizada" y una métrica registrada en épocas estandarizadas se etiquetará como metric_st. Restringiremos el experimento a 150 épocas estandarizadas, lo que significa que el entrenamiento SAM se ejecuta durante 150 épocas y el entrenamiento no SAM se ejecuta durante 300 épocas. Entrenaremos el modelo optimizado para SAM durante 50 épocas adicionales para tener una idea de cómo se comporta el modelo en un entrenamiento más prolongado.
Al intentar comprobar qué optimizador ofrece una mejor generalización, compararemos las dos métricas siguientes después de cada época de entrenamiento estandarizada:
Precisión de la prueba: rendimiento del modelo en el conjunto de datos de prueba. Brecha de generalización: diferencia entre la precisión del entrenamiento y la precisión de la prueba.
La precisión de la prueba es una medida absoluta de qué tan bien se generaliza el modelo después de un cierto número de épocas de entrenamiento. La brecha de generalización, por otro lado, es un diagnóstico que indica en qué medida se está sobreajustando un modelo en una determinada etapa de entrenamiento.
Comencemos comparando los gráficos Training_loss_st y Training_accuracy_st, como se muestra en la Figura 3. El modelo con SGD alcanza una pérdida cercana a cero y una precisión de entrenamiento cercana al 99% dentro de 150 épocas, como se esperaba de un modelo sobreparametrizado. Es evidente que SAM entrena lentamente en comparación con SGD y necesita épocas más estandarizadas para alcanzar una precisión de entrenamiento casi perfecta. Esto es evidente por el hecho de que la pérdida de entrenamiento y la precisión del entrenamiento continúan mejorando a medida que se entrena el modelo optimizado para SAM durante más épocas más allá de las 150 estipuladas.
Precisión de la prueba. Los gráficos de la Figura 4 comparan las precisiones de las pruebas para los dos casos después de cada época estandarizada.
El modelo optimizado para SGD alcanza una precisión de prueba del 92% alrededor de la época 50 y se estabiliza alrededor de ese valor durante las siguientes 100 épocas. El modelo optimizado para SAM se generaliza mal en la fase inicial del entrenamiento (hasta alrededor de 80 épocas), como se desprende de las menores precisiones de prueba en esta fase en comparación con el gráfico SGD. Sin embargo, alrededor de la época 80, alcanza el gráfico SGD y finalmente lo supera por un estrecho margen.
Para esta ejecución específica, al final de 150 épocas, la precisión de la prueba para SAM es test_SAM = 92,5 %, mientras que la de SGD es test_SGD = 92,0 %. Tenga en cuenta que esto es a pesar del hecho de que el modelo entrenado con SAM tiene una precisión de entrenamiento y una pérdida de entrenamiento mucho menores en esta etapa. Si se entrena el modelo SAM durante otras 50 épocas, la precisión de la prueba mejora ligeramente hasta el 92,7%.
Brecha de generalización. La evolución de la brecha de generalización después de cada época estandarizada en el transcurso del proceso de formación se muestra en la Figura 5.
La brecha para el modelo SGD crece de manera constante con el entrenamiento y después de 150 épocas alcanza la brecha_SGD = 6,8%, mientras que para SAM crece mucho más lentamente y alcanza la brecha_SAM = 2,3%. Tras un entrenamiento adicional durante otras 50 épocas, la brecha para SAM aumenta a alrededor del 3%, pero sigue siendo mucho menor en comparación con el valor SGD.
Si bien la diferencia en la precisión de las pruebas es pequeña entre los dos optimizadores para el conjunto de datos Fashion-MNIST, existe una diferencia no trivial en las brechas de generalización, lo que demuestra que la optimización con SAM conduce a una mejor generalización.
Comentarios finales
En este artículo, presenté una revisión pedagógica de SAM como un optimizador que mejora significativamente la generalización de modelos de aprendizaje profundo sobreparametrizados. Discutimos la motivación y la intuición detrás de SAM, analizamos un desglose paso a paso del algoritmo y estudiamos un ejemplo simple que demuestra su efectividad en comparación con un optimizador SGD estándar.
Hay varios aspectos interesantes de SAM que no tuve la oportunidad de cubrir aquí. Permítanme mencionar brevemente dos de ellos. En primer lugar, como herramienta práctica, SAM es particularmente útil para ajustar modelos previamente entrenados en pequeños conjuntos de datos, algo explorado en detalle por Foret et al. (2019) para arquitecturas tipo CNN y en muchos trabajos posteriores para arquitecturas más generales. En segundo lugar, dado que abrimos nuestra discusión con la conexión entre mínimos planos en el panorama de pérdidas y la generalización, es natural preguntar si un modelo entrenado con SAM, que mejora demostrablemente la generalización, realmente converge a un mínimo más plano. Esta es una pregunta no trivial que requiere un análisis cuidadoso del espectro hessiano del modelo entrenado y una comparación con su contraparte entrenado con SGD. ¡Pero esa es una historia para otro día!
¡Gracias por leer! Si te ha gustado el artículo y te interesa leer más artículos pedagógicos sobre aprendizaje profundo, sígueme en Medium y LinkedIn. A menos que se indique lo contrario, todas las imágenes y gráficos utilizados en este artículo fueron generados por el autor.