Aller au contenu

Attention Residuals (AttnRes)

C'est l'innovation la plus conceptuellement élégante du rapport : appliquer l'idée de l'attention à la profondeur du réseau.

L'analogie fondatrice

Le rapport construit son argument en trois temps :

  1. Un RNN compresse tout le passé temporel dans un unique état caché. C'est un goulot d'étranglement.
  2. Le Transformer a résolu ce goulot par l'attention : chaque position accède sélectivement à toutes les positions précédentes, avec des poids dépendant des données.
  3. Or la connexion résiduelle standard a exactement le même défaut, mais selon la profondeur : \(\mathbf{h}_l = \mathbf{h}_{l-1} + f_{l-1}(\mathbf{h}_{l-1})\) compresse toutes les couches précédentes dans un unique vecteur \(\mathbf{h}_l\).

La proposition

Appliquer le même remède : chaque couche récupère sélectivement des représentations de toutes les couches précédentes, au lieu de les accumuler uniformément.

Attention Residuals applies the same methodology to depth.

Full AttnRes : la version idéale

Pour chaque couche \(l\), on définit :

  • une pseudo-requête apprise \(\mathbf{q}_l = \mathbf{w}_l \in \mathbb{R}^{d}\) — un simple vecteur de paramètres, pas une projection de l'entrée ;
  • des clés et valeurs qui sont les sorties des couches précédentes :
