Reducir la memoria LLM en un 84%: una inmersión profunda en los núcleos fusionados

o has perfeccionado un LLM, probablemente te hayas topado con una pared en el último paso: la pérdida de entropía cruzada.

El culpable es el cuello de botella logit. Para predecir el siguiente token, proyectamos un estado oculto en un espacio de vocabulario masivo. Para Llama 3 (128,256 tokens), solo la matriz de peso tiene más de 525 millones de parámetros. Si bien eso es solo ~1 GB en bfloat16, el tensor logit intermedio es el verdadero problema. Para lotes grandes, puede exceder fácilmente los 80 GB de VRAM solo para calcular una única pérdida escalar.

La optimización de esta capa es la forma en que bibliotecas como Unsloth y Liger-Kernel logran reducciones de memoria tan masivas. En este artículo, construiremos un kernel fusionado de entropía cruzada y lineal desde cero en Triton. Deduciremos los cálculos e implementaremos un paso hacia adelante y hacia atrás en mosaico que reduce drásticamente el uso máximo de memoria en un 84 %.

Nota sobre el rendimiento: esta implementación es principalmente educativa. Priorizamos la claridad matemática y el código Triton legible mediante el uso de operaciones atómicas globales. Si bien resuelve el cuello de botella de la memoria, igualar velocidades de producción requeriría implementaciones significativamente más complejas que están fuera del alcance de este artículo.

Esta publicación es parte de mi serie Tritón. Usaremos conceptos como mosaico y softmax en línea que cubrimos anteriormente. Si esto no le resulta familiar, le recomiendo que se ponga al día allí primero.

El cuello de botella de Logit

Para comenzar, pongamos algunos números más en el cuello de botella logit. Consideramos una matriz de entrada X con forma [NxD], una matriz de pesos W con forma [DxV] y una matriz logit Y=X@W con forma [NxV]. En el contexto de un LLM, N sería la longitud de la secuencia multiplicada por el tamaño del lote (es decir, el número total de tokens en el lote), D el tamaño del estado oculto y V el tamaño del vocabulario.

Para un modelo Llama3 8B, tendríamos una ventana de contexto de 8192 tokens, un estado oculto con 4096 dimensiones y un tamaño de vocabulario de 128,256 tokens. Usando un tamaño de lote modesto de 8, obtenemos N = 8192×8 = 65,536.

Esto da como resultado que la matriz Y tenga la forma [NxV]=[65,536×128,256], o aproximadamente 8,4 mil millones de elementos. En bfloat16, esto ocuparía 16,8 GB de memoria. Sin embargo, si seguimos las mejores prácticas y utilizamos float32 para el cálculo de pérdidas para garantizar la estabilidad numérica, los requisitos se duplican a 33,6 GB.

Para poner este número en perspectiva, también necesitaríamos alrededor de 16 GB de memoria para mantener los pesos de Llama3 8B en la memoria en bfloat16. En la mayoría de las GPU, esto no deja espacio para la sobrecarga masiva de los estados del optimizador (por ejemplo, los momentos de Adam) y otras activaciones, lo que resulta en el infame error OOM de PyTorch.

Representación de las matrices de entrada, peso y logit junto con su huella de memoria. (Todas las ilustraciones y animaciones de este artículo fueron realizadas por el autor a menos que se especifique lo contrario)

Generalmente, este problema se soluciona utilizando:

Acumulación de gradientes: utilice un tamaño de lote más pequeño y acumule gradientes en varios lotes entre cada paso del optimizador, emulando un tamaño de lote más grande y manteniendo menos datos en la memoria. Puntos de control de activación: PyTorch almacena todas las activaciones intermedias para su reutilización en el paso hacia atrás, el punto de control borra estas activaciones y las recalcula sobre la marcha durante el paso hacia atrás. Esto conduce a un gran ahorro de memoria pero aumenta el tiempo de entrenamiento ya que se duplica el número de pases hacia adelante requeridos. Microprocesamiento de la pérdida: en lugar de calcular la pérdida en la dimensión N a la vez, podemos dividirla y acumular la pérdida en fragmentos más pequeños con tamaño n < N. Ahora, solo mantenemos una porción de tamaño [n, V] en la memoria a la vez. Entrenamiento de precisión mixta: el uso de media precisión durante el entrenamiento proporciona una reducción de memoria 2 veces mayor y aceleraciones significativas en Tensor Cores.

