Cómo abaratar la distilación de conocimiento para modelos LLM a gran escala
La distilación de conocimiento permite comprimir grandes LLMs, pero el paso de distilación suele ser el más costoso en VRAM y cómputo. Una combinación de cacheo offline de logits top-K y una versión 'chunked' y fusionada de la pérdida KL reduce el uso de memoria lo suficiente como para hacer experimentación a gran escala práctica.
Qué problema resuelve la técnica
La distilación de conocimiento —entrenar un modelo ‘student’ más pequeño para igualar el comportamiento de un ‘teacher’ grande— es una práctica común para comprimir grandes modelos de lenguaje. Con la ola reciente de LLMs open source como gpt-oss, Qwen, GLM o Kimi, la necesidad de convertir modelos enormes en versiones más manejables se volvió crítica: por ejemplo, el modelo Kimi-K3 tiene 2.8 billones de parámetros y requiere aproximadamente 3 TB de VRAM solo para cargarse. Mantener ambos modelos en memoria y calcular distribuciones sobre todo el vocabulario por cada token hace que la etapa de distilación sea extraordinariamente cara.
Este costo operativo no solo encarece despliegues en la nube, sino que también limita la capacidad de equipos e instituciones (como universidades o startups en América Latina) para experimentar y producir versiones comprimidas de alta calidad. El artículo original de Hugging Face plantea dos cambios de sistema que, combinados, reducen significativamente la VRAM necesaria y hacen viable la distilación a escala con recursos mucho más modestos.
La raíz del problema: la pérdida KL densa y la recomputación del teacher
El enfoque más utilizado en distilación es la distilación online con la pérdida de Kullback–Leibler (KL). En cada paso de entrenamiento, el teacher produce una distribución de probabilidad completa sobre el vocabulario, y el student aprende a empatarla. Esto es expresivo pero costoso: requiere guardar tensores de tamaño vocabulario × longitud-de-secuencia para cada posición y mantener al teacher activo en memoria durante todo el entrenamiento.
Para dar una idea práctica: gpt-oss-120b tiene un vocabulario de 201,088 tokens. Con una longitud de secuencia de 32K y batch size 4, el tensor de probabilidades del teacher solo tendría la forma 4 × 201,088 × 32,768. En bfloat16 eso representa cerca de 50 GB de VRAM para un solo tensor. Añadan activaciones, gradientes, pesos y estados del optimizador, y una iteración de distilación puede pico en alrededor de 250 GB de VRAM, por encima de la capacidad de GPUs como H200 o B200.
Dos cambios de sistema: cacheo offline de top-K y pérdida KL ‘chunked’ fusionada
La solución propuesta actúa en dos frentes complementarios:
-
Cacheo offline de logits top-K En lugar de recalcular el teacher en cada paso, se ejecuta el teacher una vez y se guarda una cache con los top-100 logits más probables por posición. Durante el entrenamiento, el student aprende a coincidir con esa cache. Con esto, el teacher no necesita permanecer cargado en memoria, y la misma cache puede reutilizarse para múltiples experimentos o ablations, reduciendo costos de cómputo repetido.
-
Una pérdida KL ‘chunked’ y fusionada El segundo cambio aborda el problema de memoria que surge al construir la matriz vocabulario × longitud-de-secuencia para evaluar la discrepancia entre teacher y student. El método estándar (Dense KL) reconstruye una distribución densa del teacher desde la cache y compara matriz completa contra la salida densa del student, lo que exige grandes picos de memoria.
Se proponen tres alternativas equivalentes matemáticamente:
-
Dense KL: la versión tradicional que materializa toda la cuadrícula densa y sirve como referencia de corrección, pero consume memoria en exceso.
-
Forward-chunked KL: mantiene al teacher en forma dispersa (solo los top-100 logits) y calcula la pérdida en trozos de la secuencia, procesando una porción a la vez. Esto elimina la necesidad de expandir el teacher a denso y resulta ser la más rápida en los benchmarks, pero aún requiere que los logits completos del student se computen y retengan para el backward pass, por lo que la memoria crece con la longitud de la secuencia.
-
Fused chunked KL: la contribución principal. Aquí se fusiona la proyección final del modelo (de estados ocultos a logits) directamente dentro del cálculo de la pérdida y se procesa la secuencia por trozos. El sistema proyecta los estados ocultos a logits solo para el trozo actual, incorpora ese resultado en la pérdida y descarta el logits antes de avanzar. En el backward pass, cada trozo se recomputa en lugar de almacenarse. Esto implica hacer la proyección dos veces (una en forward y otra en backward), pero evita la formación de la cuadrícula completa vocabulario × secuencia y reduce el pico de memoria a crecimiento lineal con la longitud en lugar de explotar por el producto de vocabulario y secuencia.
Resultados prácticos en VRAM
En los experimentos presentados, la versión Dense KL puede alcanzar picos cercanos a 250 GB —por encima de la capacidad de una H200 de 141 GB—. La pérdida fused chunked evita ese pico y alcanza un máximo aproximado de 128 GB. Estos números ilustran cómo reformular la computación y el flujo de datos puede hacer posible entrenar con contextos largos y realizar experimentación extensa sin requerir clusters masivos de GPUs.
Implicaciones para equipos y decisiones en América Latina
Para equipos de investigación, empresas emergentes y áreas de I+D en la región, estas técnicas tienen varias ventajas prácticas:
- Reducción de costos: al necesitar menos VRAM y evitar la recomputación constante del teacher, se disminuye el gasto en infraestructura y tiempo de cómputo en la nube.
- Mayor accesibilidad: procesos que antes requerían cientos de GPUs ahora pueden realizarse con hardware más modesto, lo que facilita la experimentación local y la adaptación de LLMs a datos y requisitos regionales.
- Reutilización de cache: el cache offline de logits top-K permite reproducibilidad y comparaciones rápidas entre variantes de entrenamiento sin volver a ejecutar el teacher.
Estas mejoras no eliminan la necesidad de buen diseño de pipelines ni de validación rigurosa, pero sí amplían la capacidad práctica de equipos con recursos limitados para contribuir y adaptar modelos grandes.
Consideraciones técnicas y trade-offs
El enfoque fused chunked introduce un trade-off clásico: recomputación por memoria. Proyectar dos veces los mismos trozos añade costo computacional, pero el beneficio en ahorro de memoria suele compensarlo si el cuello de botella es VRAM. Además, usar un cache top-K (aquí, top-100) hace implícitas ciertas aproximaciones: la distribución completa del teacher se aproxima mediante sus logits más probables. En los resultados reportados, estas aproximaciones permiten mantener calidad alta mientras reducen el costo.
Asimismo, Forward-chunked resultó ser la opción más rápida en los benchmarks aunque no reduzca tanto la memoria como la versión fusionada. Seleccionar la variante adecuada depende de la prioridad del proyecto: latencia de entrenamiento vs. uso máximo de memoria.
Conclusión
Reformular cómo se almacena y calcula la señal de supervisión en distilación de LLM permite una reducción sustancial del uso de VRAM. El cacheo offline de los top-100 logits del teacher y la pérdida KL procesada en trozos con fusión de la proyección final son cambios simples a nivel de sistema que, juntos, convierten una operación costosa en algo viable para experimentación a gran escala y para equipos con recursos más modestos. Para la región latinoamericana, esto representa una oportunidad concreta para participar en la adaptación y despliegue de LLMs sin depender exclusivamente de clusters masivos en la nube.
Si su organización planea comprimir o adaptar LLMs, evaluar estas técnicas en su pipeline puede ser el primer paso para reducir costos y acelerar ciclos de investigación y producción.
Fuente original: Hugging Face Blog