Caché KV en los LLM: comprensión e implementación desde cero
El caché KV es una técnica para la inferencia eficiente de los LLM que almacena los vectores intermedios de clave y valor para su reutilización, evitando cálculos redundantes. El artículo explica el concepto y proporciona una implementación de código desde cero en PyTorch, mostrando cómo modificar un mecanismo de atención de múltiples cabezas para almacenar en caché las claves y los valores durante la generación de texto.
La caché KV almacena los vectores de clave (K) y valor (V) de cálculos de atención previos para reutilizarlos al generar tokens posteriores, lo que acelera la inferencia al evitar recalcularlos. Sin una caché KV, cada paso de generación recalcula las claves y valores de todos los tokens anteriores; con la caché, solo se calculan los vectores del nuevo token y se añaden a la caché. El artículo proporciona una implementación de código basada en un modelo similar a GPT del libro del autor. Los cambios clave son: añadir buffers de caché (cache_k y cache_v) en la clase MultiHeadAttention, modificar el método forward para usar condicionalmente la caché, añadir un método reset y propagar el indicador use_cache a través del modelo. En la generación, cuando use_cache es True, el modelo procesa solo el nuevo token después del prompt inicial, mientras que sin caché procesa la secuencia completa en cada paso. Una comparación de rendimiento simple muestra que la caché KV aproximadamente duplica la velocidad de generación para el ejemplo probado.
Fuente: Sebastian Raschka —
original
