Implémenter KDA¶
Ce chapitre donne une implémentation de référence de Kimi Delta Attention, vérifiée, et explique comment tester la vôtre.
Stratégie
Écrire d'abord la forme récurrente naïve. Elle est lente mais évidente à lire, et sert de référence de correction pour toutes les optimisations ultérieures (forme par blocs, KCP, noyaux fusionnés).
C'est aussi la seule façon d'attraper les bugs : une forme par blocs qui se trompe donne des résultats plausibles mais faux.
La forme récurrente de référence¶
def kda_step(S, k, v, alpha, beta):
"""Un pas de la récurrence KDA (équation 1 du rapport).
S : état récurrent, matrice d_k × d_v
k : clé, d_k, normalisée L2
v : valeur, d_v
alpha : facteurs de rétention par canal, d_k, dans (e^-5, 1)
beta : force d'écriture scalaire, dans (0, 1)
"""
S = diagmul(alpha, S) # (1) oublier, canal par canal
kk = matmul(outer(k, k), S) # (2) effacer la composante alignée sur k
S = add(S, scale(kk, -beta))
return add(S, scale(outer(k, v), beta)) # (3) écrire la nouvelle association
La sortie brute est ensuite \(\tilde{\mathbf{o}}_t = \mathbf{S}_t^\top \mathbf{q}_t\).
Trois pièges à l'implémentation
- L'ordre des opérations compte. La décroissance \(\operatorname{Diag}(\alpha)\) s'applique avant l'effacement delta, pas après. Inverser les deux donne un modèle différent qui s'entraînera quand même — mal.
- \(\mathbf{k}\) doit être normalisé L2. Sans cela, \(\mathbf{I} - \beta\mathbf{k}\mathbf{k}^\top\) n'est pas une projection bien conditionnée et la récurrence peut diverger.
- La sortie lit l'état APRÈS mise à jour (\(\mathbf{S}_t\), pas \(\mathbf{S}_{t-1}\)). C'est pourquoi la forme par blocs conserve la diagonale de \(\operatorname{Tril}\).
La paramétrisation¶
def kda_project(x, W):
"""q, k, v depuis l'état caché. ShortConv (noyau 4) puis Swish,
puis L2Norm pour q et k seulement."""
q = l2norm(swish(short_conv(matmul(W.q, x), kernel=4)))
k = l2norm(swish(short_conv(matmul(W.k, x), kernel=4)))
v = swish(short_conv(matmul(W.v, x), kernel=4))
beta = sigmoid(dot(W.beta, x)) # scalaire dans (0,1)
z = matmul(W.a_up, matmul(W.a_down, x)) + W.a_bias # rang faible + biais par tête
return q, k, v, beta, z
La décroissance bornée — le point clé de K3¶
G_MIN = -5.0 # config.json : gate_lower_bound
def kda_decay(z, A_h):
"""Sigmoïde mise à l'échelle : borne la log-décroissance par le bas.
Kimi Linear utilisait g = -exp(A) * softplus(z) → non borné
Kimi K3 utilise g = g_min * sigmoid(exp(A) * z) → borné
"""
g = [G_MIN * sigmoid(math.exp(A_h) * zi) for zi in z] # dans (g_min, 0)
return [math.exp(gi) for gi in g] # alpha dans (e^-5, 1)
Pourquoi cette borne, en une ligne de calcul
\(16 \times 5 = 80\), et \(e^{80} \approx 5{,}5\times10^{34} < 3\times10^{38}\) (limite BF16). Une tuile de 16 jetons ne peut donc jamais déborder, et toutes les tuiles causales passent par les Tensor Cores.
Avec \(g_{\min} = -10\) : \(e^{160}\) → débordement. Le choix est dicté par le format numérique, pas par une recherche d'hyperparamètre.
\(A_h\) est une échelle log apprise par tête, initialisée à 0. Le biais \(\mathbf{b}_\alpha^h\) donne à chaque tête un régime d'oubli par défaut distinct.
La porte de sortie de rang plein¶
def kda_output(o_tilde, x, W):
"""Éq. 5 : RMSNorm par tête, puis porte sigmoïde, puis projection."""
return matmul(W.o, elemwise_mul(
sigmoid(matmul(W.g, x)), # porte de RANG PLEIN (K3), pas rang faible
rmsnorm_per_head(o_tilde)))
Vérifier votre implémentation¶
Trois tests, du plus simple au plus révélateur.
Test 1 — comportements limites¶
# beta = 0 → aucune écriture, seulement décroissance
S_next = kda_step(S, k, v, alpha, beta=0.0)
assert S_next == diagmul(alpha, S)
# alpha = 1 et beta = 1 → la mémoire remplace exactement l'entrée pour k
S_next = kda_step(zeros(dk, dv), k, v, ones(dk), beta=1.0)
o = matvec_T(S_next, k) # relire avec la même clé
assert allclose(o, v) # on retrouve v exactement
Test 2 — l'identité de composition (KCP) ✅ vérifié¶
C'est le test qui compte : il valide la structure affine de la récurrence.
def kda_transition(k, beta, alpha):
"""M_t = (I - beta k k^T) Diag(alpha)"""
D = diag(alpha)
P = add(eye(DK), scale(outer(k, k), -beta))
return matmul(P, D)
# 1. recurrence complete depuis un etat entrant arbitraire S0
direct = run(tokens, S0)
# 2. etat genere depuis ZERO + transition cumulee appliquee a S0
from_zero = run(tokens, zeros(DK, DV))
M = eye(DK)
for (k, v, a, b) in tokens:
M = matmul(kda_transition(k, b, a), M)
composed = add(from_zero, matmul(M, S0))
assert maxabs(direct, composed) < 1e-9
Résultat mesuré
Erreur maximale : 6,94 × 10⁻¹⁸ — la précision machine.
L'identité est exacte, pas approchée. C'est le fondement de KCP.
Test 3 — le balayage préfixe sur plusieurs rangs ✅ vérifié¶
# Chaque rang calcule ses DEUX fragments localement, sans communication
frags = []
for seg in segments:
S_local = run(seg, zeros(DK, DV))
M = eye(DK)
for (k, v, a, b) in seg:
M = matmul(kda_transition(k, b, a), M)
frags.append((M, S_local))
# after all-gather : balayage prefixe associatif
S = zeros(DK, DV)
for (M, S_local) in frags:
S = add(matmul(M, S), S_local)
assert maxabs(S, run(sequence_complete, zeros(DK, DV))) < 1e-9
Résultat mesuré
Erreur maximale : 1,39 × 10⁻¹⁷ sur 4 rangs.
Le parallélisme de contexte de KDA donne exactement le même résultat qu'une récurrence séquentielle unique. C'est une réorganisation, pas une approximation.
La loi de composition associative¶
C'est ce qui rend le balayage préfixe possible :
L'associativité est LA condition
Tout algorithme de scan parallèle exige une opération associative. Sans elle, pas de parallélisation possible — on serait condamné au séquentiel.
Cette même propriété fonde les sommes cumulées parallèles, Mamba, et toutes les architectures récurrentes modernes entraînables efficacement.
La forme par blocs¶
Une fois la forme récurrente validée, la forme par blocs est l'optimisation qui donne les performances. Ne l'écrivez qu'après, et testez-la contre la forme récurrente.
La transformation UT n'est pas dans le papier K3
Le terme \(\widetilde{\mathbf{V}}_{[t]} = \mathbf{U}_{[t]}-\mathbf{W}_{[t]}\mathbf{S}_{[t]}\) provient de la transformation UT, qui linéarise la chaîne de projections delta à l'intérieur d'un bloc.
Le rapport K3 renvoie explicitement à Kimi Linear (arXiv 2510.26692) pour sa dérivation complète. Il faut donc lire ce papier-là pour implémenter la forme par blocs.
Notre test 2 ci-dessus ne couvre pas cette partie : il valide la structure affine de la récurrence, pas la dérivation UT.
Protocole de test recommandé : générer une séquence aléatoire, calculer la sortie par la forme récurrente et par la forme par blocs, exiger un écart inférieur à \(10^{-5}\) en FP32. Tout écart supérieur signale un bug, pas une imprécision.
Les paramètres à dimensionner¶
| Élément | Forme | Params (couche K3) |
|---|---|---|
| \(\mathbf{W}_q, \mathbf{W}_k, \mathbf{W}_v\) | \(12288 \times 7168\) chacune | \(3 \times 88{,}1\) M |
| \(\mathbf{W}_o\) | \(7168 \times 12288\) | 88,1 M |
| \(\mathbf{W}_g\) (porte, rang plein) | \(12288 \times 7168\) | 88,1 M |
| \(\mathbf{W}_\alpha^{\downarrow}, \mathbf{W}_\alpha^{\uparrow}\) | rang faible | négligeable |
| \(\mathbf{W}_\beta\) | \(96 \times 7168\) | négligeable |
| ShortConv | noyau 4 par canal | négligeable |
| Total | ~440 M |
Vérification de compréhension¶
Pourquoi tester d'abord la forme récurrente naïve ?
Parce qu'elle est la seule dont la correction est évidente à la lecture. Un modèle entraîné avec une forme par blocs subtilement fausse converge quand même — moins bien — et le bug reste invisible pendant des semaines. Toute optimisation doit être testée contre une référence dont on est sûr.
Que teste réellement le test 2 ?
Que la récurrence est affine en \(\mathbf{S}\) : \(\mathbf{S}_t = \mathbf{M}_t \mathbf{S}_{t-1} + \mathbf{B}_t\).
Si votre implémentation avait une non-linéarité cachée dans la mise à jour d'état — une erreur classique —, l'identité échouerait. C'est un test très sensible pour trois lignes de code.
Chapitre précédent : Retrouver les paramètres · Chapitre suivant : Implémenter AttnRes et LatentMoE