Aller au contenu

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
C'est la seule représentation non transformée du jeton, et un chemin direct de bout en bout pour le gradient.

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
Le biais doit réguler l'aiguillage sans altérer les poids de mélange ni l'optimisation par gradient du routeur. C'est ce qui rend la méthode « sans perte auxiliaire ».

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