Aller au contenu

Kimi Delta Attention (KDA)

C'est le module central de Kimi K3 : 69 couches sur 93, soit 74 % de la profondeur du modèle.

Prérequis

Le chapitre Attention linéaire et récurrence des Fondations. Ce chapitre en reprend les notions sans les redémontrer.

L'équation fondamentale

\[ \boxed{ \mathbf{S}_t = \left(\mathbf{I}-\beta_t\mathbf{k}_t\mathbf{k}_t^{\top}\right) \operatorname{Diag}(\mathbf{\alpha}_t)\,\mathbf{S}_{t-1} + \beta_t\mathbf{k}_t\mathbf{v}_t^{\top}, \qquad \tilde{\mathbf{o}}_t = \mathbf{S}_t^{\top}\mathbf{q}_t } \]

Table des symboles

Symbole Type Signification Valeur chez K3
\(t\) entier Position du jeton —
\(\mathbf{x}_t\) \(\mathbb{R}^{d}\) État caché en entrée \(d = 7168\)
\(\mathbf{q}_t, \mathbf{k}_t\) \(\mathbb{R}^{d_k}\) Requête et clé \(d_k = 128\) par tête
\(\mathbf{v}_t\) \(\mathbb{R}^{d_v}\) Valeur \(d_v = 128\) par tête
\(\mathbf{S}_t\) \(\mathbb{R}^{d_k \times d_v}\) État récurrent (la mémoire) \(128\times128\) par tête
\(\mathbf{\alpha}_t\) \((0,1)^{d_k}\) Facteur de rétention par canal —
\(\beta_t\) \((0,1)\) Force d'écriture (scalaire) —
\(\mathbf{I}\) \(\mathbb{R}^{d_k \times d_k}\) Matrice identité —
\(\tilde{\mathbf{o}}_t\) \(\mathbb{R}^{d_v}\) Sortie brute, avant porte —

Lecture de l'équation, terme par terme

L'état se met à jour en trois opérations enchaînées :

     S_{t-1}
        │
        ├──► ① Diag(α_t) · S_{t-1}          OUBLIER, canal par canal
        │       chaque ligne de S est multipliée par un α différent
        │
        ├──► ② (I − β_t k_t k_tᵀ) · (…)     EFFACER ce qui concerne k_t
        │       projection qui retire la composante alignée sur k_t,
        │       avec une intensité β_t
        │
        └──► ③ + β_t k_t v_tᵀ                ÉCRIRE la nouvelle association
                                              (k_t → v_t), intensité β_t

Intuition en une phrase

KDA gère une mémoire associative de taille fixe : à chaque jeton, elle laisse un peu s'estomper l'ancien (①), retire spécifiquement l'entrée correspondant à la clé courante (②), puis inscrit la nouvelle valeur (③).

Sans ②, on empilerait les souvenirs et ils interféreraient. Sans ①, la mémoire n'oublierait jamais et saturerait. Sans ③, elle n'apprendrait rien.

La paramétrisation

Comment \(\mathbf{q}, \mathbf{k}, \mathbf{v}, \beta, \mathbf{\alpha}\) sont-ils produits à partir de \(\mathbf{x}_t\) ?

\[ \begin{aligned} \mathbf{q}_t^h,\mathbf{k}_t^h &= \operatorname{L_2Norm}\!\left(\operatorname{Swish}\!\left(\operatorname{ShortConv}\!\left(\mathbf{W}_{q/k}^h\mathbf{x}_t\right)\right)\right) \in \mathbb{R}^{d_k} \\[4pt] \mathbf{v}_t^h &= \operatorname{Swish}\!\left(\operatorname{ShortConv}\!\left(\mathbf{W}_v^h\mathbf{x}_t\right)\right) \in \mathbb{R}^{d_v} \\[4pt] \beta_t^h &= \operatorname{Sigmoid}\!\left(\mathbf{W}_{\beta}^h\mathbf{x}_t\right) \in (0,1) \\[4pt] \mathbf{z}_t^h &= \mathbf{W}_{\alpha}^{\uparrow}\mathbf{W}_{\alpha}^{\downarrow}\mathbf{x}_t + \mathbf{b}_{\alpha}^h \in \mathbb{R}^{d_k} \end{aligned} \]

