Aller au contenu

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 :

  1. il partitionne la séquence entre les SM d'un seul rang ;
  2. il évalue les transitions de segments en parallèle ;
  3. 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