Aller au contenu

KDA Context Parallelism (KCP)

Comment entraîner sur des séquences d'un million de jetons quand la séquence ne tient pas sur un GPU.

Le point de départ : linéaire contre softmax

Le coût de communication du parallélisme de contexte diffère fondamentalement entre les deux familles d'attention.

Attention softmax Attention linéaire
Ce qui doit circuler Des blocs clé–valeur L'état récurrent
Taille Croît avec la longueur Fixe (\(d_k \times d_v\))
Méthode de référence Ring Attention LASP, LASP-2

L'avantage structurel de l'attention linéaire

Une séquence de 1 M partitionnée sur 8 rangs : en Ring Attention, il faut faire circuler des centaines de mégaoctets de KV entre rangs. En attention linéaire, il suffit d'échanger une matrice de taille fixe.

Pourquoi la méthode standard échoue sur KDA

Les méthodes antérieures exploitent la récurrence additive de l'attention linéaire vanille :

\[ \mathbf{S}_t = \mathbf{S}_{t-1} + \mathbf{k}_t\mathbf{v}_t^\top \]

Chaque rang calcule l'état que ses jetons locaux génèrent en partant de \(\mathbf{S} = \mathbf{0}\), puis on somme les états locaux des rangs précédents pour retrouver l'état entrant. Simple et exact.

Le problème avec KDA

KDA met à jour son état ainsi :

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

La règle delta applique la matrice dépendante du jeton \(\mathbf{M}_t\) à l'état entrant, avant d'ajouter l'écriture courante.

Conséquence : l'effet d'un segment local dépend de l'état qui entre dans ce segment, et ne peut donc pas être déterminé à partir du seul état calculé depuis \(\mathbf{S}=\mathbf{0}\).

La sommation directe est fausse.

La solution : décomposer en deux quantités locales

L'idée de KCP : décomposer l'effet de chaque segment en deux quantités calculables localement :

  1. une transition cumulée qui agit sur l'état entrant ;
  2. un état généré localement depuis zéro.

Les notations

Notation Signification
\(\mathbf{S}_{[i]}^{t}\) État récurrent dans le segment du rang \(i\), après \(t\) jetons locaux
\(\mathbf{S}_{[i]}^{T_i}\) État sortant du rang \(i\), entrant du rang \(i+1\)
\(\widetilde{\mathbf{S}}_{[i]}^{t}\) Le même état, mais avec la récurrence démarrée à \(\mathbf{S}=\mathbf{0}\)
\(\mathbf{M}_{[i+1]}^{t \leftarrow 1}\) Transition cumulée des \(t\) premiers jetons locaux, \(\prod_{r=1}^{t}\mathbf{M}_r\)

La formule de composition

\[ \begin{aligned} \mathbf{S}_{[i+1]}^{t} & =\widetilde{\mathbf{S}}_{[i+1]}^{t} + \mathbf{M}_{[i+1]}^{t \leftarrow 1}\mathbf{S}_{[i]}^{T_i} \\[6pt] & = \widetilde{\mathbf{S}}_{[i+1]}^{t} + \mathbf{M}_{[i+1]}^{t \leftarrow 1}\sum_{j=1}^{i}\Big(\prod_{l=j+1}^{i}\mathbf{M}_{[l]}^{T_l \leftarrow 1}\Big)\widetilde{\mathbf{S}}_{[j]}^{T_j} \end{aligned} \]

Lecture des deux termes :

  • \(\widetilde{\mathbf{S}}_{[i+1]}^{t}\) : l'état généré par les jetons locaux ;
  • \(\mathbf{M}_{[i+1]}^{t \leftarrow 1}\mathbf{S}_{[i]}^{T_i}\) : le contexte des rangs précédents, propagé à travers les mises à jour KDA locales.

La propriété qui rend tout possible

À \(t = T_{i+1}\), les deux quantités \(\mathbf{M}_{[i+1]}^{T_{i+1}\leftarrow 1}\) et \(\widetilde{\mathbf{S}}_{[i+1]}^{T_{i+1}}\) peuvent être calculées avec les seuls jetons locaux, avant même que \(\mathbf{S}_{[i]}^{T_i}\) ne soit disponible.

