Aller au contenu

Quantile Balancing

Le mécanisme qui garantit que les 896 experts de chaque couche reçoivent chacun sa part des jetons.

Rappel du problème

Rien n'oblige le routeur à répartir équitablement. En pratique, il ne le fait pas, et les conséquences sont doubles :

Conséquence Effet
Qualité Un expert peu sollicité est mal entraîné ; un expert jamais sollicité est 33 M de paramètres perdus
Débit En parallélisme d'experts, tous les GPU attendent le plus chargé

L'état de l'art avant K3

Perte auxiliaire

Ajouter un terme au coût qui pénalise le déséquilibre. Défaut : il entre en concurrence avec l'objectif de langage. On dégrade la qualité pour gagner l'équilibre.

Biais sans perte auxiliaire (DeepSeek-V3)

Ajouter un biais \(b_j\) par expert au score, mis à jour par pas fixe :

\[ b_j^{(t+1)}=b_j^{(t)}+\gamma\operatorname{sign}\!\left(\bar{\ell}-\ell_j^{(t)}\right) \]

où \(\ell_j\) est la charge observée et \(\bar\ell\) la charge cible.

Le défaut

Le pas \(\gamma\) arbitre entre adaptation lente (si trop petit) et oscillation de charge (si trop grand). Et le rapport est explicite : maintaining balanced loads becomes more challenging as LatentMoE increases the routed expert pool to 896 per layer. Le régime de ~\(10^3\) experts sort du domaine où ces mises à jour restent bien conduites.

L'idée de Quantile Balancing

En une phrase

Au lieu de tâtonner par petits pas dans la bonne direction, calculer directement le biais qui donne exactement la charge cible.

Notations : un lot de \(m\) jetons routés vers \(n\) experts avec sélection Top-\(k\). La charge cible par expert est

\[ q := \frac{mk}{n} \]

Étape 1 — obtenir les seuils, gratuitement

Le routage passe de Top-\(k\) à Top-\((k{+}1)\) sur le score biaisé \(\mathbf{s}_i+\mathbf{b}^{(t)}\) :

  • les \(k\) premières entrées sont les routes effectivement prises ;
  • la \((k{+}1)\)-ième entrée est le seuil \(\alpha_i^{(t)}\) : le score qu'un expert doit dépasser pour entrer dans le Top-\(k\) du jeton \(i\).

L'astuce

Prendre le seuil depuis un routage Top-\((k{+}1)\) évite un calcul de quantile séparé côté jetons. Une seule passe avant fournit à la fois le routage et les seuils.

Étape 2 — inverser la relation charge ↔ biais

Les seuils étant fixés, le nombre de jetons routés vers l'expert \(j\) sous un biais candidat \(\widehat{b}_j\) vaut :

\[ \sum_{i=1}^{m}\mathbf{1}\!\left[s_{i,j}+\widehat{b}_j > \alpha_i^{(t)}\right] \]

Cette quantité est monotone décroissante en \(-\widehat{b}_j\). Il existe donc une valeur unique qui donne exactement \(q\).

En posant les marges \(s_{i,j}-\alpha_i^{(t)}\) : \(-\widehat{b}_j\) doit être la \((q{+}1)\)-ième plus grande marge, de sorte qu'exactement \(q\) marges restent au-dessus du seuil. Comme \(q/m = k/n\), c'est le quantile d'ordre \(1-k/n\) :

\[ \begin{aligned} \widehat{b}_j^{(t+1)} &\leftarrow -\operatorname{quantile}_{1-k/n}\!\left(\mathbf{s}_{:,j}-\mathbf{\alpha}^{(t)}\right) \\[4pt] \mathbf{b}^{(t+1)} &\leftarrow \widehat{\mathbf{b}}^{(t+1)} - \operatorname{mean}\!\left(\widehat{\mathbf{b}}^{(t+1)}\right)\mathbf{1} \end{aligned} \]

