Noyaux d'inférence¶
Kimi K3 introduit trois modules architecturaux nouveaux — KDA, Block AttnRes et Stable LatentMoE. Chacun exige un noyau d'inférence sur mesure.
1. Le décodage KDA¶
Un goulot différent du prefill¶
Le changement de régime
Au prefill, le problème est d'exploiter le parallélisme. Au décodage, il devient de gérer efficacement l'état récurrent évolutif, mis à jour en place à chaque pas.
Le conflit avec le décodage spéculatif¶
Cette mise à jour en place devient problématique dès qu'on utilise le décodage spéculatif par MTP :
Le problème
Si la vérification rejette une partie des jetons proposés, l'état a déjà avancé au-delà du dernier jeton accepté, et ne peut pas être ramené en arrière de façon triviale.
Maintenir un instantané d'état par position de brouillon permettrait le retour en arrière — mais multiplierait le trafic d'état, un coût qui domine aux grandes tailles de lot typiques du service en ligne.
La solution : rejouer plutôt que sauvegarder¶
L'observation clé
L'état après n'importe quel préfixe accepté de jetons brouillons est entièrement déterminé par les entrées projetées de ces jetons — qui sont bien plus petites que l'état lui-même.
Donc : ne mettre en cache que ces entrées projetées, reconstruire sur puce les états des jetons acceptés, et réécrire les états des jetons vérifiés et du jeton bonus.
Approche naïve : Approche de K3 :
sauvegarder l'état S après sauvegarder les entrées projetées
chaque jeton brouillon (petites) des jetons brouillons
│ │
▼ ▼
7 × 217 Mio de trafic d'état quelques Kio
│
▼
REJOUER la récurrence sur puce
pour reconstruire l'état accepté
Le rapport signale que cette conception a été proposée indépendamment dans le travail concurrent ReplaySSM — une convergence qui suggère qu'il s'agit de la bonne solution.
La fusion¶
Les jetons rejoués, le jeton bonus et la fenêtre de brouillon suivante partagent une seule boucle récurrente à l'intérieur d'un unique noyau fusionné qui couvre :
- la convolution courte ;
- la normalisation d'entrée ;
- le gating ;
- la récurrence KDA ;
- la normalisation de sortie.
Les propriétés obtenues
- La latence de vérification croît de façon sous-linéaire avec le nombre de jetons vérifiés, et reste inférieure aux références à cache d'état.
- Comme les caches de projection ne quittent jamais l'étage de décodage, le cache de préfixe et la désagrégation prefill/decode opèrent sur exactement la même charge utile qu'en service non spéculatif.
Ce second point est essentiel : le décodage spéculatif n'introduit aucune complication supplémentaire dans le système de cache déjà fort complexe.
2. Block AttnRes¶
Le schéma en deux phases¶
Block AttnRes suit un ordonnancement en deux temps :
- une passe inter-blocs par lot, qui lit les représentations de blocs mises en cache une fois par bloc ;
- après quoi chaque couche intègre la somme partielle intra-bloc par une fusion en softmax en ligne.
L'accès mémoire représente une fraction substantielle du coût de ces noyaux, au prefill comme au décodage. Les optimisations visent donc avant tout l'efficacité mémoire.
Au prefill : le parallélisme de séquence¶
Le problème
Matérialiser les représentations de blocs sur chaque rang de parallélisme tensoriel entraînerait une consommation mémoire redondante substantielle.
Solution : adopter le parallélisme de séquence pour les activations.
L'all-reduce du TP est décomposé en un reduce-scatter et un all-gather,
avec le noyau intra-bloc inséré entre les deux collectives, opérant sur des
états cachés partitionnés selon la séquence.
Le résultat
Les représentations de blocs de chaque jeton sont matérialisées sur exactement un rang. Cela élimine la consommation mémoire additionnelle et réduit les coûts d'entrée/sortie de Block AttnRes au prefill.
all-reduce classique :
[ calcul ] ─► all-reduce ─► [ calcul ] ← chaque rang a TOUT
décomposition K3 :
[ calcul ] ─► reduce-scatter ─► [ noyau intra-bloc ] ─► all-gather ─► [ calcul ]
↑
opère sur des données partitionnées
→ une seule copie par jeton
Au décodage : recouvrement et fusion¶
| Phase | Optimisation |
|---|---|
| Inter-blocs | Lancé sur un flux latéral, pour recouvrir du calcul indépendant sur le flux principal |
| Intra-bloc | Fusionné : la fusion de la sortie AttnRes avec sa mise à jour de somme partielle, ainsi que la RMSNorm subséquente, sont intégrées dans l'all-reduce TP précédent — supprimant un noyau dédié |
Ensemble, ces optimisations cachent la latence de la passe inter-blocs et réduisent le trafic mémoire de la phase intra-bloc.
3. Stable LatentMoE¶
Le problème¶
Stable LatentMoE augmente à la fois le nombre total d'experts et le nombre d'experts activés par jeton. Cette croissance conjointe fait monter les surcoûts d'ordonnancement et de coordination, rendant difficile pour les noyaux MoE conventionnels de maintenir une utilisation matérielle élevée.
Les GEMM latents : trois optimisations¶
| # | Optimisation |
|---|---|
| 1 | Fusionner la projection latente descendante avec le routeur MoE en un seul GEMM |
| 2 | Partitionner les matrices de poids latentes entre rangs, et fusionner l'all-gather de sortie dans l'épilogue du GEMM via des instructions multimem store |
| 3 | Recouvrir la communication résultante avec d'autres opérateurs, comme le calcul des experts partagés |
L'effet combiné
Élimination du trafic de poids redondant et du calcul dupliqué, tout en cachant la latence de communication derrière du calcul.
Le point 1 est particulièrement net : \(\mathbf{W}^{\downarrow}\) et le routeur lisent tous deux le même vecteur \(\mathbf{x}\). Les fusionner en une seule multiplication supprime une lecture complète des activations.
Les experts routés : WarpDecode¶
Le régime problématique
À petite taille de lot, les GEMM groupés se réduisent à du streaming de matrices de poids limité par la mémoire — un régime pour lequel les noyaux conventionnels centrés sur les tuiles sont mal adaptés, à cause de leur conception orientée calcul et de leurs surcoûts de prétraitement.
La solution : construire le noyau de décodage MoE sur la conception centrée sur les jetons de WarpDecode, dans laquelle chaque warp est responsable d'un neurone de sortie et diffuse les poids associés directement depuis la mémoire.
Deux raffinements ajoutés par Kimi K3 :
- Subdiviser chaque warp en équipes de voies (lane teams) plus fines, chacune traitant un sous-ensemble disjoint d'experts, suivies d'une réduction à l'échelle du warp des résultats partiels. Cela augmente le parallélisme.
- Permuter la disposition des poids hors ligne, à un coût de prétraitement unique, ce qui réduit substantiellement le surcoût de déquantification à l'exécution.
Pourquoi la déquantification coûte cher
Les poids sont stockés en MXFP4. Chaque lecture exige de les déballer et de leur appliquer le facteur d'échelle du bloc. À 16 experts × 33 M de paramètres par couche, cette opération est répétée des milliards de fois. Réorganiser la disposition hors ligne pour que le déballage suive le motif d'accès naturel du noyau élimine l'essentiel de ce coût.
Récapitulatif¶
| Module | Problème d'inférence | Solution |
|---|---|---|
| KDA (décodage) | État mis à jour en place vs retour arrière spéculatif | Cacher les entrées projetées, rejouer l'état sur puce, noyau fusionné |
| Block AttnRes (prefill) | Représentations de blocs dupliquées par rang TP | Parallélisme de séquence, noyau inséré entre reduce-scatter et all-gather |
| Block AttnRes (décodage) | Latence de la passe inter-blocs | Flux latéral + fusion dans l'all-reduce |
| LatentMoE (GEMM latents) | Trafic de poids redondant | Fusion routeur+\(\mathbf{W}^{\downarrow}\), partitionnement, multimem store |
| LatentMoE (experts) | Streaming limité par la mémoire à petit lot | WarpDecode centré jetons, équipes de voies, permutation hors ligne |
Vérification de compréhension¶
Pourquoi le décodage est-il limité par la mémoire et pas par le calcul ?
Parce qu'à petite taille de lot, chaque poids lu depuis la HBM ne sert qu'à quelques opérations. Le GPU passe l'essentiel de son temps à attendre les données. C'est l'inverse du prefill, où chaque poids est réutilisé pour des milliers de jetons — donc limité par le calcul.
Qu'est-ce qu'un « warp » et pourquoi organiser le noyau autour ?
Un warp est un groupe de 32 threads GPU qui exécutent la même instruction en parallèle. C'est l'unité d'ordonnancement matérielle réelle. Un noyau centré sur les tuiles pense en blocs de matrice ; un noyau centré sur les jetons pense en « quel warp produit quel neurone de sortie ». Dans un régime limité par la mémoire, cette seconde vue correspond mieux au motif d'accès effectif.
Le rejeu d'état de KDA ne coûte-t-il pas du calcul supplémentaire ?
Si, mais très peu : rejouer la récurrence sur 7 jetons brouillons est trivial comparé à la lecture répétée de 217 Mio d'états. Et le rapport précise que le rejeu se fait sur puce, en mémoire partagée, donc sans aller-retour vers la HBM. C'est exactement le bon arbitrage dans un régime limité par la mémoire.
Chapitre précédent : Cache de préfixe hybride · Chapitre suivant : Ordonnancement de flotte