Si bien estas soluciones parecen atractivas, todas tienen inconvenientes importantes: la acumulación de gradientes y los puntos de control de activación ralentizan el entrenamiento, la precisión mixta puede ser inestable y el microprocesamiento por lotes requiere una iteración (lenta) a nivel de PyTorch y, aunque se elige que n sea más pequeño que N, el tamaño del vocabulario sigue siendo enorme en comparación.

Más importante aún, estas soluciones no abordan el problema que hemos abordado repetidamente a lo largo de esta serie: el movimiento de datos. De hecho, todavía estamos perdiendo el tiempo escribiendo miles de millones de logits en VRAM solo para leerlos milisegundos después.

La solución del núcleo

Como veremos en un minuto, el paso hacia adelante y hacia atrás de la pérdida de entropía cruzada involucra productos escalares, multiplicación de matrices y un softmax. Como aprendimos en esta serie, todas estas son operaciones que se pueden organizar en mosaico de manera eficiente. En otras palabras, podemos realizarlos de forma iterativa manteniendo solo una pequeña parte de las entradas en la memoria en cualquier momento.

Además, la entropía cruzada generalmente va precedida de una multiplicación de matrices: la proyección lineal desde el estado oculto al espacio de vocabulario. Esta es una gran oportunidad para la fusión de operadores: fusionar múltiples operaciones dentro de un solo núcleo, lo que genera grandes aceleraciones y posibles ganancias de memoria.

En las siguientes secciones, veremos cómo derivar y fusionar eficientemente los pasos hacia adelante y hacia atrás a través de un núcleo que combina una capa lineal con entropía cruzada.

Como se mencionó en el último artículo, los núcleos de Triton no se registran de forma nativa en el autogrado de PyTorch. Por lo tanto, necesitamos derivar el gradiente nosotros mismos, una ocasión maravillosa para repasar algo de cálculo 😉

Las matemáticas detrás de la entropía cruzada lineal fusionada

Definición y pase hacia adelante

En esta sección, derivamos la expresión matemática de nuestra capa de entropía cruzada lineal fusionada para ver cómo se presta naturalmente al mosaico.

Para dos distribuciones de probabilidad discretas p y q, la entropía cruzada se define como:

En nuestro contexto, p es el vector único que representa el token objetivo, mientras que q es la distribución del modelo sobre el vocabulario. Obtenemos q aplicando un softmax a los logits l, que son en sí mismos las salidas de la capa lineal anterior.

Dado que p es positivo para un único token objetivo y, la suma colapsa. Luego podemos sustituir el softmax numéricamente estable (como se discutió en el último artículo) para derivar la expresión final:

Sustituyendo los logits l con la capa lineal x. w, vemos que el pase hacia adelante se reduce a tres cantidades principales:

El logit objetivo x . w_y. El log-sum-exp (LSE) de todos los productos escalares. El logit máximo global utilizado para la estabilidad numérica.

Gracias al algoritmo softmax en línea, podemos calcular estas cantidades sin materializar el vocabulario completo en la memoria. En lugar de un cuello de botella de memoria O(V), iteramos sobre la dimensión oculta D y el vocabulario V en mosaicos pequeños (D_block y V_block). Esto transforma el cálculo en un problema de registro O(1).

Para paralelizar esto de manera efectiva, lanzamos un programa GPU por fila de la matriz de entrada. Cada programa ejecuta de forma independiente los siguientes pasos:

Calcule previamente el logit objetivo: realice un producto escalar en mosaico entre la fila actual de X y la columna de W asociada con el token Y. Reducción en línea: itere a través de los bloques ocultos y de vocabulario para:
1. Seguimiento del máximo de carrera (m)
2. Actualice la suma acumulada de exponenciales (d) usando la fórmula softmax en línea:

