feat: add moe auxloss and metrics

This commit is contained in:
2026-08-05 18:12:28 +08:00
parent 602b5ce216
commit 9b7e6c205f
9 changed files with 583 additions and 3 deletions
+1
View File
@@ -97,6 +97,7 @@ class AutoRegressiveLMConfig(BaseModelConfig):
norm_topk_prob: bool = True
decoder_sparse_step: int = 1
mlp_only_layers: Optional[list[int]] = None
moe_aux_loss_coef: float = 0.01
@field_validator("attn_type")
def _validate_attn_type(cls, v: str) -> str:
+12 -1
View File
@@ -1,4 +1,4 @@
from typing import Optional, TypedDict
from typing import List, Optional, TypedDict
import torch
import torch.nn as nn
@@ -70,6 +70,7 @@ class DeepSeekMoE(nn.Module):
)
self.router = Linear(dim, n_routed_experts, bias=False)
self._router_probs: Optional[Tensor] = None
moe_scale = 1 / max(n_shared_experts, 1) + 1 / n_activated_experts
down_init_std = 0.02 / (2 * n_layers * moe_scale) ** 0.5
@@ -111,6 +112,7 @@ class DeepSeekMoE(nn.Module):
router_logits = self.router(x)
router_probs = torch.softmax(router_logits.float(), dim=-1).to(x.dtype)
self._router_probs = router_probs.detach()
topk_weights, topk_indices = torch.topk(router_probs, K, dim=-1)
if self.norm_topk_prob:
@@ -139,3 +141,12 @@ class DeepSeekMoE(nn.Module):
output.index_add_(0, token_idx, expert_output * weights)
return {"hidden_states": output, "aux_loss": aux_loss}
@staticmethod
def collect_router_probs(module: nn.Module) -> List[Tensor]:
"""Recursively collect router_probs from all DeepSeekMoE submodules."""
probs: List[Tensor] = []
for m in module.modules():
if isinstance(m, DeepSeekMoE) and m._router_probs is not None:
probs.append(m._router_probs)
return probs
+6 -1
View File
@@ -1,4 +1,4 @@
from typing import Any, Dict, Mapping, Optional
from typing import Any, Dict, List, Mapping, Optional
import torch
import torch.nn as nn
@@ -10,6 +10,7 @@ from astrai.model.automodel import AutoModel, ModelFactory
from astrai.model.components.decoder_block import DecoderBlock
from astrai.model.components.embedding import Embedding
from astrai.model.components.linear import Linear
from astrai.model.components.mlp import DeepSeekMoE
from astrai.model.components.norm import RMSNorm
from astrai.model.components.rope import RotaryEmbedding
@@ -133,3 +134,7 @@ class AutoRegressiveLM(AutoModel):
if aux_losses:
output["aux_loss"] = torch.stack(aux_losses).mean()
return output
def get_moe_router_probs(self) -> List[Tensor]:
"""Return router_probs from all MoE layers for strategy-side aux loss."""
return DeepSeekMoE.collect_router_probs(self)
+20
View File
@@ -88,3 +88,23 @@ def ctx_get_grad_snr(ctx):
if tracker is None:
return None
return tracker.snr
def ctx_get_moe_aux_loss(ctx):
return ctx.strategy._moe_metrics.get("aux_loss")
def ctx_get_router_entropy(ctx):
return ctx.strategy._moe_metrics.get("router_entropy")
def ctx_get_dead_expert_fraction(ctx):
return ctx.strategy._moe_metrics.get("dead_expert_fraction")
def ctx_get_load_imbalance_mean(ctx):
return ctx.strategy._moe_metrics.get("load_imbalance_mean")
def ctx_get_load_imbalance_max(ctx):
return ctx.strategy._moe_metrics.get("load_imbalance_max")
+113 -1
View File
@@ -1,7 +1,7 @@
"""Training strategy implementations with factory pattern."""
from abc import ABC, abstractmethod
from typing import Callable, Dict, Optional, TypedDict, Union
from typing import Callable, Dict, List, Optional, TypedDict, Union
import torch
import torch.nn as nn
@@ -94,6 +94,97 @@ def make_doc_boundary_mask(position_ids: Tensor) -> Tensor:
return (same_doc & causal).unsqueeze(1)
def _load_balancing_loss(router_probs: Tensor) -> Tensor:
"""Compute MoE load balancing auxiliary loss from router probabilities.
Implements the Switch Transformer load balancing loss (eq. 4-6).
Encourages tokens to be uniformly distributed across experts.
Args:
router_probs: (N, num_experts) tensor of softmax router probabilities.
Returns:
Scalar aux loss = num_experts * sum(f_i * P_i).
"""
num_experts = router_probs.size(-1)
# f_i: fraction of tokens dispatched to expert i (soft mean)
f_i = router_probs.mean(dim=0)
# P_i: average routing probability for expert i
P_i = router_probs.mean(dim=0)
return num_experts * torch.sum(f_i * P_i)
def _collect_moe_diagnostics(
router_probs_list: List[Tensor],
top_k: int,
) -> Dict[str, float]:
"""Collect MoE routing diagnostic metrics from router probabilities.
Args:
router_probs_list: List of (N, num_experts) router probability tensors,
one per MoE layer.
top_k: Number of top experts selected per token.
Returns:
Dict with keys: router_entropy, dead_expert_fraction,
load_imbalance_mean, load_imbalance_max. Values are averaged
across layers.
"""
layer_entropies: List[Tensor] = []
layer_dead_fractions: List[Tensor] = []
layer_imbalance_means: List[Tensor] = []
layer_imbalance_maxs: List[Tensor] = []
for probs in router_probs_list:
probs = probs.detach().to(dtype=torch.float32)
if probs.ndim == 0 or probs.shape[-1] == 0:
continue
probs = probs.reshape(-1, probs.shape[-1])
if probs.numel() == 0:
continue
num_experts = probs.shape[-1]
num_tokens = probs.shape[0]
# Router entropy
entropy = -(probs * torch.log(probs.clamp_min(1e-8))).sum(dim=-1).mean()
# Top-k expert selection
selected_experts = torch.topk(probs, top_k, dim=-1).indices # [tokens, top_k]
expert_mask = F.one_hot(selected_experts, num_experts) # [tokens, top_k, E]
expert_counts = expert_mask.sum(dim=(0, 1)).to(dtype=torch.float32) # [E]
# Ideal load: tokens * top_k / num_experts
ideal_load = (num_tokens * top_k) / max(num_experts, 1)
# Load imbalance ratios
load_ratios = expert_counts / max(ideal_load, 1.0)
imbalance_mean = (load_ratios - 1.0).abs().mean()
imbalance_max = load_ratios.max()
dead_fraction = (expert_counts == 0).to(dtype=torch.float32).mean()
layer_entropies.append(entropy)
layer_dead_fractions.append(dead_fraction)
layer_imbalance_means.append(imbalance_mean)
layer_imbalance_maxs.append(imbalance_max)
if not layer_entropies:
return {}
return {
"router_entropy": float(torch.stack(layer_entropies).mean().cpu().item()),
"dead_expert_fraction": float(
torch.stack(layer_dead_fractions).mean().cpu().item()
),
"load_imbalance_mean": float(
torch.stack(layer_imbalance_means).mean().cpu().item()
),
"load_imbalance_max": float(
torch.stack(layer_imbalance_maxs).mean().cpu().item()
),
}
class BaseStrategy(ABC):
"""Abstract base class for training strategies.
@@ -115,6 +206,7 @@ class BaseStrategy(ABC):
self.device = device
self.executor = kwargs.pop("executor", None)
self.moe_aux_loss_coef = kwargs.pop("moe_aux_loss_coef", 0.01)
self._moe_metrics: Dict[str, float] = {}
self.extra_kwargs = kwargs
self._rollout_runner = None
@@ -145,6 +237,7 @@ class BaseStrategy(ABC):
total_loss = total_loss + weighted_aux_loss
metrics["moe_aux_loss"] = aux_loss
metrics["moe_aux_loss_weighted"] = weighted_aux_loss
self._refresh_moe_diagnostics(aux_loss)
metrics["loss"] = total_loss
return {
"loss": total_loss,
@@ -188,6 +281,23 @@ class BaseStrategy(ABC):
"""
pass
def _refresh_moe_diagnostics(self, aux_loss: Tensor) -> None:
"""Collect MoE routing diagnostics from model router probs.
Populates ``self._moe_metrics`` with router entropy, dead expert
fraction, load imbalance, and aux_loss. Called from
:meth:`_loss_output` when an MoE aux loss is present.
"""
router_probs_list: List[Tensor] = self.model.get_moe_router_probs()
if not router_probs_list:
self._moe_metrics = {}
return
self._moe_metrics = _collect_moe_diagnostics(
router_probs_list,
self.model.config.n_activated_experts,
)
self._moe_metrics["aux_loss"] = float(aux_loss.detach().cpu().item())
def on_optimizer_step(self):
"""Advance online rollout state after a successful optimizer step."""
if self._rollout_runner is not None:
@@ -230,6 +340,7 @@ class SEQStrategy(BaseStrategy):
"""Standard next-token prediction training strategy.
Computes cross-entropy loss for next token prediction.
Optionally adds MoE load balancing auxiliary loss.
"""
def __init__(
@@ -265,6 +376,7 @@ class SFTStrategy(BaseStrategy):
"""Supervised Fine-tuning strategy with loss masking.
Applies cross-entropy loss only to tokens where loss_mask is True.
Optionally adds MoE load balancing auxiliary loss.
"""
def __init__(
+10
View File
@@ -17,10 +17,15 @@ from astrai.parallel import only_on_rank
from astrai.parallel.setup import get_current_device
from astrai.serialization import Checkpoint
from astrai.trainer.metric_util import (
ctx_get_dead_expert_fraction,
ctx_get_grad_norm,
ctx_get_grad_snr,
ctx_get_load_imbalance_max,
ctx_get_load_imbalance_mean,
ctx_get_loss,
ctx_get_lr,
ctx_get_moe_aux_loss,
ctx_get_router_entropy,
ctx_get_val_loss,
)
from astrai.trainer.train_context import TrainContext
@@ -257,6 +262,11 @@ class MetricCallback(TrainCallback):
"val_loss": ctx_get_val_loss,
"grad_norm": ctx_get_grad_norm,
"grad_snr": ctx_get_grad_snr,
"moe_aux_loss": ctx_get_moe_aux_loss,
"router_entropy": ctx_get_router_entropy,
"dead_expert_fraction": ctx_get_dead_expert_fraction,
"load_imbalance_mean": ctx_get_load_imbalance_mean,
"load_imbalance_max": ctx_get_load_imbalance_max,
}
def _metrics(self, context: TrainContext, names):