Fondamentaux de l’IA

Votre cache KV n’a pas de problème de bits, mais un problème de géométrie.

mm
Ajouter Unite.AI à vos sources préférées sur Google

À une précision de 2 bits identique, une décision concernant l’axe que vous quantifiez fait varier le score d’un benchmark de 2,88 à 63,53. Les clés et les valeurs nécessitent un traitement opposé — et la raison se trouve dans l’équation d’attention, et non dans le matériel.

Prenez Llama-2-13B. Groupez son cache de clés-valeurs par une taille de groupe de quantification de 32 en deux bits, tout en laissant tout le reste en place — même modèle, même budget de bits, mêmes tailles de groupe, mêmes benchmarks.

En fonction d’un seul élément d’une décision d’implémentation, les résultats de précision de CoQA donnent soit 2,88, soit 63,53. Le score en utilisant une précision complète est de 66,37.

La décision ne porte pas sur le nombre total de bits utilisés. La question est simplement de savoir quel axe choisir pour regrouper lors du calcul de chaque facteur d’échelle ? Lorsque vous décidez d’utiliser le canal comme dimension de regroupement (clés) et le jeton comme dimension de regroupement (valeurs), vous vous retrouvez à l’intérieur de quatre points de la performance de précision complète. Si vous inversez l’un ou l’autre de ces choix, vous subissez une perte de qualité. Si vous inversez les deux, le modèle ne fonctionne plus.

Quatre façons de dépenser les mêmes 2 bits sur le même cache. Résultats de l’ablation KIVI sur Llama-2-13B à une taille de groupe de 32.

La quantification est généralement considérée comme un seul réglage : 8 bits, 4 bits, 2 bits, avec un coût d’exactitude attaché. À l’intérieur du cache KV, ce n’est pas comme ça. C’est choisir des systèmes de coordonnées, et des systèmes différents s’appliquent aux clés et aux valeurs. Cet article explique pourquoi. En résumé : l’erreur de quantification dépend de la plage de valeurs dans les groupes ; les clés et les valeurs ont une structure très différente ; et les gens trébuchent souvent parce que vous ne pouvez pas dériver l’axe correct à partir de la distribution de valeurs du tout. Vous devez regarder comment l’erreur change après que l’attention l’ait consommée. Cela donne un principe général pour compresser les activations intermédiaires et une bonne raison de douter de l’erreur de reconstruction comme proxy de qualité.

Pourquoi le cache KV est où cela mord

Pendant la phase de génération, un transformateur stocke toutes les données de projection de clés-valeurs (KV) des jetons qu’il a traités précédemment dans un cache afin de ne pas avoir à recalculer ces données à nouveau. Ce cache croît linéairement avec la longueur du contexte et la taille du lot. Finalement, cela entraînera une croissance du cache plus grande que le modèle lui-même.

Cette augmentation de croissance peut être facilement identifiée lors de l’examen de la consommation de mémoire de différentes parties du modèle. Dans l’analyse KVQuant de LLaMA-7B, les poids représentent environ 98 pour cent de la mémoire à une longueur de séquence de 512, avec des activations à 2 pour cent. À 128K de contexte, le rapport s’inverse pour environ 16 pour cent de poids et 84 pour cent de cache KV. Lorsque nous examinons une analyse d’OPT-175B citée par les auteurs de KIVI, ils ont trouvé des résultats similaires. Plus précisément, à une taille de lot de 512 avec une invite de 512 jetons, le cache KV atteint 1,2 To — plusieurs fois la taille des poids du modèle.

Cependant, la capacité n’est qu’à moitié du problème ici. Le GPU doit lire l’ensemble du cache KV à partir de la mémoire du périphérique pour chaque jeton qu’il génère. Cela signifie que pendant que le GPU lit le cache KV, les cœurs de calcul restent inactifs. En réduisant ainsi la taille globale du cache, on augmente l’espace de traitement disponible et on réduit le temps passé à attendre les transferts de données.

Qu’est-ce que l’erreur de quantification est vraiment faite