Un ejemplo de multiplicación de matrices en mosaico para un único programa de GPU que procesa una fila de X. Los cuadrados de colores representan elementos cargados en la memoria y el contorno de color representa el mosaico completo sobre el que se itera. El mosaico sacrifica velocidad por ganancias masivas de memoria.

Ahora que comprendemos mejor el pase hacia adelante, echemos un vistazo a la derivación del pase hacia atrás.

Pase hacia atrás

Notación

Para derivar nuestros gradientes de manera eficiente, usaremos la notación de Einstein y el delta de Kronecker.

En la notación de Einstein, los índices repetidos se suman implícitamente. Por ejemplo, una multiplicación de matrices estándar Y = X@W se simplifica desde una suma detallada hasta un emparejamiento de índices limpio:

El delta de Kronecker (δ_ij) se utiliza junto con esta notación para manejar la lógica de identidad. Es igual a 1 si i=j y 0 en caso contrario. Como veremos, esto es particularmente útil para colapsar índices durante la diferenciación.

Multiplicación de matrices

En esta sección, derivamos los gradientes retropropagados para la multiplicación de matrices. Suponemos la existencia de un gradiente aguas arriba ℓ.

Para determinar cómo se propaga hacia atrás mediante la multiplicación de matrices, aplicamos la regla de la cadena a las entradas x y la matriz de peso w. Aquí y representa los resultados de la multiplicación:

Comenzamos derivando las derivadas parciales de y con respecto a x, siguiendo estos pasos:

Exprese y en términos de x y w. Observa que w es una constante con respecto a la derivada de x, por lo que podemos sacarla de la derivada. Expresar el hecho de que la derivada parcial de x_ik con respecto a x_mn es 1 sólo cuando i=m y k=n usando el delta de Kronecker. Observe que ẟ_kn impone k=n, por lo tanto w_kj * ẟ_kn se reduce a w_nj.

Luego, consideramos la expresión completa y obtenemos el gradiente. Derivamos el último paso notando una vez más que 1/y_ij * ẟ_im se reduce a 1/y_mj.

Sin embargo, la notación matricial está conceptualmente más cerca de nuestro núcleo Triton, por lo tanto, reescribimos esta expresión como una multiplicación de matrices usando la identidad X_ij = [X^T]_ji:

Seguimos exactamente los mismos pasos para derivar el gradiente con respecto a W:

Luego, el gradiente retropropagado es el siguiente:

Lo cual equivale a la notación matricial:

Entropía cruzada

En esta sección, nos centraremos en la entropía cruzada aplicada a distribuciones de probabilidad discretas. Considerando un tensor de j logits, con una etiqueta y, la entropía cruzada se calcula de la siguiente manera:

Donde x_y corresponde al logit asociado a la etiqueta.
Una vez más, nos interesa la derivada parcial de cualquier salida i con respecto a cualquier entrada k. Debido al factor de normalización, cada elemento i afecta el valor de todos los demás elementos, por lo tanto, la derivada parcial se obtiene definiendo la función por partes dependiendo del valor de i:

Sumando ambos casos obtenemos el gradiente:

Y en notación matricial:

Donde y_{one hot} es un vector de ceros con la entrada correspondiente a la etiqueta establecida en uno. Este resultado nos dice que el gradiente es simplemente la diferencia entre la predicción y la verdad fundamental.

Entropía cruzada lineal fusionada

Combinando la proyección lineal con la entropía cruzada en una sola expresión, obtenemos:

Gracias a la regla de la cadena, derivar el gradiente de esta expresión se reduce a multiplicar los gradientes que calculamos anteriormente:

Donde xey se refieren a las entradas y salidas de la capa lineal respectivamente y w a la matriz de peso asociada.

Nota: en una configuración por lotes, necesitaremos reducir los gradientes W sobre la dimensión del lote. Generalmente utilizamos una suma o una reducción media.

