De la indicación a la predicción: comprensión del prellenado, la decodificación y la caché KV en los LLM

En el artículo anterior, vimos cómo un modelo de lenguaje convierte logits en probabilidades y toma muestras del siguiente token. ¿Pero de dónde vienen estos logits?

En este tutorial, adoptamos un enfoque práctico para comprender el proceso de generación:

Cómo la fase de prellenado procesa todo el mensaje en un solo paso paralelo Cómo la fase de decodificación genera tokens uno a la vez utilizando el contexto previamente calculado Cómo la caché KV elimina el cálculo redundante para hacer que la decodificación sea eficiente

Al final, comprenderá la mecánica de dos fases detrás de la inferencia LLM y por qué la caché KV es esencial para generar respuestas largas a escala.

Empecemos.

De la indicación a la predicción: comprensión del prellenado, la decodificación y la caché KV en los LLM
Foto de Neda Astani. Algunos derechos reservados.

Descripción general

Este artículo se divide en tres partes; ellos son:

Cómo funciona la atención durante el llenado previo de la fase de decodificación de la caché KV de inferencia LLM: cómo hacer que la decodificación sea más eficiente

Cómo funciona la atención durante el llenado previo

Considere el mensaje:

El clima de hoy es tan…

Como seres humanos, podemos inferir que el siguiente símbolo debería ser un adjetivo, porque la última palabra "entonces" es una configuración. También sabemos que probablemente describe el clima, por lo que palabras como "agradable" o "cálido" son más probables que algo no relacionado como "delicioso".

Los transformadores llegan a la misma conclusión mediante la atención. Durante el llenado previo, el modelo procesa todo el mensaje en un solo paso hacia adelante. Cada token se ocupa de sí mismo y de todos los tokens anteriores, creando una representación contextual que captura las relaciones a lo largo de la secuencia completa.

El mecanismo detrás de esto es la fórmula de atención del producto escalado:

$$
text{Atención}(Q, K, V) = mathrm{softmax}left(frac{QK^top}{sqrt{d_k}}right)V
$$

Analizaremos esto concretamente a continuación.

Para que el cálculo de la atención sea rastreable, asignamos a cada token un valor escalar que representa la información que contiene:

Fichas de posición Valores 1 Hoy 10 2 tiempo 20 3 es 1 4 entonces 5

Palabras como "es" y "así" tienen menos peso semántico que "hoy" o "tiempo" y, como veremos, la atención refleja esto de forma natural.

Cabezas de atención

En los transformadores reales, los pesos de atención son valores continuos aprendidos durante el entrenamiento a través del producto escalar $Q$ y $K$. El comportamiento de las cabezas de atención es aprendido y normalmente imposible de describir. Ninguna cabeza está programada para “atender posiciones iguales”. Las cuatro reglas siguientes son una ilustración simplificada para hacer que el mecanismo de atención sea más intuitivo, mientras que la agregación ponderada sobre $V$ es la misma.

Estas son las reglas en nuestro ejemplo de juguete:

Atender a los tokens en posiciones de números pares Atender al último token Atender al primer token Atender a cada token

Para simplificar en este ejemplo, las salidas de estos cabezales se combinan (promedian).

Repasemos el proceso de precarga:

Hoy

Fichas pares → ninguna Última ficha → Hoy → 10 Primera ficha → Hoy → 10 Todas las fichas → Hoy → 10

clima

Fichas pares → clima → 20 Última ficha → clima → 20 Primera ficha → Hoy → 10 Todas las fichas → promedio (Hoy, clima) → 15

es

Fichas pares → clima → 20 Última ficha → es → 1 Primera ficha → Hoy → 10 Todas las fichas → promedio (Hoy, clima, es) → 10.33

entonces

Fichas pares → promedio(tiempo, entonces) → 12.5 Última ficha → entonces → 5 Primera ficha → Hoy → 10 Todas las fichas → promedio(Hoy, tiempo, es, entonces) → 9

Atención paralela

Si el mensaje contuviera 100.000 tokens, calcular la atención paso a paso sería extremadamente lento. Afortunadamente, la atención se puede expresar como operaciones tensoriales, lo que permite calcular todas las posiciones en paralelo.

Esta es la idea clave de la fase de prellenado en la inferencia LLM: cuando proporciona un mensaje, contiene varios tokens y se pueden procesar en paralelo. Este procesamiento paralelo ayuda a acelerar el tiempo de respuesta del primer token generado.

Para evitar que los tokens vean tokens futuros, aplicamos una máscara causal, de modo que solo puedan atender a sí mismos y a los tokens anteriores.

Producción:

Ahora podemos empezar a escribir las “reglas” para las 4 cabezas de atención.