\[ \mathbf{k}_{i} = \mathbf{v}_{i} = \begin{cases} \mathbf{h}_1 & i = 0 \quad \text{(l'embedding du jeton)} \\ f_i(\mathbf{h}_{i}) & 1 \leq i \leq l-1 \quad \text{(la sortie de la couche } i) \end{cases} \]

Les poids suivent un noyau softmax :

\[ \phi(\mathbf{q}, \mathbf{k}) = \exp\left(\mathbf{q}^\top\operatorname{RMSNorm}(\mathbf{k})\right) \]
\[ {\alpha_{i \to l}} = \frac{\phi(\mathbf{q}_{l}, \mathbf{k}_{i})}{\sum_{j=0}^{l-1} \phi(\mathbf{q}_{l}, \mathbf{k}_{j})}, \qquad \mathbf{h}_{l} = \sum_{i=0}^{l-1} {\alpha_{i \to l}} \cdot \mathbf{v}_{i} \]

Table des symboles

Symbole Signification
\(\mathbf{w}_l\) Pseudo-requête, vecteur appris propre à la couche \(l\)
\(f_i(\mathbf{h}_i)\) Sortie de la couche \(i\) (attention + MoE)
\(\mathbf{h}_1\) Embedding du jeton, toujours disponible comme source \(i = 0\)
\(\alpha_{i \to l}\) Poids que la couche \(l\) accorde à la couche \(i\)
\(L\) Profondeur totale, \(L = 93\) chez K3

Trois points de conception à comprendre

1. Pourquoi la RMSNorm dans le noyau ?

Sans elle, une couche dont la sortie a une grande amplitude dominerait mécaniquement les poids d'attention, indépendamment de sa pertinence. La normalisation force la comparaison à porter sur la direction, pas sur la norme.

Le rapport : the RMSNorm prevents layers with large-magnitude outputs from dominating the weights.

2. Pourquoi une pseudo-requête apprise et non calculée ?

Une vraie requête serait \(\mathbf{W}_q\mathbf{h}\), donc dépendante du jeton : chaque jeton choisirait ses propres profondeurs. Ici, \(\mathbf{w}_l\) est un paramètre fixe par couche. Le choix des profondeurs est donc appris une fois pour toutes, non adapté par jeton.

Le rapport ne justifie pas ce choix. Deux raisons plausibles : le coût (une requête par jeton et par couche multiplierait le calcul), et la stabilité (un routage en profondeur dépendant du jeton est un mécanisme supplémentaire à stabiliser).

C'est un point ouvert intéressant pour qui voudrait améliorer l'architecture.

3. Pourquoi l'embedding est-il toujours une source ?

Parce que \(\mathbf{h}_1\) est la seule représentation non transformée du jeton. Toutes les couches peuvent donc revenir à l'entrée brute, quel que soit leur niveau — un chemin direct de bout en bout pour le gradient.

Le coût

\(O(L^2 d)\) opérations arithmétiques. Avec \(L < 100\), c'est négligeable devant le reste du modèle.

Le vrai coût n'est pas le calcul

C'est la mémoire : il faut garder vivantes toutes les sorties de couches, soit \(O(Ld)\) par jeton. Et sous parallélisme de pipeline, ces représentations doivent traverser les frontières d'étages, ce qui multiplie la communication inter-étages.

Block AttnRes : la version déployée

Pour réduire ce coût, on regroupe les \(L\) couches en \(N\) blocs de \(S = L/N\) couches.

La construction

À l'intérieur d'un bloc \(n\), les sorties sont réduites par simple somme :

\[ \mathbf{b}_n = \sum_{j \in \mathcal{B}_n} f_j(\mathbf{h}_j) \]

et \(\mathbf{b}_n^i\) note la somme partielle sur les \(i\) premières couches du bloc. On pose \(\mathbf{b}_0 = \mathbf{h}_1\), l'embedding.

L'attention complète ne s'applique alors qu'aux \(N\) représentations de blocs :

\[ \mathbf{V} = \begin{cases} [\mathbf{b}_0, \mathbf{b}_1, \ldots, \mathbf{b}_{n-1}]^\top & \text{si } i = 1 \text{ (première couche du bloc } n) \\ [\mathbf{b}_0, \mathbf{b}_1, \ldots, \mathbf{b}_{n-1}, \mathbf{b}_n^{i-1}]^\top & \text{si } i \geq 2 \end{cases} \]

Lecture

Chaque couche voit : toutes les représentations des blocs terminés, plus la somme partielle du bloc courant. À l'intérieur d'un bloc, on revient donc à un résidu additif classique ; entre blocs, on a une attention complète.

C'est un compromis résolution/coût : de la sélectivité fine entre blocs, de l'accumulation simple à l'intérieur.

La couche de sortie finale agrège les \(N\) représentations de blocs.

Le gain

Full AttnRes Block AttnRes
Mémoire / communication \(O(Ld)\) \(O(Nd)\)
Sources visibles par couche jusqu'à 92 jusqu'à 9

Chez Kimi K3, \(L = 93\), \(N = 8\), \(S = 12\) → réduction d'un facteur ~10.

Un bénéfice inattendu : l'inférence

La structure en blocs borne l'état à l'inférence. Les résultats inter-blocs (parallélisables, calculés une fois) peuvent être fusionnés avec les sommes partielles intra-bloc (séquentielles) via un softmax en ligne (online softmax) — la même technique que FlashAttention.

Le rapport parle d'une réduction significative du coût d'inférence. C'est ce qui permet les deux noyaux spécialisés décrits en Noyaux d'inférence.

Les paramètres exacts chez K3

Grandeur Valeur Source
Taille de bloc \(S\) 12 couches attn_res_block_size: 12
Nombre de blocs 8 (le dernier partiel : \(93 = 7\times12+9\)) Rapport
Sources totales avec l'embedding 9 Rapport
Valeur de \(N\) recommandée \(\approx 8\) Papier AttnRes

Le rapport indique qu'empiriquement, \(N \approx 8\) « récupère l'essentiel du bénéfice à toutes les échelles de modèle ».

Où AttnRes réapparaît dans le reste du rapport

AttnRes n'est pas un ajout isolé. Il irrigue trois autres parties :

Endroit Rôle d'AttnRes
Mémoire d'entraînement La représentation de bloc est générée une fois à la couche frontière et partagée ; tout le calcul AttnRes est encapsulé dans un checkpoint, si bien que l'activation sauvegardée est identique à celle d'un résidu standard
Communication de pipeline Seuls les nouveaux blocs sont transférés entre étages, puis libérés dès la fin du micro-lot — le rapport revendique « la borne inférieure théorique » d'empreinte mémoire
Modèle brouillon L'entrée du brouillon EAGLE-3 fusionne les sorties des 1er, 4e et dernier blocs AttnRes — bas, moyen et haut niveau
Études de cas Le noyau AttnRes est l'une des quatre cibles d'optimisation GPU ; Kimi K3 en a réduit la latence de 283,6 ms à 114,4 ms

Une boucle qui se referme

Kimi K3 a été utilisé pour optimiser le noyau GPU d'un composant de sa propre architecture. C'est le genre de détail qui indique un cycle d'ingénierie mature — et le rapport précise qu'« un checkpoint précoce de K3 traitait déjà la majeure partie de notre travail d'optimisation de noyaux pendant le développement tardif ».

Vérification de compréhension

AttnRes ajoute-t-il beaucoup de paramètres ?

Presque aucun : une pseudo-requête \(\mathbf{w}_l \in \mathbb{R}^{7168}\) par couche, soit \(93 \times 7\,168 \approx 667\,000\) paramètres au total — 0,00002 % du modèle. C'est une modification structurelle, pas paramétrique. C'est précisément ce qui la rend intéressante : le gain ne vient pas de capacité supplémentaire.

Pourquoi sommer à l'intérieur d'un bloc plutôt que faire de l'attention partout ?

Pour borner la mémoire et la communication à \(O(Nd)\) au lieu de \(O(Ld)\). L'hypothèse implicite est que les couches voisines produisent des représentations similaires — les distinguer apporte donc peu, alors que distinguer des blocs éloignés apporte beaucoup. C'est cohérent avec le constat empirique que \(N \approx 8\) suffit.

Quel est le lien entre AttnRes et les architectures DenseNet ?

Même famille d'idée : connecter chaque couche à toutes les précédentes. Mais DenseNet concatène les sorties (la dimension croît), alors qu'AttnRes fait une moyenne pondérée apprise et normalisée (la dimension reste constante). La sélectivité par attention est ce qui distingue les deux.


Chapitre précédent : Gated MLA et NoPE · Chapitre suivant : Stable LatentMoE