Implementación del núcleo

Con la teoría establecida, podemos implementar el núcleo fusionado en Triton. Dado que la entropía cruzada suele ser la capa final de un modelo de lenguaje, podemos combinar los pasos hacia adelante y hacia atrás en un solo núcleo. Esta fusión ofrece dos ventajas: minimiza la sobrecarga de múltiples lanzamientos del kernel y mejora significativamente la localidad de los datos al mantener valores intermedios en el chip.

Analizaremos el núcleo paso a paso desde la perspectiva de una única instancia de programa que, en nuestra estrategia de paralelización, maneja una fila específica de la matriz de entrada.

1. Configuración y cálculo previo del Logit objetivo

La fase inicial implica la configuración estándar de Triton:

Identificación del programa: utilizamos tl.program_id para determinar de qué fila de la matriz de entrada es responsable el programa actual. Inicialización de parámetros: definimos mosaicos usando D_BLOCK y V_BLOCK e inicializamos el máximo de ejecución (m) y la suma (d) necesarios para el algoritmo softmax en línea. Aritmética de punteros: calculamos las direcciones de memoria base para nuestros tensores. Los punteros para X (entrada) y dX (gradiente) se compensan utilizando el paso de fila para que cada programa acceda a su vector de token único. Por el contrario, el puntero W (peso) permanece en la dirección base porque cada programa eventualmente debe recorrer todo el espacio de vocabulario. Enmascaramiento y salida anticipada: definimos un ignore_index (el valor predeterminado es -100). Si un programa encuentra esta etiqueta (por ejemplo, para tokens de relleno), finaliza antes de tiempo con una pérdida de 0 para guardar ciclos.

2. Calcular el Logit objetivo

Antes del bucle principal, debemos aislar el logit objetivo x. w_y. Iteramos sobre la dimensión oculta D en fragmentos D_BLOCK, realizando un producto escalar entre la fila de entrada X y la columna específica de W correspondiente a la etiqueta de verdad fundamental Y.

Debido a que W es una matriz 2D, calcular los punteros para estos mosaicos de columnas específicos requiere una manipulación precisa de la zancada. La siguiente ilustración ayuda a visualizar cómo "saltamos" a través de la memoria para extraer solo los pesos necesarios para el token objetivo.

Representación de la aritmética de punteros ejecutada para calcular el logit objetivo Y. Aquí, consideramos que la etiqueta es 4, lo que significa que el logit objetivo es el producto escalar de X con la quinta columna de W. Los vectores de diferentes colores representan diferentes pasos de la iteración a lo largo de D (es decir, diferentes valores de d_idx). Los números se refieren a la dirección de memoria de cada elemento asumiendo un diseño de fila principal.

Una vez cargados los mosaicos, los convertimos en float32 para garantizar la estabilidad numérica y agregamos su producto escalar a una variable acumulada antes de pasar a la siguiente iteración.

Aquí está el código hasta el momento:

A continuación, ejecutamos el pase directo, que procesa el espacio de vocabulario en dos etapas anidadas:

Cálculo Logit en mosaico: Calculamos los logits para un V_BLOCK a la vez. Esto se logra iterando sobre la dimensión de vocabulario V (bucle externo) y la dimensión oculta D (bucle interno). Dentro del bucle interno, cargamos un mosaico de X y un bloque de W, acumulando sus productos escalares parciales en un registro de alta precisión. Actualización de Softmax en línea: una vez finalizado el producto escalar completo para un mosaico logit, no lo almacenamos en VRAM. En su lugar, actualizamos inmediatamente nuestras estadísticas en ejecución: el valor máximo m y la suma acumulada de exponenciales d usando la fórmula softmax en línea. Al hacer esto "sobre la marcha", nos aseguramos de que solo mantengamos un pequeño V_BLOCK de logits en los registros de la GPU en un momento dado.

Después de estas iteraciones, los valores finales de myd se utilizan para reconstruir el LSE. Luego, la pérdida escalar final para la fila se calcula restando el logit objetivo (x. w_y) de este valor LSE.

