Load Balancing Loss (MoE)
iaDéfinition
La Load Balancing Loss (perte d'equilibrage de charge) est un terme de regularisation specifique aux architectures Mixture of Experts (MoE) qui penalise les desequilibres de distribution des tokens entre les experts, forcant le routeur a distribuer les tokens de maniere plus uniforme entre tous les experts disponibles. Le probleme qu'elle resout : sans regularisation, les MoE souffrent d'un probleme d''expert collapse' — le routeur apprend rapidement a toujours envoyer les tokens vers un petit sous-ensemble d'experts 'favoris', tandis que les autres experts restent inactifs et n'apprennent pas. Cela annule les avantages de la specialisation des experts et reduit la capacite effective du modele. La formulation de Switch Transformer (Fedus et al., 2022) : L_aux = alpha * N * sum_i(f_i * P_i), ou N = nombre d'experts, f_i = fraction de tokens routes vers l'expert i, P_i = probabilite moyenne du routeur pour l'expert i. Cette perte est maximisee quand tous les f_i = 1/N (distribution uniforme), ce qui encourage le routeur vers cet equilibre. Les valeurs de alpha : un alpha trop faible n'equilibre pas sufisamment (expert collapse), un alpha trop eleve force un equilibrage parfait qui empeche la specialisation des experts (les experts ne peuvent plus se specialiser si forcement tous recoivent le meme nombre de tokens). En pratique, alpha dans [0.01, 0.001] est utilise. Mixtral utilise alpha=0.01, DeepSeek V3 utilise une formulation plus sophistiquee avec un 'complementary load balancing loss'. L'innovation DeepSeek V3 : introduit un 'auxiliary-loss-free load balancing' via une mise a jour de biais des scores de routage, evitant l'interference entre le gradient de la loss principale et celui du load balancing. Cela ameliore la qualite du modele tout en maintenant un equilibrage adequat.
Implementation Load Balancing Loss
import torch
def load_balancing_loss(router_probs, expert_indices, n_experts, alpha=0.01):
# router_probs: (batch, seq_len, n_experts)
# expert_indices: (batch, seq_len, top_k)
batch_size, seq_len, _ = router_probs.shape
total_tokens = batch_size * seq_len
# Fraction de tokens par expert
one_hot = torch.zeros(total_tokens, n_experts)
one_hot.scatter_(1, expert_indices.reshape(-1, 1), 1)
f_i = one_hot.mean(dim=0) # fraction tokens vers chaque expert
# Probabilite moyenne par expert
P_i = router_probs.reshape(total_tokens, n_experts).mean(dim=0)
# Load balancing loss
loss = alpha * n_experts * (f_i * P_i).sum()
return loss
Articles liés
Expert en cybersécurité offensive et intelligence artificielle. Pentest, audit et développement IA sur-mesure.
Services
- Audit Infrastructure
- Audit Kubernetes
- Audit Microsoft 365
- Audit Sécurité Réseau
- Analyse de Risques
- Audit Active Directory
- Audit Application Web
- Audit Cloud (AWS/Azure/GCP)
- Audit Messagerie
- Audit API (OWASP Top 10)
- Audit DevSecOps & CI/CD
- Audit Code Source (SAST)
- Audit Postes de Travail
- Audit Sauvegarde & Résilience
- Audit OT/SCADA (IEC 62443)
- Développement IA
- Formations
Ressources
Projets & Outils
© 2026 Ayi NEDJIMI Consultants. Tous droits réservés.
Un projet cybersécurité ?
Expert dispo · Réponse 24h