La quantification entière des entiers est mathématiquement simple. Pour un groupe de nombres, vous enregistrez le plus petit nombre comme point zéro, puis divisez la plage de ce groupe par le nombre de niveaux qui peuvent être représentés pour obtenir une taille d’étape. Vous arrondissez ensuite chaque élément à l’étape la plus proche. Deux résultats immédiats suivent. Premièrement, l’erreur par élément est limitée par la moitié d’une étape. Deuxièmement, la taille de l’étape est la plage du groupe divisée par 2ᴮ − 1. À 2 bits, vous n’avez que 4 niveaux pour couvrir la dispersion qui existe dans ce groupe. Un élément qui est cent fois plus grand par rapport à ses voisins ne se comporte pas seulement mal. Il gonfle la taille de l’étape pour tous les autres éléments partageant le même groupe, et tous deviennent plus grossiers ensemble. Le groupe est l’unité de dégâts. Le choix d’un axe signifie décider quels éléments souffrent ensemble. En reformulant la question différemment, il ne s’agit plus de « combien de bits puis-je me permettre ? » mais de « où sont les valeurs extrêmes et puis-je les isoler ? »

Clés : les valeurs aberrantes vivent dans des canaux fixes

Les grands modèles de langage contiennent des activations qui sont inhabituellement grandes par rapport à la plupart des activations. Sun et ses collègues ont catalogué ces très grandes activations à travers différentes familles de modèles : dans Mixtral 8x7B, la plus grande magnitude est proche de 7000 tandis que la magnitude médiane des caractéristiques est d’environ 0,3 — environ quatre ordres de grandeur à part. Ceux-ci sont très rares ; ils restent fixes dans des dimensions qui changent rarement avec l’entrée, et ils ne sont pas accidentels. Ils agissent comme des biais implicites, et ils sont ce qui focalise l’attention sur quelques jetons : le comportement des sinks d’attention. Dans le cache des clés, cette structure est très claire : des canaux spécifiques transportent des grandeurs très grandes de manière cohérente à travers chaque jeton d’une séquence. Regroupez par jetons, et chaque groupe contient ces canaux aberrants, donc chaque groupe de taille d’étape est défini par les aberrants, et tous les canaux ordinaires paient pour cela. Regroupez par canaux, et les canaux aberrants forment leurs propres groupes. Leur plage interne est grande mais autonome ; les canaux ordinaires sont laissés seuls. Les résultats correspondent. En moyenne sur les couches et les têtes sur Llama-2-13B, KIVI rapporte une erreur de reconstruction de clé de 13,67 sous regroupement par jeton contre 4,55 par canal, et — plus important encore — une erreur de score d’attention de 47,00 contre 9,60. La quantification des clés par jeton produit environ cinq fois l’erreur de score. Les scores sont ensuite d’accord avec les métriques significatives pour les clés ; la quantification par canal excelle sur les deux fronts.

Valeurs : où l’intuition se brise

Le cache de valeurs ne montre pas de modèle de canal aberrant. Il semble être assez plat. À lui seul, en fonction de l’argument de plage, on pourrait s’attendre à ce que l’un ou l’autre de ces axes produise une qualité de compression similaire.

Ils ne le font pas. Quelle que soit la manière dont la gestion des clés est mise en œuvre (les résultats de 2,80 et 2,88), la compression par canal des valeurs fait s’effondrer le modèle.

Et voici l’astuce : si vous mesurez cette perte en utilisant l’erreur de reconstruction brute sur le tenseur d’origine pour lequel chaque valeur a été compressée, la quantification par canal des valeurs semble légèrement meilleure, à 3,73 contre 4,57. Si vous validez votre compression de la manière la plus évidente, vous choisirez la configuration qui détruit le modèle.

Erreur de quantification du cache de valeurs sur Llama-2-13B, mesurée de deux manières. La métrique de tenseur stocké et la métrique de sortie consommée diffèrent de plus d’un ordre de grandeur.

La résolution est que le cache de valeurs n’est jamais lu directement. Il est consommé par un produit matriciel : la sortie d’attention est une somme pondérée de vecteurs de valeurs à travers les jetons, avec des scores d’attention softmax comme poids. En raison de cela, l’erreur pertinente est celle introduite au cours de ce processus et non dans les tenseurs eux-mêmes. Mesurée en termes de sortie d’attention, l’ordre était complètement inversé. L’erreur relative signalée par KIVI pour la sortie d’attention due à la quantification par jeton des vecteurs de valeurs était de 3,55 par rapport à 49,89 pour la quantification par canal — plus de quatorze fois supérieure pour ce qui semblait être le meilleur choix en fonction de la façon dont il était compressé.