Ce que chaque opération apporte

ShortConv — une convolution causale sur une petite fenêtre temporelle (short_conv_kernel_size: 4 chez K3). Chaque canal est mélangé avec ses 3 prédécesseurs immédiats.

Pourquoi une convolution ?

L'attention linéaire est mauvaise sur les motifs locaux très courts (« le jeton précédent était une parenthèse ouvrante »). Une convolution de taille 4 les capture pour un coût dérisoire, et libère la récurrence pour les dépendances longues. C'est un héritage de Mamba et Gated DeltaNet.

Swish — non-linéarité, quasi linéaire pour les positifs.

L2Norm sur \(\mathbf{q}\) et \(\mathbf{k}\) — étape essentielle. La règle delta utilise \(\mathbf{I} - \beta\mathbf{k}\mathbf{k}^\top\), qui n'est une projection bien conditionnée que si \(\|\mathbf{k}\|_2 = 1\). Sans normalisation, l'opérateur peut amplifier l'état au lieu de l'atténuer, et la récurrence diverge.

Projection de rang faible pour \(\mathbf{z}\) — la décroissance a besoin d'un logit par canal de clé (128 valeurs par tête, soit 12 288 pour 96 têtes). Une matrice pleine \(7168 \times 12288\) coûterait 88 M de paramètres par couche. La factorisation \(\mathbf{W}_{\alpha}^{\uparrow}\mathbf{W}_{\alpha}^{\downarrow}\) ramène ce coût à quelques millions. Un biais par tête \(\mathbf{b}_{\alpha}^h\) permet à chaque tête d'avoir un régime de décroissance par défaut différent — certaines têtes « à mémoire longue », d'autres « à mémoire courte ».

Innovation K3 n° 1 : la décroissance bornée par le bas

C'est la modification la plus conséquente par rapport à Kimi Linear.

Le problème

Sous forme par blocs, l'algorithme calcule une décroissance cumulée \(\mathbf{\gamma}^{1\to r} = \prod_{i=1}^{r}\mathbf{\alpha}_i\) et divise les clés par elle. Comme \(\mathbf{\alpha} \in (0,1)\), ce produit tend vers zéro et son inverse vers l'infini.

Kimi Linear contournait le problème en deux temps : calculer la décroissance relative en espace logarithmique, et découper chaque bloc en tuiles secondaires de 16 jetons.

Le coût caché de ce contournement

Les tuiles hors diagonale peuvent alors utiliser des multiplications matricielles denses sur Tensor Cores. Mais les tuiles diagonales nécessitent encore un calcul explicite paire-de-positions par paire-de-positions — une opération qui n'utilise pas les Tensor Cores.

Le rapport identifie ce chemin comme the main intra-chunk bottleneck.

La solution de K3

Changer la fonction qui transforme les logits \(\mathbf{z}_t^h\) en log-décroissance \(\mathbf{g}_t^h\).

Kimi Linear (GDN, Mamba-2) Kimi K3
Formule \(\mathbf{g}_t^h=-e^{A_h}\operatorname{Softplus}(\mathbf{z}_t^h)\) \(\mathbf{g}_t^h = g_{\min}\operatorname{Sigmoid}(e^{A_h}\mathbf{z}_t^h)\)
Image \((-\infty, 0)\) — non bornée \((g_{\min}, 0)\) — bornée
Rétention \(\alpha = e^{g}\) \((0, 1)\) \((e^{g_{\min}}, 1)\)

Avec \(g_{\min} = -5\) (fixé, confirmé par gate_lower_bound: -5.0) :

\[ \alpha_{t,j}^h > e^{-5} \approx 6{,}7\times10^{-3} \]