En lugar de calcular puntuaciones de los vectores $Q$ y $K$ aprendidos, los elaboramos a mano directamente para que coincidan con nuestras cuatro reglas de atención. Cada cabeza produce una matriz de puntuación de forma (n, n), con una puntuación por par de claves de consulta, que se enmascara y se pasa a través de softmax para producir pesos de atención:

Producción:

El resultado de este paso se denomina vector de contexto, que representa un resumen ponderado de todos los tokens anteriores.

De contextos a logits

Cada cabeza de atención ha aprendido a captar diferentes patrones en la entrada. Juntos, los cuatro valores de contexto [12.5, 5.0, 10.0, 9.0] forman un resumen de lo que representa “El clima de hoy es tan…”. Luego se proyectará en una matriz, en la que cada columna codifica qué tan fuerte está asociado un vocabulario determinado con la señal de cada cabeza de atención, para dar una puntuación logit por palabra.

Para nuestro ejemplo, digamos que tenemos "agradable", "cálido" y "delicioso" en el vocabulario:

Entonces, los logits para "agradable" y "cálido" son mucho más altos que "delicioso".

La fase de decodificación de la inferencia LLM

Ahora supongamos que el modelo genera el siguiente token: "agradable". La tarea ahora es generar el siguiente token con el mensaje extendido:

El clima de hoy es tan agradable…

Las primeras cuatro palabras del mensaje extendido son las mismas que las del mensaje original. Y ahora tenemos la quinta palabra del mensaje.

Durante la decodificación, no volvemos a calcular la atención de todos los tokens anteriores ya que el resultado sería el mismo. En cambio, calculamos la atención solo para el nuevo token para ahorrar tiempo y recursos informáticos. Esto produce una única fila de atención nueva.

Producción:

Ahora aplicamos las 4 cabezas de atención y calculamos el nuevo vector de contexto:

Producción:

Sin embargo, a diferencia del prellenado, donde todo el mensaje se procesa en paralelo, la decodificación debe generar tokens uno a la vez (autorregresivamente) porque los tokens futuros aún no se han generado. Sin el almacenamiento en caché, cada paso de decodificación volvería a calcular las claves y los valores de todos los tokens anteriores desde cero, lo que haría que el trabajo total en todos los pasos de decodificación fuera $O(n^2)$ en longitud de secuencia. La caché de KV reduce esto a $O(n)$ calculando los $K$ y $V$ de cada token exactamente una vez.

Caché KV: cómo hacer que la decodificación sea más eficiente

Para que la codificación autorregresiva sea eficiente, podemos almacenar las claves ($K$) y los valores ($V$) para cada token por separado para cada cabeza de atención. En este ejemplo simplificado usaríamos solo un caché. Luego, durante la decodificación, cuando se genera un nuevo token, el modelo no vuelve a calcular las claves y los valores de todos los tokens anteriores. Calcula la consulta del nuevo token y atiende las claves y valores almacenados en caché de los tokens anteriores.

Si volvemos a mirar el código anterior, podemos ver que no es necesario volver a calcular $K$ para todo el tensor:

En su lugar, podemos simplemente calcular K para la nueva posición y adjuntarla a la matriz K que ya hemos calculado y guardado en caché:

Aquí está el código completo para la fase de decodificación usando caché KV:

Producción:

Observe que esto es idéntico al resultado que calculamos sin el caché. La caché KV no cambia lo que calcula el modelo, pero elimina cálculos redundantes.

La caché de KV se diferencia de la caché de otras aplicaciones en que el objeto almacenado no se reemplaza sino que se actualiza. Cada nuevo token agregado al mensaje agrega una nueva fila al tensor almacenado. Implementar un caché KV que pueda actualizar eficientemente el tensor es la clave para acelerar la inferencia de LLM.

Lecturas adicionales

A continuación se presentan algunos recursos que pueden resultarle útiles:

Resumen

En este artículo, recorrimos las dos fases de la inferencia LLM. Durante el llenado previo, el mensaje completo se procesa en un paso directo paralelo y las claves y valores de cada token se calculan y almacenan. Durante la decodificación, el modelo genera un token a la vez, utilizando solo la consulta del nuevo token contra las claves y valores almacenados en caché para evitar un nuevo cálculo redundante. Prefill calienta el caché KV y la decodificación lo actualiza. Un llenado previo más rápido significa que verá más rápido el primer token en la respuesta y una decodificación más rápida significa que verá más rápido el resto de la respuesta. Juntas, estas dos fases explican por qué los LLM pueden procesar solicitudes largas rápidamente pero generar salida token por token, y por qué la caché KV es esencial para que esa generación sea práctica a escala.