¿Cómo se implementa la caché KV en la inferencia de un transformador autoregresivo?
sobes.tech AI
Respuesta de la IA
En los transformadores autorregresivos, durante la inferencia, es importante utilizar eficientemente los cálculos para no volver a calcular toda la secuencia desde cero. Para ello, se utiliza la caché KV, que almacena las claves (Key) y los valores (Value) de las capas de atención (autoatención).
La idea es que, al generar el siguiente token, ya hemos calculado las claves y los valores para todos los tokens anteriores. En lugar de recalcularlos, los almacenamos en la memoria (caché) y, al calcular la atención para el nuevo token, usamos esta caché.
Esto funciona aproximadamente así:
- Para cada capa del transformador, al procesar un token, se almacenan las claves y los valores en buffers separados.
- Al generar el siguiente token, solo se calculan la clave y el valor para él, y luego se combinan con los almacenados previamente.
- El mecanismo de atención utiliza el conjunto combinado de claves y valores para calcular el contexto.
Esto acelera significativamente la inferencia, reduciendo la complejidad computacional de cuadrática en la longitud de la secuencia a lineal.
Ejemplo de pseudocódigo para una capa:
# kv_cache almacena las claves y valores de los tokens anteriores
new_key, new_value = compute_kv(new_token)
kv_cache.keys = concatenate(kv_cache.keys, new_key)
kv_cache.values = concatenate(kv_cache.values, new_value)
output = attention(query=new_key, keys=kv_cache.keys, values=kv_cache.values)
De esta manera, la caché KV es un mecanismo para guardar las claves y valores intermedios para acelerar la generación secuencial en transformadores.