La seconde ligne retire un décalage commun, qui ne change pas la sélection Top-\(k\) (elle est invariante par translation) mais garde les biais centrés.

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. Sans cette précaution, le mécanisme « verrait » sa propre distribution et introduirait une fuite d'information.

Le biais est figé à l'inférence — le routage devient un simple Top-\(k\) déterministe, sans calcul de quantile.

Le fondement théorique

L'annexe du rapport dérive QB depuis le problème d'assignation équilibrée à score maximal :

\[ \max_{x_{i,j}\in\{0,1\}} \sum_{i,j} x_{i,j}s_{i,j} \quad \text{s.c.} \quad \sum_j x_{i,j}=k, \quad \sum_i x_{i,j}=\frac{mk}{n} \]

Étape 1 — relaxation exacte. En relâchant \(x_{i,j}\in[0,1]\), on obtient un programme linéaire dont l'optimum est entier par intégralité du polytope de \(b\)-couplage biparti. La relaxation est donc sans perte.

Étape 2 — dualité. En introduisant des multiplicateurs \(\alpha_i\) (côté jetons) et \(\beta_j\) (côté experts) et en échangeant max et min (théorème du minimax, applicable car tout est linéaire sur des convexes), on obtient le dual convexe :

\[ \min_{\alpha_i,\beta_j}\; \sum_{i,j}\max\big(0,\; s_{i,j} - \alpha_i - \beta_j\big) + k\sum_i \alpha_i + \frac{mk}{n}\sum_j \beta_j \]

Étape 3 — minimisation par coordonnées. Chaque sous-problème admet une solution exacte en forme close. Pour \(\mathbf{\beta}\) fixé, le sous-problème du jeton \(i\) est linéaire par morceaux, de pente \(k\) moins le nombre de marges au-dessus de \(\alpha\) ; il est minimisé exactement quand \(k\) marges le dépassent. Symétriquement côté experts.

\[ \alpha_i^* = \operatorname{quantile}_{1-k/n}(\mathbf{s}_i - \mathbf{\beta}), \qquad \beta_j^* = \operatorname{quantile}_{1-k/n}(\mathbf{s}_{:,j} - \mathbf{\alpha}) \]

D'où vient le nom

Les deux mises à jour sont le même quantile, l'un le long de l'axe des jetons, l'autre le long de l'axe des experts. D'où Quantile Balancing.

Le lien avec la méthode antérieure

Le (sous-)gradient du dual par rapport à \(\beta_j\) vaut :

\[ \frac{\partial \mathcal{L}}{\partial \beta_j} = \frac{mk}{n} - \sum_{i=1}^{m}\chi\big(s_{i,j}-\alpha_i-\beta_j>0\big) \]

soit exactement charge cible moins charge observée.

Le résultat le plus éclairant du rapport

Un pas de SignSGD sur cet objectif redonne exactement la règle de mise à jour par signe de DeepSeek-V3.

La méthode antérieure était donc déjà une descente sur le bon objectif — mais une descente qui ne retenait que la direction de l'erreur. QB saute directement au minimiseur exact de la même fonction.

Cela explique deux choses d'un coup :

  • pourquoi QB n'a aucun hyperparamètre de type taux d'apprentissage ;
  • pourquoi il équilibre en quelques pas même pour ~\(10^3\) experts.

Le rapport situe aussi QB par rapport à BIP, qui résout la même assignation avec des contraintes d'inégalité (\(\le\) au lieu de \(=\)). Les contraintes de non-négativité induites ajoutent un écrêtage \(\max(0,\cdot)\) aux deux mises à jour, ce qui ne peut que réprimer les experts sur-sélectionnés sans promouvoir les sous-sélectionnés — d'où une équilibration nettement plus lente dans leurs expériences.

Le problème pratique : le quantile global

À l'échelle réelle, le quantile porte sur tout le lot global : des millions de marges, réparties sur des centaines de rangs et plusieurs pas d'accumulation de gradient. Les rassembler pour un quantile exact est impossible dans la boucle d'entraînement.

La solution : l'histogramme

L'observation clé

