Cache KV dans les LLM : Compréhension et implémentation de zéro
Le cache KV est une technique d'inférence efficace pour les LLM qui stocke les vecteurs intermédiaires de clés et de valeurs pour les réutiliser, évitant ainsi des calculs redondants. L'article explique le concept et fournit une implémentation de code de zéro en PyTorch, montrant comment modifier un mécanisme d'attention multi-têtes pour mettre en cache les clés et les valeurs pendant la génération de texte.
Le cache KV stocke les vecteurs de clé (K) et de valeur (V) issus des calculs d'attention précédents afin de les réutiliser lors de la génération des jetons suivants, ce qui accélère l'inférence en évitant de recalculer. Sans cache KV, chaque étape de génération recalcule les clés et les valeurs pour tous les jetons précédents ; avec le cache, seuls les vecteurs du nouveau jeton sont calculés et ajoutés au cache. L'article fournit une implémentation de code basée sur un modèle de type GPT issu du livre de l'auteur. Les modifications clés sont : l'ajout de tampons de cache (cache_k et cache_v) dans la classe MultiHeadAttention, la modification de la méthode forward pour utiliser le cache de manière conditionnelle, l'ajout d'une méthode reset, et la propagation de l'indicateur use_cache à travers le modèle. En génération, lorsque use_cache est vrai, le modèle ne traite que le nouveau jeton après la requête initiale, tandis que sans cache, il traite la séquence complète à chaque étape. Une comparaison de performance simple montre que le cache KV double approximativement la vitesse de génération pour l'exemple testé.
Source: Sebastian Raschka —
original