L’explication est la rareté de l’attention, qu’ils ont mesurée à 84,3 pour cent. La majorité de l’information contenue dans la sortie peut être attribuée à un petit nombre de jetons très importants. La quantification par jeton confine l’erreur de chaque jeton à ce jeton, donc les erreurs sur les jetons non importants sont multipliées par des poids d’attention presque nuls et disparaissent effectivement. La quantification par canal étale l’erreur de chaque jeton sur une échelle de canal partagée, donc les jetons mal représentés contaminent la représentation de ceux qui comptent. La rareté qui rend l’attention efficace est la même propriété qui rend la quantification par jeton sûre.

La leçon transférable est plus large que le cache KV : mesurez l’erreur de compression là où le tenseur est consommé, et non où il est stocké. Un postulat implicite fait par l’erreur de reconstruction est que chaque composant d’un tenseur a un poids égal lorsqu’il contribue à la sortie finale. L’attention explicite ne le fait pas. Toute opération en aval qui pondère, gâte ou espacifie son entrée brise ce postulat. Les lecteurs familiers avec mon article précédent concernant les angles morts dans les métriques d’évaluation dans les systèmes de récupération reconnaîtront que ces résultats sont similaires aux échecs précédemment décrits : des métriques facilement calculées qui signalent sur autre chose que ce qui était prévu.

Les incrustations rotatives compliquent les clés

Il y a quelques problèmes avec l’utilisation des incrustations de position rotatives (RoPE). RoPE fait pivoter des paires de canaux en fonction de la position relative de chaque jeton. Ce mélange dissout partiellement la structure de canal fixe qui a rendu la quantification par canal des clés fonctionnelle à partir de la première place — un canal aberrant est pivoté dans ses voisins, et les voisins héritent de la plage. La réponse de KVQuant est l’ordre : quantifiez les clés avant que la rotation ne soit appliquée, et appliquez RoPE après la déquantification. Avec la quantification par canal des clés, les types de données non uniformes et l’isolement d’une petite fraction de valeurs aberrantes, cela obtient moins de 0,1 de dégradation de perplexité à 3 bits, et permet de servir LLaMA-7B jusqu’à 1 million de jetons de contexte sur un seul A100-80GB.

Il est également important de comprendre le niveau d’impact de RoPE. Les auteurs de l’article « RotateKV » ont signalé une augmentation de 145 pour cent des erreurs de quantification une fois RoPE ajouté, et noté que les canaux aberrants diffèrent entre les têtes d’attention — c’est pourquoi l’application d’une matrice de rotation partagée partout est insuffisante, et les rotations adaptées à la tête font mieux.

La taxe des systèmes, et pourquoi elle n’est pas un détail

La quantification par jeton convient bien à la décoding. Chaque jeton arrive ; vous le quantifiez, vous l’ajoutez à la séquence (le long de la dimension du jeton), rien d’autre ne bouge.

Cependant, la quantification par canal ne convient pas. Puisque les statistiques d’un canal s’étendent sur des jetons qui n’ont pas encore été générés, vous ne pouvez pas calculer un facteur d’échelle lorsque le jeton arrive. Le contournement de KIVI est de conserver les jetons les plus récents — jusqu’à 128 — en précision complète dans un tampon résiduel, et de quantifier en groupes une fois qu’il y en a suffisamment accumulés.

Il se trouve que le tampon résiduel devient porteur de charge, plutôt que juste une chose incidente. Sur GSM8K avec Llama-2-7B, les scores de précision complète sont de 13,50. Complètement quantifiés à 2 bits avec les axes corrects, ils sont de 5,76. Les mêmes axes et les mêmes bits, plus le tampon résiduel de jetons récemment produits à précision complète, sont de 12,74. Une fenêtre glissante de jetons récemment produits à précision complète récupérera une grande partie de ce qui a été perdu en raison d’une quantification agressive sur des problèmes difficiles à plusieurs étapes — ce qui aurait du sens si l’on considère quels jetons étaient attentifs à une chaîne d’opérations arithmétiques.

Il y a un avantage significatif à faire toutes ces choses correctement — comme le signale KIVI, 2,6 fois moins d’utilisation de mémoire de pointe pour Llama-2-7B, permettant des tailles de lot jusqu’à 4 fois plus grandes, ainsi que 2,35 à 3,47 fois meilleure efficacité sur une tâche de service réel.

