Attention linéaire et récurrence¶
Ce chapitre est le plus important de la partie Fondations : il contient tout ce qu'il faut pour comprendre Kimi Delta Attention, le cœur de Kimi K3.
L'idée : supprimer le softmax pour factoriser¶
Reprenons la sortie de l'attention pour une position \(t\) :
Le problème vient de l'exponentielle : elle couple \(\mathbf{q}_t\) et \(\mathbf{k}_s\) de manière inséparable, il faut donc calculer chaque paire.
Remplaçons \(\exp(\mathbf{q}^\top\mathbf{k})\) par un simple produit scalaire (ou par une fonction de noyau factorisable \(\phi(\mathbf{q})^\top\phi(\mathbf{k})\)). Alors :
Le résultat clé
La somme sur les positions passées peut être précalculée dans une matrice \(\mathbf{S}_t\) de taille fixe \(d_k \times d_v\), indépendante de la longueur de la séquence.
Et cette matrice se met à jour de façon récurrente :
| Propriété | Attention softmax | Attention linéaire |
|---|---|---|
| Coût de calcul total | \(O(T^2 d)\) | \(O(T d^2)\) |
| Mémoire d'état | \(O(T d)\) — croît | \(O(d_k d_v)\) — fixe |
| Expressivité | Maximale | Réduite |
| Parallélisable en entraînement | Oui | Difficile (récurrence) |
Intuition
L'attention softmax est un carnet de notes : on garde tout, on relit tout. L'attention linéaire est une mémoire associative de taille fixe : on écrit chaque nouvelle information par-dessus, en superposition.
Le problème de la mémoire associative saturée¶
\(\mathbf{S}_t = \sum_s \mathbf{k}_s\mathbf{v}_s^\top\) est une mémoire clé–valeur additive. Pour relire l'information associée à une clé \(\mathbf{k}\), on calcule \(\mathbf{S}^\top\mathbf{k}\) : si les clés sont orthonormées, on retrouve exactement \(\mathbf{v}\). Mais dès qu'il y a plus d'écritures que de dimensions disponibles, les souvenirs interfèrent et la mémoire devient du bruit.
Trois remèdes ont été inventés successivement. Kimi K3 les utilise tous les trois.
Remède 1 — La décroissance (decay)¶
On fait oublier l'état ancien :
Après \(n\) pas, une information ancienne est atténuée d'un facteur \(\alpha^n\). La mémoire devient une fenêtre glissante douce. C'est le principe de Mamba, GLA, RWKV.
Kimi K3 pousse plus loin : la décroissance est par canal (\(\mathbf{\alpha}_t \in (0,1)^{d_k}\), une valeur par dimension de clé) et dépendante de l'entrée. Certains canaux peuvent tout retenir, d'autres tout oublier, et ce choix change à chaque jeton.
Remède 2 — La règle delta (delta rule)¶
Ajouter aveuglément est mauvais si la clé \(\mathbf{k}_t\) est déjà en mémoire. La règle delta écrit la différence entre ce qu'on veut mémoriser et ce qui est déjà là :
qu'on réécrit sous forme d'opérateur :
| Symbole | Signification |
|---|---|
| \(\mathbf{I} - \beta_t\mathbf{k}_t\mathbf{k}_t^\top\) | Opérateur qui efface la composante de l'état alignée sur \(\mathbf{k}_t\) |
| \(\beta_t \in (0,1)\) | Force d'écriture : 0 = ne rien écrire, 1 = remplacer complètement |
Intuition
C'est une correction d'erreur. Le modèle lit d'abord ce qu'il a déjà mémorisé pour cette clé, puis n'écrit que l'écart. Une mise à jour, pas un empilement.
Remède 3 — La porte de sortie¶
Même avec une mémoire propre, tout ce qu'on lit n'est pas pertinent. Une porte de sortie dépendant de l'entrée filtre le résultat, canal par canal :
Kimi Delta Attention = les trois remèdes combinés¶
Lecture de gauche à droite :
- \(\operatorname{Diag}(\mathbf{\alpha}_t)\mathbf{S}_{t-1}\) : oublier un peu, canal par canal ;
- \((\mathbf{I}-\beta_t\mathbf{k}_t\mathbf{k}_t^\top)\,\cdot\) : effacer ce qui concerne la clé courante ;
- \(+\;\beta_t\mathbf{k}_t\mathbf{v}_t^\top\) : écrire la nouvelle association.
C'est exactement l'équation 1 du rapport Kimi K3. Le chapitre Kimi Delta Attention en donne la paramétrisation complète et les modifications propres à K3.
Le calcul par blocs (chunkwise)¶
Une récurrence est intrinsèquement séquentielle, ce que les GPU détestent. La solution universelle est la forme par blocs :
- Découper la séquence en blocs de \(C\) jetons (typiquement 64 ou 128).
- Entre blocs : propager l'état \(\mathbf{S}\) séquentiellement — peu d'étapes, donc peu de sérialité.
- Dans un bloc : tout calculer en parallèle par multiplications matricielles, en partant de l'état entrant.
Intuition
On paie un peu de calcul redondant à l'intérieur des blocs pour récupérer le parallélisme massif des unités matricielles (Tensor Cores). C'est le même marché que fait FlashAttention pour l'attention softmax.
Erreur fréquente
Croire que la forme par blocs change le résultat. Non : elle est mathématiquement exacte. Seul l'ordre des opérations diffère.
Le piège numérique de la décroissance cumulée¶
Dans un bloc, on a besoin de la décroissance cumulée \(\mathbf{\gamma}^{1\to r} = \prod_{i=1}^{r}\mathbf{\alpha}_i\), et l'algorithme divise les clés par cette quantité. Or un produit de nombres \(<1\) tend rapidement vers zéro, et son inverse vers l'infini.
Exemple : avec \(\alpha = 0{,}5\) sur 64 jetons, \(\gamma = 0{,}5^{64} \approx 5\times10^{-20}\), et \(1/\gamma \approx 2\times10^{19}\) — au-delà de la précision exploitable en BF16.
La solution de Kimi K3
Borner la décroissance par le bas. Kimi Linear utilisait \(\mathbf{g}_t = -e^{A}\operatorname{Softplus}(\mathbf{z}_t)\), non borné. Kimi K3 utilise une sigmoïde mise à l'échelle :
avec \(g_{\min} = -5\) fixé (confirmé par gate_lower_bound: -5.0 dans la
configuration publiée). Chaque facteur de rétention vaut donc au minimum
\(e^{-5} \approx 6{,}7\times10^{-3}\), et sur une tuile de 16 jetons la
décroissance logarithmique cumulée reste dans \((-80, 0)\) : le facteur de
renormalisation reste sous \(e^{80}\), dans la plage BF16.
Le gain n'est pas seulement numérique : comme plus aucune tuile ne déborde, toutes les tuiles causales peuvent utiliser des multiplications matricielles denses sur Tensor Cores. Le chemin lent « paire de positions » de la diagonale disparaît complètement.
Les architectures hybrides¶
L'attention linéaire est efficace mais moins expressive : sa mémoire de taille fixe finit par perdre de l'information. L'attention softmax est expressive mais coûteuse. La solution retenue par presque tous les modèles longs contexte récents est l'hybridation par couches.
Kimi K3 : 3 couches KDA pour 1 couche Gated MLA, motif répété, avec une couche MLA supplémentaire en fin de pile.
| Couches KDA | Couches Gated MLA | |
|---|---|---|
| Nombre | 69 | 24 |
| Coût en longueur | Linéaire | Quadratique |
| Mémoire par requête | État fixe \(d_k \times d_v\) | Cache KV croissant |
| Rôle | Mélange local, récence, position | Rappel global exact |
| Encodage de position | Implicite (décroissance) | Aucun (NoPE) |
À retenir
Trois quarts des couches ont un coût linéaire. C'est ce qui rend un contexte de 1 M de jetons économiquement viable — pas une astuce d'ingénierie, un choix architectural.
Vérification de compréhension¶
Pourquoi la règle delta empêche-t-elle la parallélisation naïve entre GPU ?
Parce que l'opérateur \(\mathbf{M}_t = (\mathbf{I}-\beta_t\mathbf{k}_t\mathbf{k}_t^\top) \operatorname{Diag}(\mathbf{\alpha}_t)\) s'applique à l'état entrant. L'effet d'un segment dépend donc de ce qui l'a précédé, contrairement à l'attention linéaire additive où les états locaux se somment. C'est précisément le problème que résout KDA Context Parallelism, en transportant deux quantités par segment au lieu d'une : la transition cumulée \(\mathbf{M}\) et l'état généré localement \(\widetilde{\mathbf{S}}\).
Quelle est la taille de l'état récurrent d'une couche KDA de Kimi K3 ?
\(d_k \times d_v = 128 \times 128 = 16\,384\) valeurs par tête, et il y a 96 têtes, soit \(1{,}57\) million de valeurs par couche et par requête — environ 3,1 Mio en BF16. Multiplié par 69 couches : ~217 Mio par requête, quelle que soit la longueur du contexte. À comparer avec un cache KV qui, lui, croît linéairement.
Chapitre précédent : Le Transformer et l'attention · Chapitre suivant : Mixture-of-Experts