- Partant de l’attention softmax, l’article dérive progressivement l’attention linéaire à état de taille fixe, DeltaNet qui n’enregistre que les erreurs, Gated DeltaNet qui atténue l’état entier, puis Kimi Delta Attention (KDA), qui applique une atténuation par canal.
- L’attention linéaire de base stocke dans l’état (S_t) la somme des produits extérieurs key-value passés et fonctionne linéairement avec la longueur de séquence, mais elle souffre d’une interférence d’écriture additive : elle ajoute aux associations existantes au lieu d’affecter une nouvelle valeur.
- DeltaNet enregistre la différence entre la valeur prédite à partir de la key courante et la value cible, multipliée par (\beta_t) ; les trois interprétations — condition de reconstruction immédiate, descente de gradient en ligne et mise à jour d’état de rang 1 — mènent à la même formule.
- Gated DeltaNet atténue d’abord l’état entier avec un scalaire (\alpha_t), tandis que KDA l’étend en une matrice diagonale (D_t=\operatorname{Diag}(\alpha_t)), afin de conserver ou supprimer l’information à des taux différents selon les canaux de key.
- La même récurrence KDA s’exécute via un kernel Triton récurrent fusionné pour le décodage et via une méthode par chunks pour l’entraînement et les longs préfills ; cette dernière restaure les dépendances internes aux tokens par résolution triangulaire et les reformule en multiplications matricielles.
Notation et ordre de développement
- Dans la notation bra-ket, (\lvert q\rangle) est un vecteur colonne, (\langle k\rvert) un vecteur ligne, (\langle k\vert q\rangle) un scalaire, et (\lvert v\rangle\langle k\rvert) une matrice.
- On utilise une seule tête d’attention causale et des vecteurs réels ; on suppose que les keys de DeltaNet sont normalisées et que l’état effectue une application de l’espace des keys vers l’espace des values.
- L’ordre de développement est : attention softmax → attention linéaire → DeltaNet → Gated DeltaNet → KDA, puis raccordement final aux implémentations Triton récurrente et par chunks.
- Parmi les variantes de la famille DeltaNet, deux sont utilisées dans les familles récentes de modèles Qwen et Kimi.
De l’attention quadratique à un état linéaire
- L’attention softmax causale classique calcule la similarité entre keys et queries, normalise en distribution les scores de toutes les keys passées, puis produit une somme pondérée des vecteurs value.
- Une séquence de longueur (T) contient (T^2) paires key-query.
- En inférence autorégressive, les keys et values peuvent être mises en cache, mais la taille du cache augmente avec la séquence.
- Chaque nouvelle query doit tout de même consulter tout le passé.
- Le dénominateur du softmax dépend conjointement de la query courante et de toutes les keys précédentes, ce qui rend difficile une simple réorganisation de l’ordre de calcul.
- Si l’on supprime le softmax, on peut regrouper la sortie sous forme de somme des produits extérieurs key-value passés.
- (S_t=\sum_{i\le t}\lvert v_i\rangle\langle k_i\rvert)
- (S_t=S_{t-1}+\lvert v_t\rangle\langle k_t\rvert)
- (\lvert o_t\rangle=S_t\lvert q_t\rangle)
- L’identité centrale est ((\lvert v\rangle\langle k\rvert)\lvert q\rangle=\langle k\vert q\rangle\lvert v\rangle) : au lieu de stocker toutes les keys et values passées, on stocke le produit extérieur agrégé dans un état de taille fixe (d_v\times d_k).
- Comme les tokens sont parcourus une seule fois, le coût est linéaire en longueur de séquence, mais au prix de la normalisation et de la sélectivité du softmax.
- Les attentions linéaires plus élaborées utilisent des feature maps et des termes de normalisation.
Le problème d’écriture additive de l’attention linéaire
- Juste après avoir écrit (\lvert v_t\rangle\langle k_t\rvert) sur la key courante normalisée, une lecture avec la même key donne (S_t\lvert k_t\rangle=S_{t-1}\lvert k_t\rangle+\lvert v_t\rangle).
- La nouvelle écriture n’affecte pas la mémoire de sorte qu’elle renvoie (v_t) ; elle ajoute (v_t) avec une logique
+=à la valeur déjà renvoyée. - Si l’état précédent renvoyait déjà la bonne valeur, cette même value est doublée ; comme les keys ne sont pas mutuellement orthogonales, chaque écriture peut interférer avec les précédentes.
- L’attention linéaire fournit une mémoire associative compressée, mais elle effectue une mise à jour additive au lieu d’une mise à jour proche du
=nécessaire.
DeltaNet : écrire l’erreur de prédiction plutôt que la valeur
- DeltaNet lit d’abord la prédiction existante pour la nouvelle key, (\widehat v_t=S_{t-1}k_t), puis n’enregistre que la différence au lieu de la value complète.
- (e_t=\beta_t(v_t-S_{t-1}k_t))
- (S_t=S_{t-1}+e_tk_t^\mathsf T)
- L’intensité d’écriture apprise (\beta_t) est dans l’intervalle ([0,1]).
- Une relecture immédiate avec la même key donne ((1-\beta_t)S_{t-1}k_t+\beta_tv_t).
- Si (\beta_t=1), elle renvoie exactement (v_t).
- Une valeur plus faible ne déplace l’ancienne prédiction que partiellement vers la cible.
- La mise à jour est locale dans l’espace des keys.
- Dans une direction de query orthogonale à la key courante, la mise à jour par produit extérieur vaut 0 et la réponse ne change donc pas.
- Seule l’association dans la direction de la key courante est remplacée sélectivement.
-
Dérivation par perte de reconstruction
- Si l’on voit l’état (S) comme une application linéaire et que l’on pose, pour la paire key-value courante, la perte (\frac12\lVert Sk_t-v_t\rVert_2^2), son gradient est ((Sk_t-v_t)k_t^\mathsf T).
- Effectuer un pas de descente de gradient de taille (\beta_t) à partir de (S_{t-1}) donne exactement la règle de mise à jour de DeltaNet.
- La même mise à jour peut être interprétée de trois façons.
- En opération mémoire, (\beta_t) est l’intensité de remplacement de l’association existante.
- En apprentissage en ligne, (\beta_t) est le taux d’apprentissage.
- En algèbre linéaire, c’est le produit extérieur de rang 1 entre l’erreur de prédiction et la key.
-
Transition d’état structurée
- En développant la mise à jour, on obtient (S_t=S_{t-1}(I-\beta_tk_tk_t^\mathsf T)+\beta_tv_tk_t^\mathsf T).
- Pour une key unitaire, (I-\beta_tk_tk_t^\mathsf T) a pour valeur propre (1-\beta_t) dans la direction de la key courante, et 1 dans toutes les directions orthogonales.
- Elle supprime d’abord l’association dans la direction de key existante puis ajoute la nouvelle, mais ne résout pas encore la gestion de la durée de vie de l’état entier.
Gated DeltaNet : oublier d’abord l’état entier
- Lorsqu’on compresse tout le passé dans une seule matrice, on ne peut pas ignorer sélectivement certains tokens déjà fusionnés dans l’état.
- DeltaNet corrige le voisinage de la key courante, mais les informations anciennes dans d’autres directions restent présentes et peuvent continuer à contribuer aux lectures futures.
- Gated DeltaNet applique une porte de conservation scalaire apprise (\alpha_t\in[0,1]).
- Oubli avec (\widetilde S_t=\alpha_tS_{t-1})
- Prédiction avec (\widehat v_t=\widetilde S_tk_t)
- Correction avec (e_t=\beta_t(v_t-\widehat v_t))
- Écriture avec (S_t=\widetilde S_t+e_tk_t^\mathsf T)
- L’ordre oubli → prédiction → correction → écriture est important.
- Si l’on prédit avant l’atténuation, la mémoire utilisée pour calculer l’erreur diffère de celle effectivement mise à jour.
- La règle delta prend en charge le remplacement pour la key cible, tandis que la porte scalaire se charge de la suppression globale ; elles résolvent des problèmes différents.
- Toutefois, un unique (\alpha_t) s’applique à toute la matrice, si bien que tous les canaux de key doivent être conservés ou oubliés au même taux.
Kimi Delta Attention : atténuation par canal
- Kimi Delta Attention remplace le scalaire (\alpha_t) par un vecteur de dimension (d_k) et construit (D_t=\operatorname{Diag}(\alpha_t)).
- Comme l’état applique l’espace des keys vers l’espace des values, les canaux de key correspondent aux colonnes de (S), et la multiplication à droite (S_{t-1}D_t) applique un taux de conservation différent à chaque colonne.
- KDA fonctionne dans l’ordre suivant.
- Atténuation par canal de key avec (\widetilde S_t=S_{t-1}D_t)
- Prédiction avec (\widehat v_t=\widetilde S_tk_t)
- Correction avec (e_t=\beta_t(v_t-\widehat v_t))
- Écriture avec (S_t=\widetilde S_t+e_tk_t^\mathsf T)
- Lecture avec (o_t=S_t(d_k^{-1/2}q_t))
- Le changement conceptuel de Gated DeltaNet à KDA se résume à promouvoir (\alpha_t) en (D_t), mais cela permet d’effacer un canal tout en en conservant un autre.
-
Transition diagonale-bas rang
- En développant KDA, (S_t=S_{t-1}A_t+\beta_tv_tk_t^\mathsf T), avec (A_t=D_t(I-\beta_tk_tk_t^\mathsf T)).
- On peut écrire (A_t=D_t-b_ta_t^\mathsf T), (b_t=D_tk_t), (a_t^\mathsf T=\beta_tk_t^\mathsf T), ce qui en fait une transition diagonale-bas rang (DPLR).
- DPLR désigne une transition (d_k\times d_k) qui agit dans l’espace des keys ; l’état mémoire lui-même reste une matrice (d_v\times d_k).
- Chaque famille ajoute la fonction suivante.
- Attention linéaire : mémoire récurrente de taille fixe
- DeltaNet : remplacement sélectif dans la direction cible
- Gated DeltaNet : atténuation de l’état entier
- KDA : atténuation par canal de key
- L’implémentation stocke généralement (g_t=\log\alpha_t\le0), puis calcule le taux de conservation avec (\exp(g_t)).
- Une implémentation de référence en 5 étapes avec disposition transposée (d_k\times d_v) est disponible dans
naive_recurrent_kda.
Kernel Triton récurrent fusionné pour le décodage
- KDA a deux principaux modes d’exécution.
- Mode récurrent fusionné : adapté au décodage, aux séquences courtes et au serving avec état persistant.
- Mode par chunks : adapté à l’entraînement et aux longs préfills.
fused_recurrent_kda_fwdlance un programme Triton par séquence, tête de value et tuile de values de largeur 32.BKcouvre la dimension key dans les configurations généralement prises en charge.- Chaque programme possède une tuile
[BK, BV]de l’état transposé et parcourt les tokens dans l’ordre. - Les différentes tuiles de value, têtes et séquences s’exécutent indépendamment.
- Le kernel effectue directement la récurrence : atténuation de l’état, réduction de prédiction sur la key, calcul du residual, écriture par produit extérieur et réduction de lecture sur la query.
- Il convient au décodage, où un seul nouveau token arrive à la fois, mais il ne transforme pas les opérations vectorielles en grandes multiplications matricielles efficaces sur Tensor Cores, ce qui le rend défavorable à l’entraînement et aux longs préfills.
Chunkwise KDA : réorganiser la récurrence en multiplications matricielles
- Chunkwise KDA doit traiter (C) tokens ensemble tout en produisant exactement le même état et les mêmes sorties que le mode récurrent token par token.
- Chaque chunk calcule deux résultats.
- (S_{c+1}), l’état après traitement de tout le chunk à partir de l’état entrant (S_c)
- Les sorties causales de tous les tokens à l’intérieur du chunk
- La difficulté centrale tient au fait que l’erreur delta de chaque token dépend des écritures précédentes dans le même chunk.
-
Atténuation cumulée et erreurs temporaires
- On note (D_i) l’atténuation diagonale du token (i), et (D_{0:i}=D_0D_1\cdots D_i) l’atténuation cumulée depuis la frontière du chunk jusqu’au token (i).
- Lorsqu’une écriture du token (j) est propagée jusqu’au token (i), (D_{j+1:i}) s’applique ; comme les matrices diagonales commutent, les matrices d’atténuation sont échangeables entre elles.
- On calcule d’abord en parallèle des erreurs temporaires qui ignorent les autres écritures internes au chunk.
- (\bar e_i=\beta_i(v_i-S_cD_{0:i}k_i))
- Sauf pour le premier token, ces erreurs temporaires omettent l’effet des écritures internes précédentes au chunk et ne peuvent donc pas être utilisées telles quelles.
-
Restauration de la dépendance causale
- On définit le coefficient par lequel le token précédent (j) affecte l’erreur du token courant (i) comme (\rho_{ij}=\beta_i k_j^\mathsf TD_{j+1:i}k_i).
- L’erreur réelle possède une dépendance séquentielle de la forme (e_i=\bar e_i-\sum_{j<i}\rho_{ij}e_j).
- En plaçant (\rho_{ij}) dans une matrice strictement triangulaire inférieure (R_c), la matrice des erreurs empilées se calcule par (E_c=\bar E_c(A_c^{kk})^\mathsf T), (A_c^{kk}=(I+R_c)^{-1}).
- Aucune inverse dense générale n’est nécessaire.
- (I+R_c) est une matrice triangulaire à diagonale unitaire.
- Il suffit d’effectuer une résolution triangulaire causale pour chaque canal de value.
-
Calcul de l’état de fin de chunk
- L’état entrant traverse toutes les atténuations du chunk, et chaque écriture interne au chunk ne traverse que les atténuations qui la suivent.
- Si l’on empile en lignes dans (K_c^{\mathrm{end}}) les keys atténuées jusqu’à la fin du chunk, l’état peut s’écrire comme la multiplication matricielle suivante.
- (S_{c+1}=S_cD_{0:C-1}+E_cK_c^{\mathrm{end}})
- Plusieurs écritures par produits extérieurs de rang 1 sont fusionnées en une seule multiplication matricielle, ce qui fait avancer l’état du chunk entier d’un seul coup.
-
Calcul de toutes les sorties internes au chunk
- KDA lit après avoir écrit le token courant, donc la sortie du token (i) inclut aussi sa propre écriture.
- On définit le coefficient par lequel une écriture précédente (j) affecte la query (i) comme (\chi_{ij}=s,k_j^\mathsf TD_{j+1:i}q_i), avec (j\le i).
- Les coefficients sont placés dans la matrice de lecture triangulaire inférieure (A_c^{qk}).
- Les zéros de la partie triangulaire supérieure bloquent la contribution des tokens futurs.
- Les éléments diagonaux reflètent le fait que le token courant est lu après sa propre écriture.
- En empilant dans (Q_c^{\mathrm{boundary}}) les vecteurs atténués depuis la frontière jusqu’à chaque query, la sortie totale est :
- (O_c=sS_cQ_c^{\mathrm{boundary}}+E_c(A_c^{qk})^\mathsf T)
- La première multiplication matricielle lit l’état entrant du chunk après atténuation ; la seconde ajoute la contribution causale des écritures internes au chunk.
Pipeline Triton Chunkwise
- L’implémentation par chunks n’est pas un unique kernel géant, mais un pipeline composé de plusieurs appels de kernels.
- Elle calcule d’abord l’atténuation logarithmique cumulée à l’intérieur du chunk.
- La différence de deux prefix sums permet de représenter (D_{j+1:i}) sans multiplier de longs vecteurs de conservation.
- Elle construit ensuite les matrices d’interaction causales (A^{qk}) et (A^{kk}), puis utilise (A^{kk}) pour former une représentation WY des écritures corrigées du chunk.
- Le kernel d’état effectue l’unique parcours entre chunks.
- Il génère l’état entrant de chaque chunk.
- Il résout les erreurs delta du chunk.
- Une fois l’état entrant calculé, le kernel de sortie peut traiter en parallèle les tokens des différents chunks et tuiles.
- L’implémentation réelle calcule d’abord des blocs d’interaction diagonaux de 16 tokens, puis exécute un kernel fusionné pour les blocs non diagonaux et la résolution triangulaire.
chunk_kda_fwdorchestre les étapes ; les principaux points d’entrée sontchunk_kda_fwd_intra,chunk_gated_delta_rule_fwd_hetchunk_gla_fwd_o_gk.- Dans le code,
v_newest l’erreur résolue. hest l’état entrant du chunk.kgest la key atténuée jusqu’à la fin du chunk.
- Dans le code,
- Le mode récurrent et le mode par chunks ne sont pas des attentions différentes, mais deux plannings d’exécution de la même récurrence KDA.
- Le mode récurrent correspond à des opérations vectorielles sérielles pour le décodage à faible latence.
- Le mode par chunks correspond à des opérations matricielles centrées sur les Tensor Cores pour l’entraînement et le préfill.
1 commentaires
Avis sur Hacker News
Ces quinze dernières années, le machine learning a eu besoin d’une notation mathématique unifiée, et en aura probablement encore besoin. Avant, c’était pire : les articles de chercheurs du monde entier utilisaient des notations plus extravagantes les unes que les autres.
Quand la notation change d’un article à l’autre, cela crée de la friction dans la compréhension. Au moins, cet article explique explicitement sa notation dès le départ, ce qui est assez rare dans les articles. Au début, je n’avais même pas remarqué la fonction de changement de notation, mais elle est très utile.
∣q⟩plutôt que des symboles d’une seule lettre ou des types de données explicites. Elle a sans doute l’avantage d’être concise, mais écrire les formules sous forme de pseudo-code ou dans un vrai langage de programmation comme Python les rendrait, à mon avis, beaucoup plus faciles à comprendre.k,qetS, mais sans ce contexte, une grande partie du texte reste opaque.Dire « j’aurais pu y penser moi-même… », c’est oublier que créer ou combiner quelque chose qui n’existait pas est extrêmement difficile.
Dès que quelqu’un finit par publier un travail difficile, les réactions du genre « ce n’est pas si dur » ou « j’aurais pu le faire » apparaissent aussitôt, et tout commence à sembler simple. Il arrive souvent, en développant, de penser avoir inventé quelque chose de nouveau, puis de découvrir plus tard que cela existait déjà dans les années 1970 et était largement utilisé. C’est simplement que cela n’avait jamais croisé mon chemin, donc j’ignorais son existence.
Pour moi, la notation bra-ket rend tout simple et intuitif. Avec la notation vectorielle, je me demandais toujours quel côté était horizontal ou vertical, je finissais par suivre les blocs sans comprendre et je perdais ma concentration ; avec la notation bra-ket, l’ensemble m’a paru très intuitif.
Je pense convertir d’autres articles vers cette notation, car j’ai sans doute manqué beaucoup de bons textes. Pour référence, j’ai un doctorat en physique et une légère dyslexie.
Quand je vois un style comme « Le produit extérieur est une matrice et le produit intérieur est un nombre. Au lieu de stocker toutes les clés et valeurs passées, on stocke la somme des produits extérieurs dans un état de taille fixe
S_t», je suis convaincu que c’est un texte écrit par un LLM.–), on obtient ce genre de résultat.Il existe aussi un tutoriel visualisé : https://snowchord.com/blog/linear-attention-visualized/
Chaque fois que je vois ce genre d’article et de titre, je ressens une profonde gratitude et beaucoup d’humilité envers les très nombreuses personnes bien plus intelligentes que moi. Au lycée et à l’université, j’étais considéré comme très intelligent, et je suis plus malin que la moyenne, mais il existe certainement des millions de personnes capables de me faire passer pour un débutant.
Par intelligence, j’entends ici la capacité à garder en tête et à raisonner sur des concepts et des systèmes immenses et complexes ; cela semble être un talent particulièrement important pour les mathématiciens.
Une expérience de pensée que j’ai faite avec un ami autour d’un verre consistait à isoler les enfants des contenus grand public fournis par les écrans et les algorithmes, et à les élever dans un environnement propice à l’apprentissage où la qualité des médias et des ressources serait strictement contrôlée, comme lorsqu’on entraîne des modèles de pointe. Une sorte de monastère pour enfants, où l’on enseignerait les connaissances les plus récentes sur le réel à travers les mathématiques, l’ingénierie, l’informatique, le deep learning, etc.
Au final, pour repousser les frontières du savoir à l’aide d’outils d’IA avancés, il faudra encore des personnes très intelligentes et dont la pensée n’est pas trop contaminée. L’idée que l’IA remplacera totalement les humains me semble aller dans la mauvaise direction.
Pour référence, le nom notation bra-ket vient effectivement de « bracket », c’est-à-dire les crochets/parenthèses.
https://en.wikipedia.org/wiki/Bra-ket_notation
Au début, j’étais hésitant, mais la notation ket rend les opérations beaucoup plus claires, et elle m’a plu. Cela dit, j’aurais apprécié un bref rappel sur certaines variables, comme
d_kdans l’attention quadratique.Au début, j’étais découragé de ne pas avoir trouvé cette solution, puis je me suis souvenu que j’avais déjà du mal à écrire moi-même une recherche binaire en JavaScript, et je me suis aussitôt senti mieux. Il n’y a absolument aucune chance que j’aie pu imaginer Kimi Delta Attention.
Les boucles dépassent rarement deux ou trois niveaux de profondeur, et si c’est plus complexe que cela, mieux vaut de toute façon déléguer à une bibliothèque.