La mise à jour n'a jamais besoin des marges elles-mêmes, seulement de leur distribution par expert — qu'un histogramme résume à coût fixe.

En pratique, Kimi K3 histogramme le biais requis \(r_{i,j} := \alpha_i - s_{i,j}\) (l'opposé de la marge). La cible \(\widehat{b}_j\) devient alors le quantile d'ordre \(k/n\) de \(r_{:,j}\).

Bornes de binning. Les scores sont des sorties de sigmoïde, donc \(s_{i,j}\in(0,1)\). Le seuil \(\alpha_i\) est lui-même un score biaisé \(s_{i,j'}+b_{j'}\), donc dans \((b_{\min},\,1+b_{\max})\). Par conséquent :

\[ r_{i,j} \in [\,b_{\min}-1,\ b_{\max}+1\,] \]

L'intervalle est partitionné en \(B\) casiers uniformes, recalculé à chaque pas pour que la largeur \(w = (b_{\max}-b_{\min}+2)/B\) suive l'étalement du biais.

Accumulation. Chaque rang ajoute ses \(r_{i,j}\) dans une matrice de comptes \(\mathbf{H}\in\mathbb{N}^{n\times B}\), sans aucune communication, sur tous les micro-lots. En fin de pas, un seul all-reduce somme les comptes locaux.

Récupération. On sélectionne le premier casier dont le cumul atteint \(\lceil q\rceil\), avec interpolation linéaire à l'intérieur :

\[ \widehat{b}_j = b_{\min}-1+\Bigl(\beta_j+\operatorname{clip}\bigl(\tfrac{q-c_j}{h_j},\,0,\,1\bigr)\Bigr)w \]

Les trois propriétés qui rendent l'estimateur utilisable

Propriété Justification
Exact Les cumuls sont exacts aux bords de casiers : l'erreur est bornée par la largeur \(w\). Avec \(B = 1000\), quelques \(10^{-3}\) ; aucun déséquilibre résiduel mesurable
Bon marché Un seul all-reduce entier de \(nB\) valeurs par couche et par pas, indépendant de \(m\). Moins de 1 % du coût de l'alternative naturelle
Correct Les comptes étant additifs, l'histogramme global est invariant à la façon dont les jetons sont partitionnés. On obtient le quantile du lot global — et non une moyenne de quantiles par rang, qui en diffère en général

Raffinement mentionné

Maintenir une moyenne mobile exponentielle des quantiles estimés entre les pas réduit le bruit d'échantillonnage et peut améliorer encore l'équilibre.

Vérification de compréhension

Pourquoi la moyenne de quantiles par rang diffère-t-elle du quantile global ?

Parce que le quantile n'est pas une fonction linéaire. Exemple : deux rangs de deux valeurs, \(\{1, 100\}\) et \(\{2, 3\}\). Médiane du rang 1 : 50,5. Médiane du rang 2 : 2,5. Moyenne : 26,5. Or la médiane globale de \(\{1,2,3,100\}\) vaut 2,5. L'écart est majeur — d'où l'importance de l'invariance par partitionnement.

Combien de communication QB ajoute-t-il par pas d'entraînement ?

Un all-reduce d'entiers de \(896 \times 1000 = 896\,000\) valeurs par couche MoE, soit ~82,4 M de valeurs sur 92 couches — environ 330 Mio en int32 par pas. C'est ce que le rapport chiffre à moins de 1 % du coût de l'échange des marges brutes.

Pourquoi figer le biais à l'inférence ?

Parce qu'à l'inférence, il n'y a pas de « lot » sur lequel calculer un quantile : les requêtes arrivent une par une. Figer le biais garde la cohérence entraînement/inférence, et rend le routage déterministe et reproductible. C'est exactement ce que permet la structure du dual : seuls les seuils côté experts \(\mathbf{\beta}\) sont nécessaires au routage ; les seuils côté jetons \(\mathbf{\alpha}\) sont des variables intermédiaires, liées au lot, qu'on jette.


Chapitre précédent : SiTU-GLU · Chapitre suivant : Vision native : MoonViT-V2