La inferencia de Large Language Models (LLMs) en entornos de producción o locales presenta una variabilidad significativa en la salida de logits, incluso cuando se utilizan los mismos pesos de modelo. Esta divergencia no es un artefacto menor, sino una consecuencia directa de las complejidades inherentes a la pila de software y hardware, y las decisiones de cuantificación que se toman para optimizar el rendimiento. El problema fundamental de la computación que se aborda es la reproducibilidad y la fidelidad de la inferencia de modelos complejos en sistemas distribuidos y heterogéneos, donde pequeñas variaciones en las operaciones de punto flotante o entero pueden propagarse y alterar significativamente la distribución de probabilidad del siguiente token.
La relevancia de este problema es crítica en la era actual de LLMs, donde la precisión en tareas como la generación de código, la interacción con herramientas (tool-calling) o la respuesta a preguntas con contexto largo, depende de la estabilidad de la salida del modelo. Una pequeña divergencia en los logits puede llevar a una selección de token diferente, lo que a su vez puede ramificar la generación en trayectorias completamente distintas, resultando en errores funcionales o respuestas incoherentes. Este fenómeno es particularmente agudo en configuraciones de hardware heterogéneas y con el uso de cuantificación agresiva para reducir la huella de memoria y mejorar la latencia, lo que introduce compromisos inherentes entre rendimiento y precisión.
Históricamente, la reproducibilidad de cálculos de punto flotante ha sido un desafío conocido en la computación paralela y distribuida, exacerbado por las optimizaciones del compilador y las arquitecturas de hardware. En el contexto de los LLMs, donde las operaciones de multiplicación de matrices (GEMM) son omnipresentes y se realizan en hardware especializado (GPUs) con diversas precisiones (BF16, FP8, INT8, FP4), estas diferencias se magnifican. La necesidad de comprender y mitigar esta divergencia es fundamental para construir sistemas de IA fiables y predecibles.
Arquitectura del Sistema
El sistema de inferencia de LLMs se descompone en varios componentes clave que interactúan para generar tokens. En el corazón de la inferencia se encuentra el motor de inferencia, como vLLM, que orquesta la ejecución del modelo. Este motor gestiona la carga del modelo, la cuantificación de pesos y activaciones, y la gestión del KV cache. Durante la fase de prefill (procesamiento del prompt), el motor selecciona un backend de atención, como FlashAttention 2, Flash Inference o Triton Attention, que son implementaciones optimizadas de los algoritmos de atención para GPUs. Estos backends utilizan kernels CUDA específicos para cada familia de GPU y capacidad de cómputo (SM compute capability), lo que introduce variaciones en la precisión de las operaciones de multiplicación de matrices (GEMM) y acumulación (MMA).
La arquitectura del modelo, como Qwen3.6-27B, que es un modelo denso con capas Gated DeltaNet y capas de atención completa, influye en qué partes del modelo son sensibles a estas variaciones. Las capas de atención completa son particularmente críticas ya que utilizan los backends seleccionables. La cuantificación es otro componente clave, aplicándose a los pesos del modelo (W8A16, W4A16, FP8) y al KV cache (BF16, INT8, INT4). Cada esquema de cuantificación implica diferentes métodos lineales (ej. UnquantizedLinearMethod, Fp8LinearMethod, CompressedTensorsWNA16) y kernels GEMM/MMA específicos (ej. CutlassFp8BlockScaledMMKernel, MarlinLinearKernel, FlashInferFP8ScaledMMLinearKernel), que son implementaciones de bajo nivel que operan con diferentes precisiones y optimizaciones. La interacción entre estos componentes (motor de inferencia, backends de atención, esquemas de cuantificación y kernels CUDA) es la fuente principal de la divergencia observada en los logits. La salida final, los logits, son las puntuaciones del modelo para cada posible token siguiente, que luego se normalizan en probabilidades y se utilizan para la selección del siguiente token mediante un sampler configurado.
Flujo de Inferencia de LLM (Simplificado)
- 1 Prompt Input Texto de entrada del usuario o agente
- 2 Prefill (Prompt Processing) Procesamiento del prompt, cálculo de estados ocultos iniciales
- 3 Attention Backend Selección de FlashAttention 2, Flash Inference o Triton Attention
- 4 KV Cache Almacenamiento de claves y valores para tokens previos
- 5 GEMM/MMA Kernels Operaciones de multiplicación de matrices (ej. Cutlass, Marlin)
- 6 Logit Generation Cálculo de puntuaciones para cada posible token siguiente
- 7 Sampler Normalización de logits a probabilidades y selección del siguiente token
- 8 Detokenizer Conversión del token seleccionado a texto
| Capa | Tecnología | Justificación |
|---|---|---|
| compute | vLLM | Motor de inferencia de LLMs, orquesta la ejecución del modelo y la gestión de recursos. Pinned nightly build, eager execution, CUDA graphs disabled, prefix caching disabled, MTP disabled, TP1 |
| compute | CUDA | Plataforma de computación paralela y API para GPUs NVIDIA, fundamental para la ejecución de kernels de bajo nivel. Kernels específicos para cada GPU family / SM compute capability |
| compute | FlashAttention 2 | Backend de atención optimizado para GPUs, mejora la velocidad y eficiencia de la atención. vs Flash Inference, Triton Attention |
| compute | Flash Inference | Backend de atención optimizado para GPUs, alternativa a FlashAttention 2. vs FlashAttention 2, Triton Attention |
| compute | Triton Attention | Backend de atención optimizado para GPUs, utilizado como línea base en los experimentos. vs FlashAttention 2, Flash Inference |
| storage | KV Cache | Almacena las claves y valores de los tokens previamente procesados para acelerar la inferencia en contexto largo. BF16, INT8, INT4 (cuantificación) |
| compute | Qwen3.6-27B | Modelo de lenguaje base utilizado para los experimentos de inferencia. BF16, FP8, INT8, NVFP4, AWQ W4A16 (cuantificación de pesos) |
Trade-offs
Ganancias
- ▲ Reducción de huella de memoria (cuantificación)
- ▲ Mejora de latencia (cuantificación y backends optimizados)
Costes
- ▲ Precisión y fidelidad de logits (cuantificación)
- ▲ Reproducibilidad de la inferencia
- ▲ Capacidad de tool-calling (con cuantificación agresiva)
Fundamentos Teóricos
El problema de la divergencia de logits y la sensibilidad a la precisión numérica en la inferencia de LLMs se conecta directamente con los fundamentos de la aritmética de punto flotante y la estabilidad numérica en algoritmos computacionales. Un principio teórico clave es la propagación de errores en cálculos numéricos, un tema central en el análisis numérico. Pequeñas diferencias en la representación o el orden de las operaciones de punto flotante pueden acumularse y llevar a resultados significativamente diferentes, un fenómeno bien documentado por autores como Goldberg en su paper 'What Every Computer Scientist Should Know About Floating-Point Arithmetic' (1991).
La cuantificación, que es una técnica para reducir la precisión de los pesos y activaciones del modelo, se basa en principios de compresión de datos y teoría de la información, buscando un equilibrio entre la reducción de la huella de memoria y el impacto en la fidelidad del modelo. La divergencia KL (Kullback-Leibler Divergence) es una medida fundamental de la teoría de la información, introducida por Solomon Kullback y Richard Leibler en 'On Information and Sufficiency' (1951), que cuantifica la diferencia entre dos distribuciones de probabilidad. En este contexto, se utiliza para medir cuánto se desvía la distribución de probabilidad de los logits de una implementación respecto a una línea base, proporcionando una métrica rigurosa para la 'distancia' entre las salidas de diferentes configuraciones de inferencia. La sensibilidad de los modelos de atención a la precisión numérica también se relaciona con la estabilidad de los algoritmos de multiplicación de matrices a gran escala, un área activa de investigación en álgebra lineal numérica y optimización de hardware.