Aller au contenu

Per-Head Muon

L'optimiseur. Ce n'est pas une couche du modèle, mais le rapport le classe dans la partie architecture — et à juste titre : sans lui, cette architecture ne s'entraînerait probablement pas de façon stable à cette échelle.

Rappel : ce qu'est Muon

Les optimiseurs classiques (SGD, Adam) traitent les paramètres comme une longue liste de nombres indépendants. Muon traite les paramètres matriciels comme des matrices.

À chaque pas :

  1. calculer le momentum \(\mathbf{M}\) du gradient (une matrice, même forme que le paramètre) ;
  2. l'orthogonaliser par itérations de Newton–Schulz ;
  3. appliquer la matrice orthogonalisée comme mise à jour.

Que fait l'orthogonalisation ?

Toute matrice se décompose en valeurs singulières : \(\mathbf{M} = \mathbf{U}\mathbf{\Sigma}\mathbf{V}^\top\). Orthogonaliser revient à remplacer \(\mathbf{\Sigma}\) par l'identité :

\[ \operatorname{ortho}(\mathbf{M}) = \mathbf{U}\mathbf{V}^\top \]

Intuition

Un gradient matriciel est généralement dominé par quelques directions (les grandes valeurs singulières). Une mise à jour brute pousse fort dans ces directions et néglige les autres.

L'orthogonalisation égalise toutes les directions : la mise à jour devient une rotation d'amplitude uniforme. On garde la direction du gradient sans son déséquilibre d'échelle.

Newton–Schulz est une itération polynomiale qui approche \(\mathbf{U}\mathbf{V}^\top\) en quelques passes, sans calculer la SVD — cette dernière étant beaucoup trop coûteuse à l'échelle de matrices de millions d'éléments.

Kimi K2 avait déjà démontré Muon à l'échelle du billion de paramètres. Kimi K3 le conserve, avec le weight clipping introduit par K2.

L'apport de K3 : orthogonaliser tête par tête

Le problème

Une projection d'attention \(\mathbf{W}_Q\) chez Kimi K3 est de taille \(12\,288 \times 7\,168\) — mais elle représente en réalité 96 têtes empilées, chacune \(128 \times 7\,168\).

Ce qui se passe si on orthogonalise le tout

L'orthogonalisation en bloc traite les 96 têtes comme un seul objet couplé. Or les têtes n'ont pas les mêmes échelles de gradient : certaines sont très sollicitées, d'autres presque inertes.

Conséquence, telle que formulée par le rapport : heads with larger gradient or momentum scales dominate the shared update direction, while smaller-scale heads receive insufficiently normalized updates.

Les grosses têtes imposent la direction commune ; les petites reçoivent une mise à jour mal normalisée.

La solution

Partitionner les matrices de momentum le long de la dimension des têtes, et orthogonaliser chaque bloc de tête séparément.

        AVANT (Muon standard)              APRÈS (Per-Head Muon)
   ┌──────────────────────────┐      ┌──────────────────────────┐
   │  tête 1                  │      │  tête 1     → ortho ──┐  │
   │  tête 2                  │      │  tête 2     → ortho ──┤  │
   │  …           ortho global│      │  …          → ortho ──┼──► mise à jour
   │  tête 96                 │      │  tête 96    → ortho ──┘  │
   └──────────────────────────┘      └──────────────────────────┘
     une seule SVD implicite            96 SVD implicites, indépendantes

S'applique aux projections \(Q\), \(K\) et \(V\).

Les bénéfices revendiqués

Bénéfice Explication
Échelle de mise à jour égalisée entre têtes Chaque tête est normalisée par sa propre échelle
Dynamique d'apprentissage plus équilibrée Les têtes faibles ne sont plus noyées
Meilleure stabilité aux grandes échelles Constaté empiriquement par l'équipe
Coût d'optimiseur légèrement réduit Newton–Schulz sur 96 blocs élancés est moins cher que sur la matrice complète

Le dernier point est contre-intuitif et vaut d'être noté

On s'attendrait à ce que 96 orthogonalisations coûtent plus cher qu'une. En réalité non : le coût de Newton–Schulz croît de façon superlinéaire avec la dimension. Sur des blocs \(128 \times 7168\) (très élancés), l'itération est plus économique que sur \(12\,288 \times 7\,168\).

Une amélioration de qualité qui est aussi une réduction de coût — cas rare, généralement le signe qu'on a exploité une structure jusqu'alors ignorée.

L'implémentation distribuée

L'optimiseur distribué répartit les paramètres uniformément entre les rangs de parallélisme de données. Mais Newton–Schulz a besoin de la matrice complète (ou, ici, du bloc de tête complet). Il faut donc rassembler avant chaque mise à jour.

L'approche naïve et son coût

Faire un all-gather sur tout le tampon de paramètres, sur chaque rang (comme dans Moonlight). Cela crée une empreinte mémoire considérable et fait de la communication le goulot d'étranglement principal à l'échelle.

La solution de K3 : chaque rang ne récupère que les tranches des paramètres dont il est propriétaire, par communication P2P avec les rangs détenteurs.

all-gather naïf P2P ciblé (K3)
Tampon de paramètres complet Requis sur chaque rang Éliminé
Volume de communication \(O(N \times R)\) \(O(N)\)
Recouvrement Difficile Pipeliné par tampon de morceau de modèle

Communication et calcul sont en outre pipelinés à la granularité des tampons de morceaux de modèle, ce qui cache le surcoût de communication.

Le reste de la recette d'optimisation

Élément Valeur
Optimiseur Per-Head Muon (paramètres matriciels)
Weight clipping Oui, mécanisme introduit dans Kimi K2
Weight decay 0,1 constant
Calendrier de LR Cosinus, warmup linéaire 1 %
Équilibrage MoE Quantile Balancing

Ce qui n'est pas publié

Ni le taux d'apprentissage de pointe, ni la taille de lot, ni le nombre d'itérations de Newton–Schulz, ni les détails du weight clipping. Le rapport indique seulement que ces valeurs ont été retunées par des études de loi d'échelle dédiées, sans donner les résultats.

Voir Ce qui est public et ce qui ne l'est pas.

Vérification de compréhension

Pourquoi Muon ne s'applique-t-il qu'aux paramètres matriciels ?

Parce que l'orthogonalisation n'a de sens que pour une matrice. Les paramètres vectoriels — gains de RMSNorm, biais, plongements — n'ont pas de structure singulière à égaliser. Ils sont traités par un optimiseur classique de type AdamW.

Per-Head Muon s'applique-t-il aussi aux experts MoE ?

Le rapport ne l'indique pas ; il parle spécifiquement des « projections d'attention ». Les matrices d'experts (\(3584 \times 3072\)) n'ont pas de structure en têtes, donc l'argument ne s'y transpose pas. Elles sont vraisemblablement traitées par Muon standard.

Quel lien avec la stabilité observée dans le reste du rapport ?

Kimi K3 accumule les mécanismes de stabilisation : SiTU-GLU borne les activations, la RMSNorm du LatentMoE borne les échelles, la décroissance bornée de KDA borne la plage numérique, MoonViT depuis zéro borne les normes de gradient, et Per-Head Muon borne l'écart d'échelle entre têtes.

La stabilité est le fil rouge de tout le travail architectural de K3. À 2,8 T de paramètres, une divergence coûte des semaines de calcul : il est rationnel d'investir massivement dans des garanties structurelles plutôt que dans des correctifs a posteriori.


Chapitre précédent : Vision native : MoonViT-V2

Fin de la partie Architecture. Suite : Entraînement.