refactor: stateless MoE routing with grouped dispatch
- replace per-expert mask scan with sort+bincount grouped dispatch - carry router stats in forward output instead of module state - keep MoE diagnostics working under DDP/FSDP wrappers - remove unused _load_balancing_loss helper
This commit is contained in:
@@ -6,13 +6,14 @@ from torch import Tensor
|
|||||||
|
|
||||||
from astrai.inference.core.cache import KVCache
|
from astrai.inference.core.cache import KVCache
|
||||||
from astrai.model.components.attention import AttnFactory
|
from astrai.model.components.attention import AttnFactory
|
||||||
from astrai.model.components.mlp import FFNFactory
|
from astrai.model.components.mlp import FFNFactory, RouterStats
|
||||||
from astrai.model.components.norm import RMSNorm
|
from astrai.model.components.norm import RMSNorm
|
||||||
|
|
||||||
|
|
||||||
class DecoderOutput(TypedDict):
|
class DecoderOutput(TypedDict):
|
||||||
hidden_states: Tensor
|
hidden_states: Tensor
|
||||||
aux_loss: Optional[Tensor]
|
aux_loss: Optional[Tensor]
|
||||||
|
router_stats: Optional[RouterStats]
|
||||||
|
|
||||||
|
|
||||||
class DecoderBlock(nn.Module):
|
class DecoderBlock(nn.Module):
|
||||||
@@ -66,4 +67,8 @@ class DecoderBlock(nn.Module):
|
|||||||
mlp_output = self.mlp(normalized)
|
mlp_output = self.mlp(normalized)
|
||||||
x = mlp_output["hidden_states"] + x
|
x = mlp_output["hidden_states"] + x
|
||||||
|
|
||||||
return {"hidden_states": x, "aux_loss": mlp_output["aux_loss"]}
|
return {
|
||||||
|
"hidden_states": x,
|
||||||
|
"aux_loss": mlp_output["aux_loss"],
|
||||||
|
"router_stats": mlp_output.get("router_stats"),
|
||||||
|
}
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
from typing import List, Optional, TypedDict
|
from typing import Optional, TypedDict
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
@@ -13,14 +13,26 @@ class FFNFactory(BaseFactory[nn.Module]):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class RouterStats(TypedDict):
|
||||||
|
"""Per-layer MoE routing statistics for training diagnostics.
|
||||||
|
|
||||||
|
Both tensors are detached monitoring data produced during forward.
|
||||||
|
"""
|
||||||
|
|
||||||
|
probs: Tensor
|
||||||
|
topk_indices: Tensor
|
||||||
|
|
||||||
|
|
||||||
class FFNOutput(TypedDict):
|
class FFNOutput(TypedDict):
|
||||||
hidden_states: Tensor
|
hidden_states: Tensor
|
||||||
aux_loss: Optional[Tensor]
|
aux_loss: Optional[Tensor]
|
||||||
|
router_stats: Optional[RouterStats]
|
||||||
|
|
||||||
|
|
||||||
class RoutedOutput(TypedDict):
|
class RoutedOutput(TypedDict):
|
||||||
hidden_states: Tensor
|
hidden_states: Tensor
|
||||||
aux_loss: Optional[Tensor]
|
aux_loss: Optional[Tensor]
|
||||||
|
router_stats: Optional[RouterStats]
|
||||||
|
|
||||||
|
|
||||||
@FFNFactory.register("mlp")
|
@FFNFactory.register("mlp")
|
||||||
@@ -34,7 +46,7 @@ class MLP(nn.Module):
|
|||||||
def forward(self, x: Tensor) -> FFNOutput:
|
def forward(self, x: Tensor) -> FFNOutput:
|
||||||
gated = self.up(x) * F.silu(self.gate(x))
|
gated = self.up(x) * F.silu(self.gate(x))
|
||||||
out = self.down(gated)
|
out = self.down(gated)
|
||||||
return {"hidden_states": out, "aux_loss": None}
|
return {"hidden_states": out, "aux_loss": None, "router_stats": None}
|
||||||
|
|
||||||
|
|
||||||
@FFNFactory.register("moe")
|
@FFNFactory.register("moe")
|
||||||
@@ -70,7 +82,6 @@ class DeepSeekMoE(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.router = Linear(dim, n_routed_experts, bias=False)
|
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
|
moe_scale = 1 / max(n_shared_experts, 1) + 1 / n_activated_experts
|
||||||
down_init_std = 0.02 / (2 * n_layers * moe_scale) ** 0.5
|
down_init_std = 0.02 / (2 * n_layers * moe_scale) ** 0.5
|
||||||
|
|
||||||
@@ -96,7 +107,11 @@ class DeepSeekMoE(nn.Module):
|
|||||||
routed_output = self._routed_forward(x_flat, include_aux_loss)
|
routed_output = self._routed_forward(x_flat, include_aux_loss)
|
||||||
|
|
||||||
out = (shared_out + routed_output["hidden_states"]).view(bsz, seq_len, dim)
|
out = (shared_out + routed_output["hidden_states"]).view(bsz, seq_len, dim)
|
||||||
return {"hidden_states": out, "aux_loss": routed_output["aux_loss"]}
|
return {
|
||||||
|
"hidden_states": out,
|
||||||
|
"aux_loss": routed_output["aux_loss"],
|
||||||
|
"router_stats": routed_output["router_stats"],
|
||||||
|
}
|
||||||
|
|
||||||
def _shared_forward(self, x: Tensor) -> Tensor:
|
def _shared_forward(self, x: Tensor) -> Tensor:
|
||||||
if self.n_shared_experts == 0:
|
if self.n_shared_experts == 0:
|
||||||
@@ -109,44 +124,54 @@ class DeepSeekMoE(nn.Module):
|
|||||||
def _routed_forward(self, x: Tensor, include_aux_loss: bool) -> RoutedOutput:
|
def _routed_forward(self, x: Tensor, include_aux_loss: bool) -> RoutedOutput:
|
||||||
N, D = x.shape
|
N, D = x.shape
|
||||||
K = self.n_activated_experts
|
K = self.n_activated_experts
|
||||||
|
E = self.n_routed_experts
|
||||||
|
|
||||||
router_logits = self.router(x)
|
router_logits = self.router(x)
|
||||||
router_probs = torch.softmax(router_logits.float(), dim=-1).to(x.dtype)
|
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)
|
topk_weights, topk_indices = torch.topk(router_probs, K, dim=-1, sorted=False)
|
||||||
if self.norm_topk_prob:
|
if self.norm_topk_prob:
|
||||||
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
|
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
|
||||||
|
|
||||||
aux_loss = None
|
aux_loss = None
|
||||||
|
router_stats = None
|
||||||
if include_aux_loss:
|
if include_aux_loss:
|
||||||
expert_load = F.one_hot(
|
expert_load = F.one_hot(topk_indices, num_classes=E).float()
|
||||||
topk_indices, num_classes=self.n_routed_experts
|
|
||||||
).float()
|
|
||||||
expert_load = expert_load.mean(dim=(0, 1))
|
expert_load = expert_load.mean(dim=(0, 1))
|
||||||
router_prob = router_probs.float().mean(dim=0)
|
router_prob = router_probs.float().mean(dim=0)
|
||||||
aux_loss = self.n_routed_experts * (expert_load * router_prob).sum()
|
aux_loss = E * (expert_load * router_prob).sum()
|
||||||
|
router_stats = {
|
||||||
|
"probs": router_probs.detach(),
|
||||||
|
"topk_indices": topk_indices,
|
||||||
|
}
|
||||||
|
|
||||||
|
# Grouped dispatch: sort (token, slot) pairs by expert so each expert
|
||||||
|
# consumes one contiguous slice instead of a per-expert mask scan.
|
||||||
|
flat_experts = topk_indices.reshape(-1)
|
||||||
|
sorted_experts, order = torch.sort(flat_experts)
|
||||||
|
flat_tokens = x.repeat_interleave(K, dim=0)[order]
|
||||||
|
flat_weights = topk_weights.reshape(-1, 1)[order]
|
||||||
|
boundaries = torch.cumsum(
|
||||||
|
torch.bincount(sorted_experts, minlength=E), dim=0
|
||||||
|
).tolist()
|
||||||
|
|
||||||
output = torch.zeros(N, D, device=x.device, dtype=x.dtype)
|
output = torch.zeros(N, D, device=x.device, dtype=x.dtype)
|
||||||
for expert_idx in range(self.n_routed_experts):
|
start = 0
|
||||||
expert_mask = topk_indices == expert_idx
|
for expert_idx, end in enumerate(boundaries):
|
||||||
token_idx, k_idx = expert_mask.nonzero(as_tuple=True)
|
if end == start:
|
||||||
if token_idx.numel() == 0:
|
|
||||||
continue
|
continue
|
||||||
expert = self.routed_experts[expert_idx]
|
expert_output = self.routed_experts[expert_idx](flat_tokens[start:end])[
|
||||||
expert_input = x[token_idx]
|
"hidden_states"
|
||||||
expert_output = expert(expert_input)["hidden_states"]
|
]
|
||||||
|
output.index_add_(
|
||||||
|
0,
|
||||||
|
order[start:end] // K,
|
||||||
|
expert_output * flat_weights[start:end],
|
||||||
|
)
|
||||||
|
start = end
|
||||||
|
|
||||||
weights = topk_weights[token_idx, k_idx].unsqueeze(-1)
|
return {
|
||||||
output.index_add_(0, token_idx, expert_output * weights)
|
"hidden_states": output,
|
||||||
|
"aux_loss": aux_loss,
|
||||||
return {"hidden_states": output, "aux_loss": aux_loss}
|
"router_stats": router_stats,
|
||||||
|
}
|
||||||
@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
|
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
from typing import Any, Dict, List, Mapping, Optional
|
from typing import Any, Dict, Mapping, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
@@ -10,7 +10,6 @@ from astrai.model.automodel import AutoModel, ModelFactory
|
|||||||
from astrai.model.components.decoder_block import DecoderBlock
|
from astrai.model.components.decoder_block import DecoderBlock
|
||||||
from astrai.model.components.embedding import Embedding
|
from astrai.model.components.embedding import Embedding
|
||||||
from astrai.model.components.linear import Linear
|
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.norm import RMSNorm
|
||||||
from astrai.model.components.rope import RotaryEmbedding
|
from astrai.model.components.rope import RotaryEmbedding
|
||||||
|
|
||||||
@@ -115,6 +114,7 @@ class AutoRegressiveLM(AutoModel):
|
|||||||
use_sdpa_causal_mask = attn_mask is None
|
use_sdpa_causal_mask = attn_mask is None
|
||||||
|
|
||||||
aux_losses = []
|
aux_losses = []
|
||||||
|
router_stats_list = []
|
||||||
for layer in self.layers:
|
for layer in self.layers:
|
||||||
layer_output = layer(
|
layer_output = layer(
|
||||||
x,
|
x,
|
||||||
@@ -124,8 +124,10 @@ class AutoRegressiveLM(AutoModel):
|
|||||||
use_sdpa_causal_mask,
|
use_sdpa_causal_mask,
|
||||||
)
|
)
|
||||||
x = layer_output["hidden_states"]
|
x = layer_output["hidden_states"]
|
||||||
if layer_output["aux_loss"] is not None:
|
stats = layer_output.get("router_stats")
|
||||||
|
if stats is not None:
|
||||||
aux_losses.append(layer_output["aux_loss"])
|
aux_losses.append(layer_output["aux_loss"])
|
||||||
|
router_stats_list.append(stats)
|
||||||
|
|
||||||
hidden_states = self.norm(x)
|
hidden_states = self.norm(x)
|
||||||
logits = self.lm_head(hidden_states)
|
logits = self.lm_head(hidden_states)
|
||||||
@@ -133,8 +135,5 @@ class AutoRegressiveLM(AutoModel):
|
|||||||
output = {"logits": logits, "hidden_states": hidden_states}
|
output = {"logits": logits, "hidden_states": hidden_states}
|
||||||
if aux_losses:
|
if aux_losses:
|
||||||
output["aux_loss"] = torch.stack(aux_losses).mean()
|
output["aux_loss"] = torch.stack(aux_losses).mean()
|
||||||
|
output["router_stats"] = router_stats_list
|
||||||
return output
|
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)
|
|
||||||
|
|||||||
+50
-59
@@ -9,6 +9,7 @@ import torch.nn.functional as F
|
|||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
|
from astrai.model.components.mlp import RouterStats
|
||||||
from astrai.parallel.executor import broadcast_state_dict
|
from astrai.parallel.executor import broadcast_state_dict
|
||||||
from astrai.trainer.rollout import RolloutResult
|
from astrai.trainer.rollout import RolloutResult
|
||||||
|
|
||||||
@@ -21,6 +22,7 @@ class LossOutput(TypedDict):
|
|||||||
class LogprobsOutput(TypedDict):
|
class LogprobsOutput(TypedDict):
|
||||||
logprobs: Tensor
|
logprobs: Tensor
|
||||||
aux_loss: Optional[Tensor]
|
aux_loss: Optional[Tensor]
|
||||||
|
router_stats: Optional[List[RouterStats]]
|
||||||
|
|
||||||
|
|
||||||
def move_to_device(batch: Dict[str, Tensor], device: str) -> Dict[str, Tensor]:
|
def move_to_device(batch: Dict[str, Tensor], device: str) -> Dict[str, Tensor]:
|
||||||
@@ -75,7 +77,11 @@ def get_logprobs(
|
|||||||
logprobs = (token_logprobs * shifted_loss_mask).sum(dim=-1)
|
logprobs = (token_logprobs * shifted_loss_mask).sum(dim=-1)
|
||||||
else:
|
else:
|
||||||
logprobs = token_logprobs * shifted_loss_mask
|
logprobs = token_logprobs * shifted_loss_mask
|
||||||
return {"logprobs": logprobs, "aux_loss": outputs.get("aux_loss")}
|
return {
|
||||||
|
"logprobs": logprobs,
|
||||||
|
"aux_loss": outputs.get("aux_loss"),
|
||||||
|
"router_stats": outputs.get("router_stats"),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
def make_doc_boundary_mask(position_ids: Tensor) -> Tensor:
|
def make_doc_boundary_mask(position_ids: Tensor) -> Tensor:
|
||||||
@@ -94,36 +100,14 @@ def make_doc_boundary_mask(position_ids: Tensor) -> Tensor:
|
|||||||
return (same_doc & causal).unsqueeze(1)
|
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(
|
def _collect_moe_diagnostics(
|
||||||
router_probs_list: List[Tensor],
|
router_stats_list: List[RouterStats],
|
||||||
top_k: int,
|
|
||||||
) -> Dict[str, float]:
|
) -> Dict[str, float]:
|
||||||
"""Collect MoE routing diagnostic metrics from router probabilities.
|
"""Collect MoE routing diagnostic metrics from per-layer router stats.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
router_probs_list: List of (N, num_experts) router probability tensors,
|
router_stats_list: One :class:`RouterStats` dict per MoE layer with
|
||||||
one per MoE layer.
|
keys ``probs`` (N, E) and ``topk_indices`` (N, K), both detached.
|
||||||
top_k: Number of top experts selected per token.
|
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Dict with keys: router_entropy, dead_expert_fraction,
|
Dict with keys: router_entropy, dead_expert_fraction,
|
||||||
@@ -135,33 +119,26 @@ def _collect_moe_diagnostics(
|
|||||||
layer_imbalance_means: List[Tensor] = []
|
layer_imbalance_means: List[Tensor] = []
|
||||||
layer_imbalance_maxs: List[Tensor] = []
|
layer_imbalance_maxs: List[Tensor] = []
|
||||||
|
|
||||||
for probs in router_probs_list:
|
for stats in router_stats_list:
|
||||||
probs = probs.detach().to(dtype=torch.float32)
|
probs = stats["probs"].float()
|
||||||
if probs.ndim == 0 or probs.shape[-1] == 0:
|
topk_indices = stats["topk_indices"]
|
||||||
|
num_experts = probs.shape[-1]
|
||||||
|
if num_experts == 0:
|
||||||
continue
|
continue
|
||||||
probs = probs.reshape(-1, probs.shape[-1])
|
probs = probs.reshape(-1, num_experts)
|
||||||
if probs.numel() == 0:
|
if probs.numel() == 0:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
num_experts = probs.shape[-1]
|
|
||||||
num_tokens = probs.shape[0]
|
|
||||||
|
|
||||||
# Router entropy
|
# Router entropy
|
||||||
entropy = -(probs * torch.log(probs.clamp_min(1e-8))).sum(dim=-1).mean()
|
entropy = -(probs * torch.log(probs.clamp_min(1e-8))).sum(dim=-1).mean()
|
||||||
|
|
||||||
# Top-k expert selection
|
# Load from the actual dispatch: one-hot sum of top-k assignments.
|
||||||
selected_experts = torch.topk(probs, top_k, dim=-1).indices # [tokens, top_k]
|
expert_counts = F.one_hot(topk_indices, num_experts).sum(dim=(0, 1)).float()
|
||||||
expert_mask = F.one_hot(selected_experts, num_experts) # [tokens, top_k, E]
|
ideal_load = expert_counts.mean() # N*K / E
|
||||||
expert_counts = expert_mask.sum(dim=(0, 1)).to(dtype=torch.float32) # [E]
|
load_ratios = expert_counts / max(float(ideal_load), 1.0)
|
||||||
|
|
||||||
# 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_mean = (load_ratios - 1.0).abs().mean()
|
||||||
imbalance_max = load_ratios.max()
|
imbalance_max = load_ratios.max()
|
||||||
dead_fraction = (expert_counts == 0).to(dtype=torch.float32).mean()
|
dead_fraction = (expert_counts == 0).float().mean()
|
||||||
|
|
||||||
layer_entropies.append(entropy)
|
layer_entropies.append(entropy)
|
||||||
layer_dead_fractions.append(dead_fraction)
|
layer_dead_fractions.append(dead_fraction)
|
||||||
@@ -230,6 +207,7 @@ class BaseStrategy(ABC):
|
|||||||
task_loss: Tensor,
|
task_loss: Tensor,
|
||||||
metrics: Dict[str, Tensor],
|
metrics: Dict[str, Tensor],
|
||||||
aux_loss: Optional[Tensor] = None,
|
aux_loss: Optional[Tensor] = None,
|
||||||
|
router_stats: Optional[List[RouterStats]] = None,
|
||||||
) -> LossOutput:
|
) -> LossOutput:
|
||||||
total_loss = task_loss
|
total_loss = task_loss
|
||||||
if aux_loss is not None:
|
if aux_loss is not None:
|
||||||
@@ -237,7 +215,7 @@ class BaseStrategy(ABC):
|
|||||||
total_loss = total_loss + weighted_aux_loss
|
total_loss = total_loss + weighted_aux_loss
|
||||||
metrics["moe_aux_loss"] = aux_loss
|
metrics["moe_aux_loss"] = aux_loss
|
||||||
metrics["moe_aux_loss_weighted"] = weighted_aux_loss
|
metrics["moe_aux_loss_weighted"] = weighted_aux_loss
|
||||||
self._refresh_moe_diagnostics(aux_loss)
|
self._refresh_moe_diagnostics(aux_loss, router_stats)
|
||||||
metrics["loss"] = total_loss
|
metrics["loss"] = total_loss
|
||||||
return {
|
return {
|
||||||
"loss": total_loss,
|
"loss": total_loss,
|
||||||
@@ -281,21 +259,18 @@ class BaseStrategy(ABC):
|
|||||||
"""
|
"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def _refresh_moe_diagnostics(self, aux_loss: Tensor) -> None:
|
def _refresh_moe_diagnostics(
|
||||||
"""Collect MoE routing diagnostics from model router probs.
|
self,
|
||||||
|
aux_loss: Tensor,
|
||||||
|
router_stats: Optional[List[RouterStats]] = None,
|
||||||
|
) -> None:
|
||||||
|
"""Collect MoE routing diagnostics from the latest forward pass.
|
||||||
|
|
||||||
Populates ``self._moe_metrics`` with router entropy, dead expert
|
Populates ``self._moe_metrics`` with router entropy, dead expert
|
||||||
fraction, load imbalance, and aux_loss. Called from
|
fraction, load imbalance, and aux_loss. Called from
|
||||||
:meth:`_loss_output` when an MoE aux loss is present.
|
:meth:`_loss_output` when an MoE aux loss is present.
|
||||||
"""
|
"""
|
||||||
router_probs_list: List[Tensor] = self.model.get_moe_router_probs()
|
self._moe_metrics = _collect_moe_diagnostics(router_stats or [])
|
||||||
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())
|
self._moe_metrics["aux_loss"] = float(aux_loss.detach().cpu().item())
|
||||||
|
|
||||||
def on_optimizer_step(self):
|
def on_optimizer_step(self):
|
||||||
@@ -368,7 +343,12 @@ class SEQStrategy(BaseStrategy):
|
|||||||
label_smoothing=self.label_smoothing,
|
label_smoothing=self.label_smoothing,
|
||||||
)
|
)
|
||||||
|
|
||||||
return self._loss_output(loss, {"task_loss": loss}, outputs.get("aux_loss"))
|
return self._loss_output(
|
||||||
|
loss,
|
||||||
|
{"task_loss": loss},
|
||||||
|
outputs.get("aux_loss"),
|
||||||
|
outputs.get("router_stats"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@StrategyFactory.register("sft")
|
@StrategyFactory.register("sft")
|
||||||
@@ -416,7 +396,12 @@ class SFTStrategy(BaseStrategy):
|
|||||||
label_smoothing=self.label_smoothing,
|
label_smoothing=self.label_smoothing,
|
||||||
)
|
)
|
||||||
|
|
||||||
return self._loss_output(loss, {"task_loss": loss}, outputs.get("aux_loss"))
|
return self._loss_output(
|
||||||
|
loss,
|
||||||
|
{"task_loss": loss},
|
||||||
|
outputs.get("aux_loss"),
|
||||||
|
outputs.get("router_stats"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@StrategyFactory.register("dpo")
|
@StrategyFactory.register("dpo")
|
||||||
@@ -491,7 +476,12 @@ class DPOStrategy(BaseStrategy):
|
|||||||
ratio_diff = pi_log_ratio - ref_log_ratio
|
ratio_diff = pi_log_ratio - ref_log_ratio
|
||||||
dpo_loss = -F.logsigmoid(self.beta * ratio_diff).mean()
|
dpo_loss = -F.logsigmoid(self.beta * ratio_diff).mean()
|
||||||
|
|
||||||
return self._loss_output(dpo_loss, {"dpo_loss": dpo_loss}, aux_loss)
|
return self._loss_output(
|
||||||
|
dpo_loss,
|
||||||
|
{"dpo_loss": dpo_loss},
|
||||||
|
aux_loss,
|
||||||
|
policy_output.get("router_stats"),
|
||||||
|
)
|
||||||
|
|
||||||
def supports_online(self) -> bool:
|
def supports_online(self) -> bool:
|
||||||
return True
|
return True
|
||||||
@@ -661,6 +651,7 @@ class GRPOStrategy(BaseStrategy):
|
|||||||
task_loss,
|
task_loss,
|
||||||
{"policy_loss": policy_loss, "kl_loss": kl_penalty},
|
{"policy_loss": policy_loss, "kl_loss": kl_penalty},
|
||||||
aux_loss,
|
aux_loss,
|
||||||
|
policy_output.get("router_stats"),
|
||||||
)
|
)
|
||||||
|
|
||||||
def supports_online(self) -> bool:
|
def supports_online(self) -> bool:
|
||||||
|
|||||||
@@ -265,10 +265,9 @@ def test_moe_defaults_preserve_normalized_routing():
|
|||||||
assert model.layers[0].mlp.norm_topk_prob is True
|
assert model.layers[0].mlp.norm_topk_prob is True
|
||||||
|
|
||||||
|
|
||||||
def test_moe_router_probs_populated_after_forward():
|
def test_moe_router_stats_in_output_during_training():
|
||||||
"""Verify DeepSeekMoE._router_probs is set after forward in training mode."""
|
"""Verify forward output carries per-layer router_stats in training mode."""
|
||||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||||
from astrai.model.components.mlp import DeepSeekMoE
|
|
||||||
|
|
||||||
config = AutoRegressiveLMConfig(
|
config = AutoRegressiveLMConfig(
|
||||||
**TINY_CONFIG,
|
**TINY_CONFIG,
|
||||||
@@ -283,19 +282,18 @@ def test_moe_router_probs_populated_after_forward():
|
|||||||
input_ids = torch.randint(0, config.vocab_size, (2, 8))
|
input_ids = torch.randint(0, config.vocab_size, (2, 8))
|
||||||
|
|
||||||
with torch.enable_grad():
|
with torch.enable_grad():
|
||||||
model(input_ids)
|
outputs = model(input_ids)
|
||||||
|
|
||||||
# All MoE layers should have router_probs set
|
stats = outputs["router_stats"]
|
||||||
moe_layers = [m for m in model.modules() if isinstance(m, DeepSeekMoE)]
|
assert isinstance(stats, list)
|
||||||
assert len(moe_layers) > 0
|
assert len(stats) == config.num_hidden_layers
|
||||||
for layer in moe_layers:
|
for s in stats:
|
||||||
assert layer._router_probs is not None
|
assert s["probs"].shape == (2 * 8, 4) # (N, n_routed_experts)
|
||||||
assert layer._router_probs.ndim == 2
|
assert s["topk_indices"].shape == (2 * 8, 2) # (N, n_activated_experts)
|
||||||
assert layer._router_probs.shape[-1] == 4 # n_routed_experts
|
|
||||||
|
|
||||||
|
|
||||||
def test_get_moe_router_probs_moe_model():
|
def test_moe_router_stats_absent_in_eval():
|
||||||
"""Verify get_moe_router_probs() returns a list of tensors for MoE models."""
|
"""Verify no router_stats are emitted outside training."""
|
||||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||||
|
|
||||||
config = AutoRegressiveLMConfig(
|
config = AutoRegressiveLMConfig(
|
||||||
@@ -306,60 +304,27 @@ def test_get_moe_router_probs_moe_model():
|
|||||||
n_activated_experts=2,
|
n_activated_experts=2,
|
||||||
)
|
)
|
||||||
model = AutoRegressiveLM(config)
|
model = AutoRegressiveLM(config)
|
||||||
model.train()
|
model.eval()
|
||||||
|
|
||||||
with torch.enable_grad():
|
with torch.no_grad():
|
||||||
model(torch.randint(0, config.vocab_size, (2, 8)))
|
outputs = model(torch.randint(0, config.vocab_size, (2, 8)))
|
||||||
|
|
||||||
probs = model.get_moe_router_probs()
|
assert "router_stats" not in outputs
|
||||||
assert isinstance(probs, list)
|
|
||||||
assert len(probs) == 2 # num_hidden_layers
|
|
||||||
for p in probs:
|
|
||||||
assert p.ndim == 2
|
|
||||||
assert p.shape[-1] == 4
|
|
||||||
|
|
||||||
|
|
||||||
def test_get_moe_router_probs_non_moe_model():
|
def test_no_router_stats_for_mlp_model():
|
||||||
"""Verify get_moe_router_probs() returns empty list for non-MoE models."""
|
"""Verify pure MLP models emit no router_stats and no aux_loss."""
|
||||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||||
|
|
||||||
config = AutoRegressiveLMConfig(**TINY_CONFIG, ffn_type="mlp")
|
config = AutoRegressiveLMConfig(**TINY_CONFIG, ffn_type="mlp")
|
||||||
model = AutoRegressiveLM(config)
|
model = AutoRegressiveLM(config)
|
||||||
|
|
||||||
probs_untrained = model.get_moe_router_probs()
|
|
||||||
assert probs_untrained == []
|
|
||||||
|
|
||||||
model.train()
|
model.train()
|
||||||
|
|
||||||
with torch.enable_grad():
|
with torch.enable_grad():
|
||||||
model(torch.randint(0, config.vocab_size, (2, 8)))
|
outputs = model(torch.randint(0, config.vocab_size, (2, 8)))
|
||||||
|
|
||||||
probs = model.get_moe_router_probs()
|
assert "router_stats" not in outputs
|
||||||
assert probs == []
|
assert "aux_loss" not in outputs
|
||||||
|
|
||||||
|
|
||||||
def test_collect_router_probs_static_method():
|
|
||||||
"""Verify DeepSeekMoE.collect_router_probs static method."""
|
|
||||||
from astrai.model.components.mlp import DeepSeekMoE
|
|
||||||
|
|
||||||
moe = DeepSeekMoE(
|
|
||||||
dim=8,
|
|
||||||
dim_ffn=16,
|
|
||||||
n_routed_experts=4,
|
|
||||||
n_shared_experts=1,
|
|
||||||
n_activated_experts=2,
|
|
||||||
)
|
|
||||||
moe.train()
|
|
||||||
with torch.enable_grad():
|
|
||||||
moe(torch.randn(2, 8, 8))
|
|
||||||
|
|
||||||
# collect_router_probs should find the MoE layer
|
|
||||||
probs = DeepSeekMoE.collect_router_probs(moe)
|
|
||||||
assert len(probs) == 1
|
|
||||||
assert probs[0].shape[-1] == 4
|
|
||||||
|
|
||||||
# On a plain MLP module, should return empty
|
|
||||||
mlp_module = MLP(8, 16)
|
|
||||||
assert DeepSeekMoE.collect_router_probs(mlp_module) == []
|
|
||||||
|
|
||||||
|
|
||||||
def test_moe_aux_loss_only_emitted_during_training():
|
def test_moe_aux_loss_only_emitted_during_training():
|
||||||
|
|||||||
@@ -9,19 +9,15 @@ import pytest
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from astrai.config.model_config import AutoRegressiveLMConfig
|
from astrai.config.model_config import AutoRegressiveLMConfig
|
||||||
from astrai.model.components.mlp import DeepSeekMoE
|
|
||||||
from astrai.model.transformer import AutoRegressiveLM
|
from astrai.model.transformer import AutoRegressiveLM
|
||||||
from astrai.trainer.strategy import (
|
from astrai.trainer.strategy import (
|
||||||
SEQStrategy,
|
SEQStrategy,
|
||||||
SFTStrategy,
|
SFTStrategy,
|
||||||
StrategyFactory,
|
StrategyFactory,
|
||||||
_collect_moe_diagnostics,
|
_collect_moe_diagnostics,
|
||||||
_load_balancing_loss,
|
|
||||||
)
|
)
|
||||||
from tests.helpers import TINY_CONFIG
|
from tests.helpers import TINY_CONFIG
|
||||||
|
|
||||||
# ── helpers ──────────────────────────────────────────────────────────
|
|
||||||
|
|
||||||
|
|
||||||
def _make_tiny_moe_config(**overrides) -> AutoRegressiveLMConfig:
|
def _make_tiny_moe_config(**overrides) -> AutoRegressiveLMConfig:
|
||||||
return AutoRegressiveLMConfig(
|
return AutoRegressiveLMConfig(
|
||||||
@@ -43,14 +39,16 @@ def _make_model(config=None) -> AutoRegressiveLM:
|
|||||||
return AutoRegressiveLM(config)
|
return AutoRegressiveLM(config)
|
||||||
|
|
||||||
|
|
||||||
# ── _collect_moe_diagnostics unit tests ─────────────────────────────
|
def _router_stats(probs, topk_indices):
|
||||||
|
return {"probs": probs, "topk_indices": topk_indices}
|
||||||
|
|
||||||
|
|
||||||
def test_collect_moe_diagnostics_returns_all_keys():
|
def test_collect_moe_diagnostics_returns_all_keys():
|
||||||
"""_collect_moe_diagnostics should return the four expected keys."""
|
"""_collect_moe_diagnostics should return the four expected keys."""
|
||||||
# Simulate two MoE layers with uniform routing probabilities
|
# Simulate two MoE layers with uniform routing probabilities
|
||||||
probs = torch.ones(128, 4) / 4.0
|
probs = torch.ones(128, 4) / 4.0
|
||||||
diag = _collect_moe_diagnostics([probs, probs], top_k=2)
|
topk = torch.zeros(128, 2, dtype=torch.long)
|
||||||
|
diag = _collect_moe_diagnostics([_router_stats(probs, topk)] * 2)
|
||||||
|
|
||||||
assert set(diag.keys()) == {
|
assert set(diag.keys()) == {
|
||||||
"router_entropy",
|
"router_entropy",
|
||||||
@@ -64,11 +62,11 @@ def test_collect_moe_diagnostics_returns_all_keys():
|
|||||||
|
|
||||||
def test_collect_moe_diagnostics_empty_list():
|
def test_collect_moe_diagnostics_empty_list():
|
||||||
"""Empty list returns empty dict."""
|
"""Empty list returns empty dict."""
|
||||||
assert _collect_moe_diagnostics([], top_k=2) == {}
|
assert _collect_moe_diagnostics([]) == {}
|
||||||
|
|
||||||
|
|
||||||
def test_collect_moe_diagnostics_uniform_routing():
|
def test_collect_moe_diagnostics_uniform_routing():
|
||||||
"""Uniform routing probabilities with top_k=2 → tie-breaking by index.
|
"""Uniform routing with top_k=2 → tie-breaking by index.
|
||||||
|
|
||||||
torch.topk breaks ties by index, so with equal probabilities
|
torch.topk breaks ties by index, so with equal probabilities
|
||||||
experts 0 and 1 always win over experts 2 and 3:
|
experts 0 and 1 always win over experts 2 and 3:
|
||||||
@@ -77,7 +75,8 @@ def test_collect_moe_diagnostics_uniform_routing():
|
|||||||
- load_imbalance_max = 2.0
|
- load_imbalance_max = 2.0
|
||||||
"""
|
"""
|
||||||
probs = torch.ones(128, 4) / 4.0
|
probs = torch.ones(128, 4) / 4.0
|
||||||
diag = _collect_moe_diagnostics([probs], top_k=2)
|
topk = torch.tensor([[0, 1]] * 128)
|
||||||
|
diag = _collect_moe_diagnostics([_router_stats(probs, topk)])
|
||||||
|
|
||||||
assert diag["dead_expert_fraction"] == pytest.approx(0.5, abs=1e-6)
|
assert diag["dead_expert_fraction"] == pytest.approx(0.5, abs=1e-6)
|
||||||
assert diag["load_imbalance_mean"] == pytest.approx(1.0, abs=1e-6)
|
assert diag["load_imbalance_mean"] == pytest.approx(1.0, abs=1e-6)
|
||||||
@@ -88,38 +87,41 @@ def test_collect_moe_diagnostics_max_entropy():
|
|||||||
"""Uniform probabilities should give log(num_experts) entropy."""
|
"""Uniform probabilities should give log(num_experts) entropy."""
|
||||||
num_experts = 4
|
num_experts = 4
|
||||||
probs = torch.ones(128, num_experts) / num_experts
|
probs = torch.ones(128, num_experts) / num_experts
|
||||||
diag = _collect_moe_diagnostics([probs], top_k=2)
|
topk = torch.zeros(128, 2, dtype=torch.long)
|
||||||
|
diag = _collect_moe_diagnostics([_router_stats(probs, topk)])
|
||||||
expected_entropy = float(torch.log(torch.tensor(num_experts, dtype=torch.float32)))
|
expected_entropy = float(torch.log(torch.tensor(num_experts, dtype=torch.float32)))
|
||||||
assert diag["router_entropy"] == pytest.approx(expected_entropy, abs=1e-5)
|
assert diag["router_entropy"] == pytest.approx(expected_entropy, abs=1e-5)
|
||||||
|
|
||||||
|
|
||||||
# ── _load_balancing_loss unit tests ──────────────────────────────────
|
def test_moe_metrics_flow_through_wrapped_model(device):
|
||||||
|
"""DDP-like wrappers (no .config / get_moe_router_probs) still collect MoE metrics."""
|
||||||
|
import torch.nn as nn
|
||||||
|
|
||||||
|
from astrai.trainer.strategy import SEQStrategy
|
||||||
|
|
||||||
def test_load_balancing_loss_shape_and_range():
|
class ForwardOnlyWrapper(nn.Module):
|
||||||
"""Verify _load_balancing_loss returns a non-negative scalar tensor."""
|
def __init__(self, model):
|
||||||
probs = torch.randn(64, 8).softmax(dim=-1)
|
super().__init__()
|
||||||
loss = _load_balancing_loss(probs)
|
self.module = model
|
||||||
assert loss.ndim == 0
|
|
||||||
assert loss.item() >= 0
|
|
||||||
|
|
||||||
|
def forward(self, *args, **kwargs):
|
||||||
|
return self.module(*args, **kwargs)
|
||||||
|
|
||||||
def test_load_balancing_loss_uniform_minimum():
|
config = _make_tiny_moe_config()
|
||||||
"""Uniform routing gives the lowest possible load balancing loss."""
|
model = AutoRegressiveLM(config).to(device)
|
||||||
probs = torch.ones(64, 8) / 8.0
|
wrapped = ForwardOnlyWrapper(model)
|
||||||
loss = _load_balancing_loss(probs).item()
|
wrapped.train()
|
||||||
|
|
||||||
# Very skewed routing should give higher loss
|
strategy = SEQStrategy(wrapped, device, moe_aux_loss_coef=0.01)
|
||||||
skewed = torch.zeros(64, 8)
|
output = strategy.compute_loss_output(
|
||||||
skewed[:, 0] = 1.0
|
{
|
||||||
skewed[:, 1] = 1.0
|
"input_ids": torch.randint(0, config.vocab_size, (2, 8)),
|
||||||
skewed = skewed / skewed.sum(dim=-1, keepdim=True)
|
"target_ids": torch.randint(0, config.vocab_size, (2, 8)),
|
||||||
skewed_loss = _load_balancing_loss(skewed).item()
|
}
|
||||||
|
)
|
||||||
|
|
||||||
assert loss < skewed_loss
|
assert "moe_aux_loss" in output["metrics"]
|
||||||
|
assert "router_entropy" in strategy._moe_metrics
|
||||||
|
|
||||||
# ── SEQStrategy integration tests ────────────────────────────────────
|
|
||||||
|
|
||||||
|
|
||||||
class TestSEQStrategyMoE:
|
class TestSEQStrategyMoE:
|
||||||
@@ -255,9 +257,6 @@ class TestSEQStrategyMoE:
|
|||||||
assert strategy._moe_metrics == {}
|
assert strategy._moe_metrics == {}
|
||||||
|
|
||||||
|
|
||||||
# ── SFTStrategy integration tests ────────────────────────────────────
|
|
||||||
|
|
||||||
|
|
||||||
class TestSFTStrategyMoE:
|
class TestSFTStrategyMoE:
|
||||||
"""End‑to‑end tests for SFTStrategy with MoE aux loss."""
|
"""End‑to‑end tests for SFTStrategy with MoE aux loss."""
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user