Implémenter AttnRes et Stable LatentMoE¶
Block AttnRes¶
C'est le module le plus simple de Kimi K3 à implémenter, et celui qui apporte le plus de structure pour le moins de paramètres.
La structure¶
class BlockAttnRes:
def __init__(self, n_layers=93, block_size=12, d=7168):
self.S = block_size
# UNE pseudo-requête apprise par couche. C'est tout.
self.w = [randn(d) for _ in range(n_layers)]
self.block_repr = [] # b_0 … b_{n-1} : sorties de blocs terminés
self.partial = None # b_n^{i-1} : somme partielle du bloc courant
def read(self, layer_idx):
"""Ce qu'une couche lit, à la place du simple h_{l-1}."""
sources = list(self.block_repr) # blocs terminés
if self.partial is not None:
sources.append(self.partial) # + somme partielle courante
q = self.w[layer_idx]
# noyau softmax avec RMSNorm sur les CLÉS
logits = [dot(q, rmsnorm(s)) for s in sources]
a = softmax(logits)
return sum(ai * si for ai, si in zip(a, sources))
def write(self, layer_idx, output):
"""Accumule la sortie de couche dans la représentation du bloc."""
self.partial = output if self.partial is None else add(self.partial, output)
if (layer_idx + 1) % self.S == 0: # fin de bloc
self.block_repr.append(self.partial)
self.partial = None
Les trois détails à ne pas manquer¶
1. La RMSNorm porte sur les clés, pas sur la sortie
\(\phi(\mathbf{q}, \mathbf{k}) = \exp(\mathbf{q}^\top\operatorname{RMSNorm}(\mathbf{k}))\)
Sans elle, une couche dont la sortie a une grande amplitude domine mécaniquement les poids, indépendamment de sa pertinence. La normalisation force la comparaison à porter sur la direction.
Mais la moyenne pondérée porte sur les valeurs NON normalisées : les clés sont normalisées, les valeurs ne le sont pas.
2. \(\mathbf{b}_0\) est l'embedding, toujours présent
self.block_repr = [token_embedding] # b_0, initialisé AVANT la couche 1
3. La pseudo-requête est un PARAMÈTRE, pas une projection
\(\mathbf{q}_l = \mathbf{w}_l\), un vecteur appris, pas \(\mathbf{W}_q\mathbf{h}\). Le choix des profondeurs est appris une fois pour toutes, il n'est pas adapté par jeton.
Le dimensionnement chez K3¶
| Paramètre | Valeur |
|---|---|
| Taille de bloc | 12 (attn_res_block_size) |
| Nombre de blocs | 8 (\(93 = 7\times12 + 9\), dernier partiel) |
| Sources maximales | 9 (8 blocs + embedding) |
| Paramètres ajoutés | \(93 \times 7168 \approx 667\) K, soit 0,00002 % du modèle |
Le rapport coût/bénéfice
AttnRes est une modification structurelle, pas paramétrique. Le gain ne vient pas de capacité supplémentaire mais d'un meilleur acheminement de l'information. C'est ce qui la rend intéressante.
Stable LatentMoE¶
La couche complète¶
def latent_moe(x, W, experts, shared, router_bias, k=16):
"""Éq. 10 du rapport.
x : état caché, dimension d = 7168
"""
# --- branche partagée : pleine largeur, toujours active ---------
y = sum(E(x) for E in shared) # 2 experts, d → d
# --- routage ----------------------------------------------------
s = sigmoid(matmul(W.router, x)) # 896 scores INDÉPENDANTS
top = argtopk(add(s, router_bias), k) # le biais entre ICI…
p = normalize([s[j] for j in top]) # …mais PAS dans les poids
# --- branche routée : espace latent -----------------------------
z = matmul(W.down, x) # 7168 → 3584
u = sum(pj * experts[j](z) for pj, j in zip(p, top))
u = rmsnorm(u) # ← AJOUT DE K3, essentiel
y = add(y, matmul(W.up, u)) # 3584 → 7168
return y
Les deux erreurs qui cassent tout
1. Mettre le biais dans les poids de mélange.
p = normalize([s[j] + router_bias[j] for j in top]) # ❌ FAUX
p = normalize([s[j] for j in top]) # ✅ correct
2. Oublier la RMSNorm avant \(\mathbf{W}^{\uparrow}\). C'est l'ajout spécifique de K3, et le rapport indique qu'il est nécessaire pour éviter les explosions d'activation dans la chaîne de quatre multiplications matricielles quasi consécutives, à 2,8 T de paramètres.
SiTU-GLU dans chaque expert¶
B1, B2 = 4.0, 25.0 # activation_situ_beta, activation_situ_linear_beta
def situ_glu(x, Wg, Wu):
"""Éq. 11 : plafond doux sur les DEUX branches."""
g = matmul(Wg, x)
u = matmul(Wu, x)
gate = elemwise(lambda t: B1 * tanh(t / B1), g) * sigmoid(g)
up = elemwise(lambda t: B2 * tanh(t / B2), u)
return gate * up
Test de la borne ✅ vérifié
Sur 200 000 tirages dans \([-500, 500]^2\) :
max|SiTU-GLU| = 100.000 ≤ β₁ × β₂ = 100
La borne est atteinte exactement et jamais dépassée. C'est une garantie structurelle, pas statistique.
Et près de l'origine, l'écart relatif à SwiGLU est inférieur à 0,53 % — l'accord au premier ordre est bien vérifié.
Quantile Balancing¶
def qb_update(S, b, k, n):
"""Une mise à jour QB. AUCUN hyperparamètre.
S : matrice de scores, m tokens × n experts
b : biais courant, n
"""
m = len(S)
alphas = []
for i in range(m):
# Top-(k+1) : les k premiers sont les routes, le (k+1)-ième est le seuil
order = sorted(range(n), key=lambda j: -(S[i][j] + b[j]))
alphas.append(S[i][order[k]] + b[order[k]])
# b_j = -quantile_{1-k/n}(s_:,j - alpha)
nb = [-quantile([S[i][j] - alphas[i] for i in range(m)], 1.0 - k / n)
for j in range(n)]
mean = sum(nb) / n
return [x - mean for x in nb] # centrage : invariant pour le Top-k
La causalité
La mise à jour ne prend effet qu'au pas suivant. Un lot n'est jamais routé avec un biais dérivé de lui-même — sinon le mécanisme « verrait » sa propre distribution.
L'estimateur par histogramme, pour l'échelle réelle¶
Le quantile exact sur des millions de marges réparties sur des centaines de rangs est impossible dans la boucle d'entraînement. On histogramme le biais requis \(r_{i,j} := \alpha_i - s_{i,j}\) :
B = 1000 # nombre de casiers
def qb_histogram_update(S_local, alphas_local, b, k, n):
# bornes : s ∈ (0,1) et alpha ∈ (b_min, 1+b_max) ⟹ r ∈ [b_min-1, b_max+1]
lo, hi = min(b) - 1.0, max(b) + 1.0
w = (hi - lo) / B
H = [[0] * B for _ in range(n)] # comptes par expert
for i, ai in enumerate(alphas_local): # accumulation LOCALE, sans communication
for j in range(n):
idx = clamp(int((ai - S_local[i][j] - lo) / w), 0, B - 1)
H[j][idx] += 1
H = all_reduce_sum(H) # UNE seule communication par couche et par pas
q = total_tokens * k / n
out = []
for j in range(n):
c = 0
for bin_idx in range(B):
if c + H[j][bin_idx] >= q:
frac = clamp((q - c) / H[j][bin_idx], 0.0, 1.0)
out.append(lo + (bin_idx + frac) * w)
break
c += H[j][bin_idx]
mean = sum(out) / n
return [x - mean for x in out]
Les trois propriétés
| Propriété | Pourquoi |
|---|---|
| Exact | Les cumuls sont exacts aux bords de casiers ; erreur ≤ largeur de casier. Avec \(B = 1000\) : quelques \(10^{-3}\) |
| Bon marché | Un all-reduce d'entiers de \(nB\) valeurs, indépendant du nombre de jetons |
| Correct | Les comptes sont additifs, donc l'histogramme global est invariant au partitionnement. On obtient le quantile du lot global, pas une moyenne de quantiles par rang |
Résultat de notre reproduction¶
Ce que nous avons mesuré
Sur 4 096 jetons et 16 experts (cible : 512 jetons par expert) :
QB (sans hyperparamètre) déséquilibre relatif : 10,5 → 6,1 → 4,3 → 3,9 → 3,5 → 2,7 %
règle par signe (γ=1e-3) déséquilibre relatif : 10,5 → 7,4 → 4,9 → 3,1 → 3,5 → 2,9 %
Confirmé : QB réduit fortement le déséquilibre dès le premier pas, et ne demande aucun réglage.
Non reproduit : l'avantage net sur la règle par signe. À cette échelle, les deux méthodes convergent de façon comparable — la règle par signe bénéficiant ici d'un \(\gamma\) choisi à la main.
Le régime où le rapport situe l'avantage de QB — 896 experts, lots de millions de jetons — n'est pas atteignable dans cette reproduction. Nous rapportons ce que nous mesurons, y compris là où cela ne confirme pas le rapport.
Vérification de compréhension¶
Pourquoi la sigmoïde plutôt que le softmax pour le routeur ?
Un softmax force les scores à sommer à 1 : les experts sont mis en compétition au niveau du score, ce qui rend l'ajout d'un biais additif difficile à interpréter. Une sigmoïde donne des scores indépendants : le biais devient un simple déplacement de seuil par expert, ce qui est exactement le cadre de la dérivation duale de QB.
Que se passe-t-il si on oublie le centrage des biais ?
Fonctionnellement rien : la sélection Top-\(k\) est invariante par translation globale. Mais les biais dériveraient sans borne au fil des pas, ce qui élargirait indéfiniment la plage de binning de l'histogramme, et donc dégraderait la résolution du quantile.
Chapitre précédent : Implémenter KDA · Chapitre suivant : Un modèle jouet : nano-K3