¿Cómo haces que un modelo de inteligencia artificial gigante enseñe lo que sabe sin necesitar una sala llena de GPUs? Una investigación de Multiverse Computing propone una respuesta práctica: guardar solo las predicciones más importantes del modelo grande y calcular la función de entrenamiento por partes, sin construir matrices enormes en memoria.
El resultado es una reducción considerable del uso de VRAM, especialmente cuando se trabaja con contextos largos. En una demostración con un modelo GPT-OSS de 20.000 millones de parámetros, el sistema pasó de necesitar cuatro nodos de GPU a funcionar con uno, mientras el tiempo por paso bajó de 57 a 12,23 segundos.
Por qué la destilación de modelos es tan costosa
La destilación de conocimiento consiste en entrenar un modelo pequeño, llamado estudiante, para que imite las respuestas de otro más grande, conocido como profesor. La idea es sencilla: en lugar de ejecutar siempre el modelo gigante, transfieres parte de sus capacidades a una versión más barata y rápida.
Pero el proceso tradicional tiene un problema. Durante cada paso de entrenamiento deben permanecer cargados tanto el modelo profesor como el estudiante. Además, el profesor produce una distribución de probabilidad sobre todo su vocabulario para cada token de la entrada.
¿Parece exagerado? Considera este caso: el modelo gpt-oss-120b utiliza un vocabulario de 201.088 tokens. Con secuencias de 32.768 tokens y un lote de cuatro ejemplos, una sola matriz de probabilidades del profesor puede ocupar cerca de 50 GB en bfloat16.
Cuando se suman los pesos, las activaciones, los gradientes y los estados del optimizador, una iteración puede alcanzar aproximadamente 250 GB de VRAM. Eso supera la capacidad de una GPU H200 de 141 GB y obliga a distribuir el entrenamiento entre varios dispositivos.
El desafío no está únicamente en cargar el modelo grande. También está en almacenar todas las predicciones que produce para cada posición de una secuencia extensa.
La propuesta: separar al profesor del entrenamiento
La primera modificación consiste en cambiar la destilación en línea por una estrategia fuera de línea, conocida como offline distillation.
En lugar de mantener al profesor activo durante todo el entrenamiento, el sistema lo ejecuta una sola vez. Para cada token guarda únicamente sus top-k predicciones, es decir, las opciones con mayor probabilidad. En el experimento principal se almacenan las 100 más probables por posición.
Después, el estudiante se entrena usando ese archivo de predicciones. El profesor ya no necesita ocupar memoria ni volver a calcular sus resultados, y la misma caché puede reutilizarse en múltiples experimentos.
Guardar solo las 100 mejores opciones parece una simplificación importante frente a un vocabulario de más de 200.000 tokens. Sin embargo, esas predicciones concentran la mayor parte de la información útil para que el estudiante aprenda el comportamiento del profesor.
Una función KL que no construye toda la matriz
La segunda modificación afecta la función de pérdida basada en la divergencia de Kullback-Leibler, o pérdida KL. Esta función mide cuánto se aleja la distribución de predicciones del estudiante de la del profesor.
El método tradicional crea una matriz completa con una dimensión para cada token del vocabulario y otra para cada posición de la secuencia. Con contextos de 32K, 64K o incluso 256K tokens, esa matriz crece rápidamente y puede provocar un pico de memoria imposible de manejar.
El equipo propone una implementación que procesa la secuencia en fragmentos. En vez de crear y conservar toda la matriz, calcula una parte, incorpora su resultado a la pérdida acumulada y descarta ese fragmento antes de continuar.
Tres formas de calcular la pérdida
El trabajo compara tres alternativas para la destilación fuera de línea:
- KL densa: reconstruye una distribución completa del profesor a partir de los
top-100y la compara con las probabilidades completas del estudiante. Es la referencia más cercana al método tradicional, pero consume mucha memoria. - KL dividida hacia adelante: mantiene las predicciones del profesor en formato disperso y procesa la secuencia por partes. Reduce el consumo, aunque todavía conserva en memoria los logits completos del estudiante.
- KL fusionada y dividida: integra la proyección de salida del estudiante directamente dentro del cálculo de la pérdida. Nunca materializa todos sus logits al mismo tiempo y vuelve a calcular cada fragmento durante la retropropagación para evitar almacenarlo.
Esta última versión realiza parte del trabajo dos veces, una durante la propagación hacia adelante y otra durante la retropropagación. A cambio, reduce de forma drástica el pico de memoria cuando aumenta la longitud del contexto.
Menos memoria sin perder calidad de entrenamiento
En una comparación realizada con una GPU H200, un modelo Llama 3.1 8B Instruct como profesor y un estudiante Llama de 3.2B parámetros, las cuatro configuraciones alcanzaron curvas de pérdida prácticamente idénticas con un contexto de 8K tokens.
| Método | Memoria máxima | Tiempo por iteración | Rendimiento |
|---|---|---|---|
| Destilación en línea | 102,8 GB | 25,9 s | 237 TFLOP/s |
| Offline con KL densa | 78,3 GB | 18,5 s | 331 TFLOP/s |
| Offline con KL dividida | 61,8 GB | 18,4 s | 335 TFLOP/s |
| Offline con KL fusionada y dividida | 58,3 GB | 20,2 s | 304 TFLOP/s |
En contextos de 8K, la variante fusionada no es la más rápida. Su ventaja aparece cuando las secuencias se hacen mucho más largas.
En una prueba aislada con una red de proyección de salida, el consumo a 32K tokens bajó de 85,2 GiB con la pérdida densa a 5,45 GiB con la versión completamente dividida. Eso representa una reducción de 15,6 veces. Además, el método denso dejó de funcionar a partir de 64K tokens por falta de memoria.
A 256K tokens, la versión fusionada utilizó 11,6 GiB, frente a 134,2 GiB de la siguiente alternativa más eficiente. En esa longitud también fue aproximadamente 3,3 veces más rápida por iteración.
La diferencia se nota en modelos grandes
El impacto más concreto apareció al destilar un modelo GPT-OSS de 20.000 millones de parámetros con un contexto de 32.768 tokens. La reducción de memoria permitió pasar de cuatro nodos de GPU a uno solo.
El tiempo por paso cayó de 57 a 12,23 segundos, cerca de cinco veces menos. Al mismo tiempo, el rendimiento por GPU aumentó de 74,2 a 345,7 TFLOP/s.
Estos números importan porque la destilación no suele ser un experimento único. Los equipos necesitan probar distintos datos, funciones de pérdida, longitudes de contexto y configuraciones de entrenamiento. Si cada intento requiere cientos de gigabytes de VRAM, la investigación se vuelve lenta y costosa.
Con una implementación más eficiente, comprimir modelos puede convertirse en un proceso iterativo accesible para más equipos, no solo para laboratorios con grandes clústeres.
Modelos pequeños con buena parte de las capacidades
En el experimento de recuperación, un estudiante de aproximadamente 3.200 millones de parámetros fue destilado a partir de un Llama 3.1 8B Instruct. El modelo reducido conservó la mayor parte de la precisión del profesor en BoolQ y HellaSwag.
En MMLU quedó a unos nueve puntos del modelo original, pese a tener menos de la mitad de sus parámetros. Esto no significa que la compresión sea gratuita: el estudiante pierde parte de la capacidad, especialmente en tareas complejas. Pero sí muestra que una reducción importante puede mantener un rendimiento útil.
La técnica también puede ayudar en lo que el equipo llama healing, o recuperación de capacidades después de comprimir un modelo. ¿El objetivo? Obtener sistemas más pequeños, baratos de ejecutar y capaces de conservar habilidades que normalmente se perderían durante la reducción.
Código abierto para probar la técnica
Multiverse Computing publicó la implementación de la pérdida KL fusionada y dividida en GitHub: Full-Chunked-KL-Loss.
La investigación también analiza otros factores, como la función de pérdida utilizada y la forma de empaquetar secuencias durante el entrenamiento. Para quienes trabajan con modelos grandes, el aporte principal no es una nueva arquitectura, sino una manera más inteligente de administrar la memoria.
La IA generativa suele parecer una carrera por construir modelos cada vez más grandes. Sin embargo, hacerlos útiles en el mundo real también exige aprender a reducirlos. Si una técnica permite trasladar capacidades a modelos pequeños usando menos GPUs y contextos más largos, la innovación deja de depender tanto del tamaño del laboratorio que la desarrolla.
Fuente original
https://huggingface.co/blog/MultiverseComputingCAI/efficient-knowledge-distillation