Ce sont exactement les fragments que chaque rang échange avec les autres.

L'algorithme

Sur chaque rang i, en parallèle et sans communication :
  ① calculer M_[i]^{T_i←1}     (transition cumulée locale)
  ② calculer S̃_[i]^{T_i}       (état local depuis zéro)

Une seule communication :
  ③ all-gather des deux tenseurs (taille FIXE)

Sur chaque rang, localement :
  ④ reconstruire l'état entrant par balayage préfixe :
       S ← 0
       pour chaque fragment j précédent, dans l'ordre :
           S ← M_[j]^{T_j←1} · S + S̃_[j]^{T_j}

  ⑤ poursuivre le calcul local normalement

Pourquoi le balayage préfixe fonctionne

Les mises à jour au niveau des rangs se composent de façon associative :

\[ (\mathbf{M}_2, \widetilde{\mathbf{S}}_2) \circ (\mathbf{M}_1, \widetilde{\mathbf{S}}_1) = (\mathbf{M}_2\mathbf{M}_1,\ \mathbf{M}_2\widetilde{\mathbf{S}}_1 + \widetilde{\mathbf{S}}_2) \]

L'associativité est la condition d'existence d'un balayage préfixe parallèle. C'est le même théorème qui fonde tous les scans parallèles, depuis les sommes cumulées jusqu'à Mamba.

Le résultat

Propriété Valeur
Communication Un seul all-gather de taille fixe
Dépendance à la longueur de séquence Aucune
Passage à l'échelle du calcul Linéaire
Exactitude Exacte — pas une approximation

Filiation et disponibilité

La construction s'appuie sur le parallélisme de contexte de DeltaNet (Wang et al., 2025). L'implémentation KDA est disponible dans la PR #691 de flash-linear-attention.

Un détail d'implémentation : la reconstruction traite les fragments précédents du même document, dans l'ordre. Le parallélisme de contexte doit donc rester conscient des frontières de documents dans un lot empaqueté.

Comparaison des trois parallélismes de contexte

Méthode Attention Ce qui circule Taille
Ring Attention Softmax Blocs KV \(O(T)\)
LASP / LASP-2 Linéaire additive États locaux \(O(d_k d_v)\)
KCP KDA (delta + décroissance) Transition + état local \(O(d_k^2 + d_k d_v)\)

Le surcoût de KCP par rapport à LASP

KCP transmet une matrice de transition \(\mathbf{M} \in \mathbb{R}^{d_k \times d_k}\) en plus de l'état. Chez Kimi K3 (\(d_k = d_v = 128\)), cela double approximativement le volume échangé par rapport à LASP.

C'est le prix de la règle delta — et il reste indépendant de la longueur de séquence, ce qui est l'essentiel.

Vérification de compréhension

Pourquoi ne peut-on pas simplement calculer les états séquentiellement entre rangs ?

On le pourrait, mais cela sérialiserait complètement les rangs : le rang 2 attendrait le rang 1, etc. Sur 8 rangs, on n'aurait aucun gain de temps par rapport à un seul GPU — seulement un gain de mémoire. KCP permet à tous les rangs de calculer simultanément leurs fragments, puis de reconstituer.

Que représente concrètement \(\mathbf{M}_{[i]}^{T_i \leftarrow 1}\) ?

C'est la composition de toutes les transformations que le segment du rang \(i\) ferait subir à un état arbitraire qui y entrerait. Une réponse à la question : « si un état quelconque entre ici, dans quel état ressort-il, en ignorant ce que ce segment écrit lui-même ? »

KCP s'applique-t-il aussi aux couches MLA ?

Non — les couches MLA sont des attentions softmax et relèvent des méthodes classiques (le rapport cite DeepSpeed-Ulysses). Une passe avant de Kimi K3 utilise donc deux mécanismes de parallélisme de contexte différents selon le type de couche.


Chapitre précédent : Noyaux KDA et FlashKDA · Chapitre suivant : MoonEP : l'équilibrage parfait