Noyaux KDA et FlashKDA¶
Le conflit fondamental¶
Le problème en une phrase
KDA remplace le cache clé–valeur croissant de l'attention softmax par un état récurrent de taille fixe \(\mathbf{S}\in\mathbb{R}^{d_k\times d_v}\).
C'est excellent pour la mémoire. Mais sa mise à jour est série, et le GPU préfère un parallélisme large et uniforme.
Le rapport structure sa réponse selon deux propriétés opposées de l'état :
- sa sérialité est un obstacle → à contourner par des noyaux dédiés ;
- sa taille fixe est un atout → à exploiter (transferts bon marché, réutilisation entre requêtes).
Et selon deux niveaux d'exécution : au sein d'un appareil (noyaux fusionnés) et entre appareils (KCP).
Un noyau par régime d'exécution¶
La sérialité de KDA se manifeste comme un goulot différent dans chaque régime. Kimi K3 conçoit donc un noyau dédié pour chacun.
| Régime | Goulot | Solution |
|---|---|---|
| Entraînement / prefill | Alternance calcul parallèle ↔ propagation série | FlashKDA |
| Prefill ultra-long sous TP | SM inactifs quand chaque rang n'a que peu de têtes | CP intra-appareil |
| Décodage | Gestion de l'état évolutif, mis à jour en place | Voir Noyaux d'inférence |
FlashKDA : le noyau par blocs¶
Le problème précis¶
La forme par blocs de KDA est parallèle à l'intérieur d'un bloc mais série entre blocs, puisque l'état doit se propager de bloc en bloc.
Ce qui se passe naïvement
Les deux phases alternent :
[calcul intra-bloc massivement parallèle] ← les SM travaillent
[propagation d'état série] ← les SM sont INACTIFS
[calcul intra-bloc massivement parallèle]
[propagation d'état série] ← les SM sont INACTIFS
...
Une fraction significative du temps, la machine ne fait rien.
La solution¶
FlashKDA est un noyau par blocs fondé sur CUTLASS qui recouvre le calcul intra-bloc et la propagation d'état inter-blocs.
Il décompose le travail en :
- des étages parallèles aux jetons ;
- une récurrence parallèle aux têtes.
Chacun est ordonnancé et réglé indépendamment.
L'idée clé
Les 96 têtes sont indépendantes les unes des autres. Pendant que la récurrence progresse pour la tête 1, le calcul intra-bloc peut avancer pour la tête 2.
La sérialité de KDA est intra-tête ; il reste 96 fils de travail parallèles à exploiter. Le noyau reformule le problème pour que la partie série n'immobilise jamais toute la machine.
Résultat annoncé : dépasse substantiellement l'implémentation de référence Triton.
FlashKDA sert à la fois l'entraînement et le prefill d'inférence, et est auto-dispatché comme backend de la bibliothèque flash-linear-attention.
Le rôle de la décroissance bornée¶
Rappel du lien avec l'architecture
Sans le plancher \(g_{\min} = -5\), les tuiles diagonales exigeraient un calcul explicite paire-de-positions, qui n'utilise pas les Tensor Cores.
Avec le plancher, toutes les tuiles causales passent par des multiplications matricielles denses. Le noyau n'a plus qu'un seul chemin de calcul, ce qui le simplifie autant que ça l'accélère.
Le changement d'équation existe pour rendre ce noyau possible.
Le parallélisme de contexte intra-appareil¶
Le problème¶
Le parallélisme tensoriel partitionne les têtes entre appareils. Mais il ne raccourcit jamais la récurrence : chaque rang doit toujours parcourir toute la séquence.
La conséquence
Sous déploiement en TP pur, prefiller une séquence ultra-longue laisse la plupart des SM inactifs, parce que chaque rang ne détient plus que quelques têtes — donc peu de fils de travail parallèles, pour une récurrence toujours aussi longue.
L'observation clé¶
The state transition of each segment can be evaluated independently of the incoming state and composed exactly afterward.
C'est la même propriété algébrique qui fonde KCP, mais appliquée à l'intérieur d'un seul GPU.
La solution¶
Un planificateur de parallélisme de contexte au niveau des SM, automatique :
- il partitionne la séquence entre les SM d'un seul rang ;
- il évalue les transitions de segments en parallèle ;
- il les fusionne pour récupérer l'état initial exact de chaque segment.
L'intérêt
Contrairement au KCP inter-appareils, ce parallélisme est entièrement intra-appareil et n'engage aucune communication entre GPU. C'est du parallélisme gratuit.
Vérification de compréhension¶
Pourquoi CUTLASS plutôt que Triton ?
Triton est un DSL de haut niveau, productif mais qui laisse moins de contrôle sur l'ordonnancement fin, le pipeline mémoire et le placement en mémoire partagée. CUTLASS est une bibliothèque C++/CUDA de bas niveau qui expose ces leviers. Pour un noyau où le recouvrement calcul/propagation est le cœur du problème, ce contrôle est décisif.
Ironie notable : Kimi K3 a lui-même écrit MiniTriton, un compilateur de type Triton, dont le noyau de prefill KDA « dépasse nettement une référence Triton équivalente ».
Pourquoi la propagation d'état ne peut-elle pas être simplement recouverte par du prefetch ?
Parce qu'il ne s'agit pas d'un transfert de données mais d'une dépendance de calcul : l'état du bloc \(n+1\) ne peut littéralement pas être calculé avant celui du bloc \(n\). Le seul recouvrement possible est avec du travail indépendant — d'où l'exploitation du parallélisme entre têtes.
Chapitre suivant : KDA Context Parallelism