Aquí hay una representación visual del pase hacia adelante:

Representación visual de la multiplicación de matrices en mosaico con actualizaciones de estadísticas en ejecución. En cada paso, cargamos elementos coloreados en verde o azul oscuro y calculamos los productos escalares de los vectores resaltados en verde. Los elementos de Y se acumulan iterando sobre la dimensión D, cuando se hace esto (es decir, las celdas son verdes), actualizamos myd según el mosaico recién calculado.

Aquí está el código para el pase hacia adelante:

Ahora llegamos a la última parte del núcleo: el pase hacia atrás. Nuestro objetivo es calcular los gradientes con respecto a X y W usando la expresión que derivamos anteriormente:

Para mantener la eficiencia de la memoria, una vez más procesamos el vocabulario en mosaicos utilizando un enfoque de dos etapas:

Recalcular las probabilidades normalizadas (P): debido a que no almacenamos la matriz logit completa durante el pase directo, debemos recalcular las activaciones para cada mosaico. Al reutilizar el Log-Sum-Exp calculado en el pase hacia adelante, podemos normalizar estas activaciones sobre la marcha. Restar la etiqueta de verdad fundamental Y del logit objetivo dentro de este mosaico nos da una porción local del logit de gradiente, P.
2. Acumulación de gradientes: Con una ficha de P en mano, calculamos los gradientes parciales. Para dX, realizamos un producto escalar con bloques de W^T; para dW, multiplicamos por mosaicos de X^T. Para agregar de forma segura estos valores en todo el lote, utilizamos tl.atomic_add de Triton.
Esta operación actúa como un += seguro para subprocesos, lo que garantiza que diferentes programas que actualizan el mismo gradiente de peso no se sobrescriban entre sí.

Aquí hay algunos detalles adicionales sobre la implementación:

The Stride Swap: al calcular P . W_T, en realidad no necesitamos transponer físicamente la enorme matriz W en la memoria. En cambio, invertimos las formas y los pasos en el puntero del bloque de W para leer las filas de W como columnas de W^T. Esto da como resultado una transposición "gratuita" que ahorra tiempo y VRAM. Precisión numérica: vale la pena señalar que, si bien X y W pueden estar en bfloat16, la acumulación de dW y dX mediante atomic_add generalmente se realiza en float32 para evitar la acumulación de pequeños errores de redondeo en miles de filas. Nota de contención: si bien atomic_add es necesario para dW (porque cada programa actualiza los mismos pesos), dX es privado para cada programa, lo que significa que no hay contención entre los ID de programa para ese tensor específico. Enmascaramiento de adición atómica: atomic_add no admite punteros de bloque. Por lo tanto, implementamos explícitamente la lógica de puntero y máscara para dW.

La siguiente figura es una representación del paso hacia atrás para una iteración del bucle externo (es decir, un bloque a lo largo de V y todos los bloques a lo largo de D):

Representación del paso hacia atrás para un solo paso a lo largo de la dimensión V y una iteración completa a lo largo de la dimensión D. En la etapa 4, resaltamos cómo dX se acumula a lo largo de iteraciones (cada programa actualiza su fila privada una vez por paso a lo largo de V) mientras que dW se acumula a lo largo de programas (N programas actualizan los valores de un solo bloque en dW en cada paso a lo largo de V).

Aquí está el código completo para el pase hacia atrás:

¡Esto concluye la implementación de nuestro kernel! El código completo, incluido el kernel y el script de referencia, está disponible aquí.

Punto de referencia de memoria

Finalmente, comparamos nuestro kernel con la línea base de PyTorch usando hiperparámetros inspirados en Llama3 y una GPU A100. Específicamente, consideramos una longitud de secuencia de S = 16,384, un tamaño de lote de B = 1 y una dimensión de incrustación de D = 4096; el tamaño del vocabulario se establece en V=128,256.

