La formación previa de LLM a escala de frontera en el FP8 es ahora una práctica estándar, pero pasar al punto flotante de 4 bits sigue siendo un problema de investigación abierto porque los formatos más estrechos comprimen el rango dinámico y amplifican el error de cuantificación en horizontes de tokens largos. Una nueva investigación de NVIDIA describe una metodología de preentrenamiento basada en NVFP4, un formato de microescalado de 4 bits compatible de forma nativa con Blackwell Tensor Cores, y la valida entrenando previamente un Mamba-Transformer híbrido de 12 mil millones de parámetros en 10 billones de tokens. El equipo de investigación afirma que este es el entrenamiento más largo documentado públicamente con precisión de 4 bits hasta la fecha. El modelo resultante alcanza el 62,58% en MMLU-Pro 5-shot frente al 62,62% de la línea base FP8, y es compatible con Transformer Engine de NVIDIA.
Qué es realmente NVFP4
Para comprender por qué NVFP4 es importante, es útil revisar cómo funcionan los formatos de microescala. En un formato de microescala (MX), un bloque contiguo de elementos de baja precisión comparte un factor de escala único, que se utiliza para mapear el bloque nuevamente en un rango numérico más amplio durante la multiplicación de la matriz. MXFP4 utiliza bloques de 32 elementos donde cada elemento se almacena como E2M1 (1 bit de signo, 2 bits de exponente, 1 bit de mantisa) que codifica solo los valores ±0, ±0,5, ±1, ±1,5, ±2, ±3, ±4 y ±6. Los factores de escala de bloque se almacenan en UE8M0, lo que los restringe a potencias de dos.
NVFP4 cambia tres cosas. Primero, el tamaño del bloque cae de 32 a 16 elementos, reduciendo el rango dinámico que cada escala debe cubrir. En segundo lugar, los factores de escala de bloque se almacenan en E4M3 en lugar de UE8M0, intercambiando el rango de exponente por la precisión de mantisa, de modo que el amax (máximo absoluto) por bloque se pueda mapear mucho más cerca del máximo representable del FP4. En tercer lugar, NVFP4 agrega un segundo nivel de escala: una escala por tensor FP32 que reasigna los valores para que las escalas del bloque E4M3 permanezcan dentro del rango. El resultado es que al menos el 6,25% de los valores en cada bloque (el amax por bloque) están representados con una precisión cercana al FP8, mientras que el resto se ubica en el FP4.
En NVIDIA Blackwell, los GEMM del FP4 funcionan con un rendimiento de BF16 de 4× en GB200 y de 6× en GB300, lo que se traduce en aceleraciones de aproximadamente 2× y 3× con respecto a FP8. El uso de memoria de operandos se reduce aproximadamente a la mitad en comparación con el FP8.
Qué está cuantificado y qué no
Solo los GEMM dentro de las capas lineales (totalmente conectadas) Fprop, Dgrad y Wgrad realmente se ejecutan en NVFP4. Las incrustaciones, el cabezal de proyección de salida, las capas de normalización, las no linealidades y todos los componentes de atención (softmax y los GEMM por lotes de clave de consulta y valor de puntuación de atención) permanecen en BF16 o FP32. Los pesos del modelo, los gradientes de peso utilizados para la acumulación entre microlotes y réplicas de datos paralelos, y los estados del optimizador se mantienen en FP32. Las reducciones paralelas del tensor se ejecutan en BF16.
La metodología de formación de cuatro partes
La cuantización de cada GEMM de capa lineal a NVFP4 con la configuración predeterminada (escalado de bloque de 1 × 16 en todas partes, redondeo al más cercano, incluso en cada tensor, sin transformaciones) diverge al principio del entrenamiento. El enfoque de NVIDIA lo estabiliza con cuatro componentes y los estudios de ablación en el modelo 12B muestran que cada uno de ellos es necesario.
Alta precisión selectiva: las capas lineales en los dos primeros y los últimos ocho de los 62 bloques (aproximadamente el 16% de todas las capas lineales) se mantienen en BF16. Ablaciones indicó que los bloques finales son los sensibles porque requieren más rango dinámico que el que proporciona el 4PM; mantener sólo los últimos cuatro bloques en BF16 también fue suficiente para una convergencia estable.
Transformadas aleatorias de Hadamard (RHT): los valores atípicos en los gradientes de peso se distribuyen en una distribución aproximadamente gaussiana multiplicando los mosaicos de entrada con una matriz de Hadamard de 16 × 16 combinada con un vector de signo aleatorio de ±1. Debido a que las transformaciones ortogonales se cancelan dentro del producto escalar, no se necesita corrección matemática en el GEMM. El tamaño d=16 se eligió empíricamente: d=4 perjudica la convergencia, d=128 dio resultados similares. RHT se aplica solo a las entradas del GEMM de gradiente de peso (Wgrad), y se comparte un único vector de signo aleatorio en todas las capas lineales. La aleatorización en sí misma no fue operativa en la escala de 1.200 millones, pero mejoró considerablemente en la escala de 12.000 millones.
Escalado de bloques bidimensionales (2D) para pesos: el NVFP4 estándar escala bloques de 1×16 a lo largo de la dimensión del producto escalar. Debido a que el paso hacia atrás transpone el tensor de peso, los pases hacia adelante y hacia atrás terminan con pesos cuantificados diferentes, rompiendo la regla de la cadena. La solución de NVIDIA es escalar los pesos en bloques de 16×16 para que se utilice la misma representación cuantificada en ambas pasadas. Las activaciones y gradientes mantienen una escala de 1×16, ya que son menos sensibles a esta inconsistencia.
Redondeo estocástico en gradientes: el redondeo al par más cercano introduce un sesgo sistemático cuando se aplica a tensores de gradiente. El redondeo estocástico redondea probabilísticamente en función de la distancia a los dos valores representables más cercanos, eliminando ese sesgo. El equipo de investigación señala explícitamente en un artículo de investigación que el redondeo estocástico es perjudicial cuando se aplica a tensores de paso hacia adelante, por lo que se limita a gradientes.
Resultados del transformador Mamba híbrido 12B
El modelo 12B utiliza la arquitectura Nemotron-Nano-12B-v2-Base: 62 bloques (6 Self-Attention, 28 FFN, 28 Mamba-2), dimensión oculta 5120, dimensión FFN 20480, entrenados con un programa de calentamiento-estable-decaimiento (LR constante durante el 80% del entrenamiento, decaimiento en el 20% final), tamaño de lote 736, longitud de secuencia 8192. La línea base de referencia del FP8 sigue la metodología DeepSeek-V3: elementos E4M3, bloques de peso de 128 × 128, bloques de activación y gradiente de 1 × 128, con el primer bloque y los dos últimos bloques mantenidos en BF16.
La pérdida de validación de NVFP4 se mantiene dentro del 1% de la línea de base del FP8 durante la fase estable y se amplía a ligeramente por encima del 1,5% durante el decaimiento. La precisión descendente es comparable en la mayoría de los puntos de referencia: MMLU 76,57 % frente a 77,36 %, GSM8K CoT 92,27 % frente a 89,08 %, MATH 81,48 % frente a 83,32 %, AGIEval English CoT 70,31 % frente a 67,01 %. La codificación muestra la brecha más grande (HumanEval+ 57,43% frente a 59,93%, MBPP+ 55,91% frente a 59,11%) que el equipo de investigación atribuye en parte a la ruidosa evaluación del punto de control final. El equipo de investigación también documenta una técnica de cambio de precisión: la transición del pase directo de NVFP4 a BF16 a partir de 8,2 T de tokens (aproximadamente el 18 % del cronograma) redujo el error de pérdida relativa del 1,5 % al 0,5 %.
NVFP4 frente a MXFP4
En un Mamba-Transformer híbrido 8B separado entrenado en tokens 1T, NVFP4 alcanzó un error de pérdida relativa de aproximadamente el 1,5% frente a BF16, mientras que MXFP4 se mantuvo cerca del 2,5%. Para cerrar la brecha, MXFP4 requirió 1,36T de tokens para igualar la pérdida de 1T de NVFP4: una sobrecarga de token del 36%. El equipo de investigación atribuye la diferencia al tamaño de bloque más pequeño de NVFP4 y a las escalas E4M3, que conservan más rango dinámico de FP4 que las escalas UE8M0 de potencia de dos de MXFP4 (que pueden desperdiciar hasta una binada y las ±4, ±6 muestras en el peor de los casos).
Explicador visual de Marktechpost
MARKTECHPOST · Investigación en IA, explicada en profundidad.
Conclusiones clave
El equipo de investigación de NVIDIA preentrenó un Mamba-Transformer híbrido de 12B en tokens de 10T en NVFP4 (la ejecución de entrenamiento de 4 bits más larga documentada públicamente), igualando a FP8 en MMLU-Pro con un 62,58 % frente a un 62,62 %. NVFP4 utiliza bloques de 16 elementos con escalas E4M3 más una escala por tensor FP32, preservando las muestras ±4 y ±6 que el diseño UE8M0 de 32 elementos de MXFP4 puede perder debido al redondeo de potencia de dos. Se requieren cuatro técnicas para la convergencia; ninguna es opcional: ~16 % de capas lineales en BF16, transformaciones aleatorias de Hadamard de 16 × 16 en entradas Wgrad, escalado de peso 2D de 16 × 16 y redondeo estocástico solo en gradientes. Solo los GEMM de capa lineal se ejecutan en NVFP4: la atención, las incrustaciones, la normalización, las no linealidades, los pesos maestros, los gradientes y los estados del optimizador permanecen en BF16 o FP32. En un modelo 8B, MXFP4 necesitaba 1,36T de tokens (36% más) para igualar la pérdida de NVFP4 con 1T de tokens, mientras que los GEMM de FP4 ofrecen un rendimiento 2× de FP8 en GB200 y 3× en GB300.
Consulte el documento aquí. Además, no dude en seguirnos en Twitter y no olvide unirse a nuestro SubReddit de más de 150.000 ML y suscribirse a nuestro boletín. ¡Esperar! estas en telegrama? Ahora también puedes unirte a nosotros en Telegram.
¿Necesita asociarse con nosotros para promocionar su repositorio de GitHub O su página principal de Hugging O su lanzamiento de producto O seminario web, etc.? Conéctate con nosotros