Tabla de Contenidos
Si alguna vez ha leído la implementación de un modelo de aprendizaje profundo, es probable que ya se haya encontrado con BatchNorm (Normalización por lotes). Esta es una operación muy común que se utiliza para acelerar el entrenamiento de modelos grandes y para estabilizar los inestables. Sin embargo, si es un profesional, es muy posible que también haya tenido dificultades con esa operación, que notoriamente plantea muchos problemas. En este artículo, revisaremos los problemas que a menudo encontramos y propondremos algunas soluciones.
¿Qué es una capa de normalización por lotes?
BatchNorm tiene como objetivo resolver el problema del desplazamiento de covariables. Esto significa que, para una capa dada en una red profunda, la salida tiene una media y una desviación estándar en todo el conjunto de datos. Durante el entrenamiento, esta media y desviación estándar no están restringidas y pueden evolucionar aleatoriamente, lo que puede plantear algunos problemas de estabilidad numérica. La operación BatchNorm intenta eliminar este problema normalizando la salida de la capa. Sin embargo, es demasiado costoso evaluar la media y la desviación estándar en todo el conjunto de datos, por lo que solo las evaluamos en un lote de datos.
Esto funciona bien en la práctica, pero no podemos hacer lo mismo en el momento de la inferencia, porque recibimos los datos uno por uno, por lo que los promedios ya no tienen sentido. Para resolver este problema, las implementaciones modernas proponen calcular una media móvil sobre los datos.
El problema
En resumen, el comportamiento es diferente entre el entrenamiento y la inferencia. En el momento del entrenamiento t,mt, y σt se utilizan, pero en el momento de la inferencia mt y σt se utilizan. Esta diferencia es la raíz de todos los males, ya que las métricas en la validación y en el entrenamiento pueden ser muy diferentes. Más precisamente, a medida que la cantidad real evoluciona durante el entrenamiento, la media móvil a menudo se quedará atrás, lo que puede causar una diferencia significativa. En principio, si el lote es grande y si el modelo converge bien, entonces esas cantidades deberían ser las mismas. Pero en la práctica, a menudo es incorrecto o poco práctico. Por ejemplo, no será obvio si una gran discrepancia entre la pérdida de entrenamiento y la de validación se debe a un sobreajuste severo, o porque esas cantidades aún no han convergido.
Más peligrosamente, observamos regularmente que, aunque la pérdida de entrenamiento converge a algún valor, la pérdida de validación puede permanecer considerablemente más alta, debido a que la media y la desviación estándar de BatchNorm nunca se estabilizan. Nosotros, los autores, no estamos completamente seguros de la causa del problema, pero creemos que esto puede ocurrir cuando el mínimo está fuertemente degenerado. Por ejemplo, en un paisaje de pérdida como el ilustrado a continuación, el modelo se moverá aleatoriamente en el valle circular, haciendo que la media móvil se quede atrás para siempre.

La solución
Lo primero que hay que hacer si te encuentras con este problema es probar algunos trucos estándar. Aquí tienes algunos típicos:
- Intenta usar otra solución de normalización (es decir, LayerNorm, InstanceNorm…);
- Aumenta el tamaño del lote (batch size), lo que puede estabilizar la estimación de la media y la desviación estándar dentro del lote;
- Juega con el parámetro de momento de la media móvil. Este indica cuánto persisten los lotes anteriores en la media móvil, es decir, cuánto pueden "retrasarse" las estimaciones;
- Mezcla tu conjunto de entrenamiento en cada época para evitar la correlación entre los puntos de datos.
Sin embargo, a veces esos trucos básicos no serán suficientes. En ese caso, proponemos usar una estrategia más potente.
Ten en cuenta que la capa BatchNorm presenta dos comportamientos diferentes:
- En lo que llamaremos Modo de Estimación por Lotes (Batch Estimation Mode), la media y las desviaciones estándar se estiman en el lote. Este es el modo utilizado durante el entrenamiento;
- En lo que llamaremos Modo de Inferencia (Inference Mode), la media y la desviación estándar se basan en estimaciones previas, es decir, en la media móvil. Esto es lo que se suele usar durante la validación y la inferencia.
¡Nuestra solución consta de dos pasos! Primero, desactivamos la diferencia entre el entrenamiento y la validación utilizando siempre el modo de estimación por lotes. En segundo lugar, para poder usar el modelo en producción, aún necesitamos estimar la media y la desviación estándar para poder usar el modo de inferencia. Así, una vez entrenado el modelo, calculamos la media y la desviación estándar que se utilizarán. Al hacerlo, se evalúan en un modelo con pesos fijos y evitamos el efecto de "retraso" descrito anteriormente. Más concretamente, después del entrenamiento, congelamos todos los pesos del modelo y ejecutamos una época para estimar la media móvil en todo el conjunto de datos.
Experimentando con la solución
Para mostrar la ventaja de nuestra solución, hagamos un pequeño experimento. Usamos intencionadamente una arquitectura muy deficiente y la entrenamos con una tasa de aprendizaje relativamente alta, lo que resultó en un modelo con BatchNorm inestable. El código escrito en Python y que utiliza PyTorch está disponible.
La red es una pila de 3 capas de convolución, con activación BatchNorm y ReLU, seguida de una capa de agrupamiento promedio global (global average pooling). La entrenamos en MNIST durante 10 épocas utilizando el algoritmo de optimización Adam. En la siguiente figura se muestran la precisión de entrenamiento y validación por época en 4 modos:
- Modo 0: No se utilizan capas BatchNorm.
- Modo 1: BatchNorm básico sin modificaciones.
- Modo 2: BatchNorm casi inteligente: activamos las estadísticas en ejecución para la inferencia, pero no ejecutamos el modelo durante 1 época para estimar la media móvil de las estadísticas.
- Modo 3: BatchNorm inteligente: estimamos en 1 época las estadísticas promedio del conjunto de datos antes del modo de inferencia.

Observamos dos cosas. Primero, BatchNorm ayuda a aumentar la precisión. En segundo lugar, sin nuestra solución, la métrica de validación es errática y poco informativa. Finalmente, proporcionamos la precisión de prueba para las 4 situaciones.


Como puede ver, podríamos obtener mejores resultados utilizando nuestra solución. El tercer modo es realmente malo: activamos las estadísticas en ejecución (modo de inferencia) pero no estimamos esas estadísticas en el conjunto de datos, por lo que al probar en condiciones de inferencia con un tamaño de lote de 1, obtenemos malos resultados. Esto demuestra la necesidad de combinar las estadísticas en ejecución en el momento de la inferencia con la estimación de las estadísticas del conjunto de datos en una época completa del conjunto de datos antes de usar el modelo para la inferencia.
¿Es perfecta esa solución?
¡No, obviamente no! Todavía pueden ocurrir muchas cosas malas. Lo más complicado es que su media y desviación estándar estimadas seguirán siendo diferentes de la estimación del lote y algunos fenómenos realmente extraños aún pueden afectarle gravemente. Por ejemplo, se ha demostrado que algunos modelos pueden codificar información en el ruido estadístico. Afortunadamente, esos casos extremos son muy escasos y la experiencia ha demostrado que esta solución es bastante robusta, solo debería mejorar su rendimiento y ahorrarle muchos dolores de cabeza. Si desea evitar comportamientos extraños con sus capas de BatchNorm, adelante.
Imagen destacada de Pietro Jeng
Acerca de




.webp)
.webp)
.webp)