Como se esperaba, la línea base de PyTorch asigna un tensor intermedio masivo para almacenar las activaciones, lo que resulta en un uso máximo de memoria de 36,02 GB. En comparación, nuestro kernel Triton reduce el uso máximo de memoria en un 84 % al asignar solo 5,04 GB usando D_BLOCK=64 y V_BLOCK=64.

El uso de tamaños de bloques aún más pequeños permitiría mayores ganancias de memoria a costa de la eficiencia.

Limitaciones atómicas y escalamiento de la producción

En este artículo, nos centramos en la intuición técnica y matemática detrás de los núcleos fusionados de entropía cruzada lineal. Usamos operaciones atómicas como tl.atomic_add para mantener el código mínimo y legible. Sin embargo, si bien nuestro kernel redujo con éxito el uso de memoria en un asombroso 86%, el kernel Triton es significativamente más lento que el PyTorch nativo.

Desafortunadamente, las mismas operaciones atómicas que hacen que este núcleo sea más fácil de escribir y comprender tienen el costo de un atasco de tráfico masivo, ya que miles de subprocesos intentan modificar la misma dirección de memoria a la vez. Generalmente, tl.atomic_add tiene buen rendimiento cuando la contención es baja. En nuestra implementación actual, tenemos:

Alta contención: para el gradiente de peso, todos los programas del lote (hasta 16,384 en nuestra prueba) intentan actualizar los mismos mosaicos de memoria simultáneamente. El hardware debe serializar estas actualizaciones, lo que obliga a miles de subprocesos a esperar en fila. No asociatividad numérica: en las computadoras, la suma de punto flotante no es asociativa. Los errores de redondeo pueden acumularse de manera diferente según el orden de las operaciones, por lo que las pruebas de corrección pueden pasar en un T4 pero fallar en un A100; este último tiene más multiprocesadores (SM) de transmisión que realizan más adiciones concurrentes y no deterministas.

Nota sobre la precisión: en Ampere y arquitecturas más nuevas, el formato TF32 puede contribuir aún más a estas discrepancias. Para una paridad numérica estricta, se debe establecer enable_tf32=False o utilizar tipos de mayor precisión durante los pasos de acumulación.

Camino a la producción

Para ir más allá de esta implementación educativa y hacia un kernel listo para producción (recomiendo mirar la implementación Liger-Kernel), se podrían implementar varias optimizaciones:

Reemplazo de dX Atomics: dado que cada programa "es dueño" de su fila de X, podemos usar una acumulación de registros simple seguida de un tl.store, eliminando por completo los átomos para los gradientes de entrada. Un kernel dW dedicado: para optimizar el cálculo de dW, los kernels de producción generalmente usan una estrategia de cuadrícula diferente donde cada programa maneja un bloque de W e itera a través de la dimensión del lote, acumulando gradientes localmente antes de una única escritura global. Microprocesamiento por lotes: las implementaciones avanzadas, como las de la biblioteca Liger-Kernel, procesan la secuencia por bloques a lo largo de la dimensión N, haciendo que el escalado de la memoria sea constante en la longitud de la secuencia en lugar de lineal. Esto permite utilizar tamaños de lote mucho mayores con un coste de memoria reducido.

Conclusión

Esto concluye nuestra inmersión profunda en los núcleos de entropía cruzada lineal fusionados. Gracias por leer hasta el final y espero que este artículo le haya brindado la intuición y la comprensión práctica necesarias para aprovechar estas ideas y explorarlas más a fondo.

Si esto le resultó útil, considere compartir el artículo; Realmente ayuda a respaldar el tiempo y el esfuerzo que se dedican a producir este trabajo. Y como siempre, no dude en ponerse en contacto conmigo si tiene preguntas, pensamientos o ideas para seguimiento.

¡Hasta la próxima! 👋

Fuentes

Presentamos Meta Llama 3: el LLM disponible abiertamente más capaz hasta la fecha LigerKernel (conferencia) Implementación de entropía cruzada lineal de LigerKernel Implementación de Unsloth (solo entropía cruzada)