Sur une tuile de 16 jetons, la log-décroissance cumulée est donc bornée :

\[ 16 \times (-5) = -80 \;\le\; \sum_{i=1}^{16} g_i \;\le\; 0 \]

Le facteur de renormalisation inverse est donc inférieur à \(e^{80} \approx 5{,}5\times10^{34}\) — dans la plage dynamique de BF16 (qui va jusqu'à \(\sim 3\times10^{38}\)).

Le gain réel

Le gain n'est pas seulement d'éviter un débordement. Comme plus aucune tuile ne peut déborder, toutes les tuiles causales, diagonale comprise, peuvent passer par des multiplications matricielles denses sur Tensor Cores.

Le chemin lent paire-de-positions disparaît complètement. C'est un exemple de co-conception algorithme–système : on modifie une équation pour débloquer une unité matérielle.

\(A_h\) est une échelle logarithmique apprise par tête, initialisée à 0. Le biais \(\mathbf{b}_{\alpha}^h\) suit l'initialisation de Kimi Linear / Mamba-2 / GDN.

Parenté

Le rapport note que cette paramétrisation est proche des portes récurrentes bornées par le bas de HGRN2, Griffin et RWKV-7. L'idée de borner une porte récurrente n'est donc pas neuve ; son application ici pour débloquer les Tensor Cores l'est davantage.

Innovation K3 n° 2 : la porte de sortie de rang plein

Kimi Linear utilisait une porte de sortie de rang faible. Kimi K3 passe à une projection de rang plein dépendante de l'entrée :

\[ \mathbf{y}_t = \mathbf{W}_o\!\left[\operatorname{Sigmoid}(\mathbf{W}_g\mathbf{x}_t) \odot \operatorname{RMSNorm}(\tilde{\mathbf{o}}_t)\right] \]

Ordre des opérations :

  1. RMSNorm par tête sur la sortie récurrente \(\tilde{\mathbf{o}}_t\) ;
  2. multiplication coordonnée par coordonnée par la porte \(\operatorname{Sigmoid}(\mathbf{W}_g\mathbf{x}_t)\) ;
  3. projection de sortie \(\mathbf{W}_o\).

\(\mathbf{W}_g\) est de taille \(12\,288 \times 7\,168\), soit 88 M de paramètres par couche KDA — soit un cinquième des paramètres de la couche, consacrés uniquement au filtrage de sa propre sortie.

Pourquoi c'est justifié

Le rapport cite les travaux sur l'attention à porte (Gated Attention, Qiu et al., 2025), qui montrent que ces portes apportent non-linéarité, parcimonie, et suppriment le phénomène d'attention sink.

Intuition : la mémoire récurrente contient beaucoup de choses ; la porte laisse chaque jeton décider, canal par canal, de ce qu'il en lit réellement. À 96 têtes et 128 canaux, une porte de rang faible ne peut pas exprimer des choix indépendants entre canaux — d'où le passage au rang plein.

Confirmé par la configuration : use_full_rank_gate: true.

La forme par blocs (chunkwise)

Pour l'entraînement et le prefill, la récurrence est reformulée en calcul par blocs, exact mathématiquement.

Pour un bloc de taille \(C\), avec \(\mathbf{\Gamma}_{[t]}^{1\to C}\) la matrice des décroissances cumulées :

\[ \begin{aligned} \mathbf{A}_{[t]} &= \operatorname{Tril}\!\left[(\mathbf{Q}_{[t]}\odot \mathbf{\Gamma}_{[t]}^{1\rightarrow C})(\mathbf{K}_{[t]}/\mathbf{\Gamma}_{[t]}^{1\rightarrow C})^{\top}\right] \\[4pt] \mathbf{O}_{[t]} &= \underbrace{(\mathbf{\Gamma}_{[t]}^{1\rightarrow C}\odot\mathbf{Q}_{[t]})\mathbf{S}_{[t]}}_{\text{inter-blocs}} + \underbrace{\mathbf{A}_{[t]}\widetilde{\mathbf{V}}_{[t]}}_{\text{intra-bloc}} \end{aligned} \]
  • \(\operatorname{Tril}\) met à zéro la partie strictement triangulaire supérieure. La diagonale est conservée, car chaque sortie lit l'état après la mise à jour du jeton courant.
  • \(\widetilde{\mathbf{V}}_{[t]} := \mathbf{U}_{[t]}-\mathbf{W}_{[t]}\mathbf{S}_{[t]}\) est le terme de « pseudo-valeur », issu de la transformation UT, qui linéarise la chaîne de projections delta à l'intérieur du bloc. Sa dérivation complète est dans le papier Kimi Linear, pas dans celui de K3.
  • Le terme inter-blocs apporte l'information des blocs précédents ; le terme intra-bloc les interactions internes.

Erreur fréquente

Croire que la forme par blocs est une approximation. Elle est exacte : même résultat, ordre d'opérations différent, choisi pour saturer les unités matricielles.

Ce que KDA coûte et rapporte

Grandeur Valeur
Paramètres par couche KDA ~440 M
État récurrent par tête \(128 \times 128 = 16\,384\) valeurs
État récurrent par couche et par requête ~1,57 M valeurs (~3,1 Mio en BF16)
État total (69 couches) par requête ~217 Mio, indépendant de \(T\)
Coût de calcul \(O(T)\)

À comparer à un cache KV MLA qui, à 1 M de jetons, se compte en dizaines de gigaoctets.

Les conséquences système (renvois)

KDA crée quatre problèmes d'infrastructure, chacun traité dans sa page :

Problème Solution Page
Récurrence série vs parallélisme GPU FlashKDA, noyau CUTLASS par blocs Noyaux KDA
Séquence trop longue pour un GPU KDA Context Parallelism (KCP) KCP
État mis à jour en place vs décodage spéculatif Rejeu des entrées projetées Noyaux d'inférence
État volumineux vs cache de préfixe fin Découplage des granularités Cache de préfixe

Vérification de compréhension

Que se passe-t-il si \(\beta_t = 0\) pour un jeton ?

L'opérateur d'effacement devient l'identité et le terme d'écriture s'annule : \(\mathbf{S}_t = \operatorname{Diag}(\mathbf{\alpha}_t)\mathbf{S}_{t-1}\). Le jeton n'écrit rien en mémoire ; il ne fait que laisser l'état décroître. C'est le comportement attendu pour un jeton sans information nouvelle (ponctuation, remplissage).

Pourquoi \(\mathbf{\alpha}_t\) est-il un vecteur alors que \(\beta_t\) est un scalaire ?

\(\mathbf{\alpha}_t\) agit sur les lignes de \(\mathbf{S}\), indexées par les canaux de clé : chaque canal peut avoir sa propre vitesse d'oubli. C'est ce que le rapport appelle channel-wise forget gate, et c'est l'apport principal de KDA sur Gated DeltaNet.

\(\beta_t\) contrôle l'intensité d'un événement d'écriture unique (l'association \(\mathbf{k}_t \to \mathbf{v}_t\)), pour lequel une seule valeur suffit.

Pourquoi \(g_{\min} = -5\) et pas \(-10\) ou \(-3\) ?

Le rapport donne la contrainte, pas l'optimisation. Avec des tuiles de 16 jetons, il faut \(16 \cdot |g_{\min}|\) dans la plage de l'exponentielle BF16 : \(|g_{\min}| = 5 \Rightarrow e^{80}\) (sûr) ; \(|g_{\min}| = 10 \Rightarrow e^{160}\) (dépassement).

En sens inverse, \(g_{\min}\) trop proche de 0 empêcherait tout oubli rapide. \(-5\) est donc le plus grand oubli possible compatible avec BF16 sur des tuiles de 16. Le choix est dicté par le format numérique, pas par une recherche d'hyperparamètre.


Chapitre précédent : L'attention hybride 3:1 · Chapitre suivant : Gated MLA et NoPE