Que faire avec cela

  1. N’utilisez jamais un seul quantifieur pour les deux. Utilisez des quantifieurs différents pour les clés (par canal) et pour les valeurs (par jeton). Un pipeline qui applique un seul quantifieur au « cache KV » a probablement déjà sacrifié la plupart de la qualité possible lors de l’utilisation d’un petit nombre de bits pour représenter chaque valeur.
  2. Quantifiez les clés avant RoPE. C’est une question de correction plutôt que de préférence.
  3. Stockez une fenêtre de précision complète de jetons récemment générés. Bien que le stockage d’une telle fenêtre prenne très peu de mémoire par rapport à la taille du cache, c’est précisément cette zone qui génère une grande partie de la précision pour les tâches difficiles.
  4. Ne validez pas sur l’erreur de reconstruction. Validez toujours en fonction de la sortie d’attention ou de la performance de la tâche finale. La métrique de stockage n’est pas seulement bruyante — pour les valeurs, elle pointe dans la mauvaise direction.
  5. Ne validez pas sur des benchmarks de choix multiple à court contexte. Les auteurs de KIVI évitent délibérément les tâches à fermeture comme MMLU pour cette évaluation, car une seule étape de décodage lisant les logits de sortie à peine exerce le cache. Toute évaluation qui ne construit pas un cache au fil du temps et n’exécute pas la génération à partir de celui-ci ne pourra jamais observer les défaillances inhérentes à la conception de votre système.

Where the Work Is Headed

Bien qu’il y ait encore quelque chose à faire concernant la nature géométrique du problème, de nombreux chercheurs continuent à étudier les moyens par lesquels les canaux aberrants sont répartis entre les différentes têtes de transformateur, et comment les limitations matérielles affectent les regroupements les moins chers : InnerQ plie la normalisation des clés par canal dans les poids des clés et des requêtes pendant le préremplissage. Par conséquent, aucun surcoût n’est encouru au moment de l’exécution. De plus, InnerQ stocke des fenêtres de haute précision à la fois pour les jetons récemment générés et pour les jetons d’attention. En faisant cela, InnerQ élimine la possibilité pour les aberrants du canal de contaminer les canaux voisins.

D’autres proposent que, au lieu de stocker l’ensemble du cache, nous devrions stocker uniquement suffisamment d’informations pour être en mesure de rematérialiser la clé et/ou la valeur(s) à la demande à partir d’une représentation du cache plus petite.

Enfin, il est important de se rappeler que la précision n’est pas le seul paramètre que la quantification affecte. Des recherches récemment publiées ont démontré une dégradation d’alignement résultant de la quantification des caches KV. De plus, ces recherches ont documenté une dégradation d’alignement même dans les environnements de service de production vLLM utilisant des caches FP8 avec un protocole de récupération sans formation qui a restauré jusqu’à 97 pour cent de ce qui a été perdu en termes d’alignement. Ainsi, même si une configuration conserve ses résultats de benchmark, cela ne signifie pas nécessairement qu’elle conserve tous les autres paramètres dont vous vous souciez.

Le principe général

L’idée de quantification a été formulée comme un « budget de précision » : combien de bits puis-je me permettre de sacrifier ? Le cache KV montre que la question la plus utile est structurelle. La précision est allouée en groupes ; le groupe est l’unité de dégâts, et l’axe que vous regroupez détermine quels éléments partagent leur sort. L’axe correct est celui sur lequel votre tenseur est consommé, c’est-à-dire la façon dont vous utilisez votre tenseur et non la façon dont votre tenseur apparaît lorsqu’il est stocké en mémoire. Les clés sont utilisées via un calcul de produit scalaire contre la requête. Un canal corrompu unique empoisonnera tous les scores. Les valeurs sont consommées par un calcul de moyenne pondérée espacée à travers les jetons. Par conséquent, un jeton unique corrompu est simplement pondéré.

Deux tenseurs de dimensions identiques et générés par deux couches consécutives sont traités différemment. Il vaut la peine de se demander, pour toute activation que vous prévoyez de compresser, quelle opération éloigne cela, et mon regroupement respecte-t-il cela ?

Himanshu Goel est un chercheur en IA/ML spécialisé dans la génération augmentée de récupération pour des domaines à enjeux élevés, notamment les flux de travail de documents biomédicaux, financiers et réglementaires.