diff --git a/astrai/config/model_config.py b/astrai/config/model_config.py index 911ebd1..b982769 100644 --- a/astrai/config/model_config.py +++ b/astrai/config/model_config.py @@ -65,7 +65,7 @@ class AutoRegressiveLMConfig(BaseModelConfig): topk_method (Optional[str]): Top-k routing method, MoE only. Defaults to None. moe_intermediate_size (Optional[int]): Expert hidden dim, defaults to intermediate_size if None. MoE only. shared_expert_intermediate_size (Optional[int]): Shared expert hidden dim, defaults to intermediate_size if None. MoE only. - norm_topk_prob (bool): Normalize top-k routing probabilities. Defaults to False. + norm_topk_prob (bool): Normalize top-k routing probabilities. Defaults to True. decoder_sparse_step (int): Frequency of MoE layers, 1=every layer. Defaults to 1. mlp_only_layers (Optional[list[int]]): Layer indices using dense MLP instead of MoE. Defaults to None. """ @@ -94,7 +94,7 @@ class AutoRegressiveLMConfig(BaseModelConfig): topk_method: Optional[str] = None moe_intermediate_size: Optional[int] = None shared_expert_intermediate_size: Optional[int] = None - norm_topk_prob: bool = False + norm_topk_prob: bool = True decoder_sparse_step: int = 1 mlp_only_layers: Optional[list[int]] = None @@ -112,6 +112,12 @@ class AutoRegressiveLMConfig(BaseModelConfig): raise ValueError(f"ffn_type must be one of {sorted(_FFN_TYPES)}, got {v!r}") return v + @field_validator("decoder_sparse_step") + def _validate_decoder_sparse_step(cls, v: int) -> int: + if v < 1: + raise ValueError(f"decoder_sparse_step must be at least 1, got {v}") + return v + @dataclass @ConfigFactory.register("embedding") diff --git a/astrai/model/components/mlp.py b/astrai/model/components/mlp.py index 95f7d8f..e294fb8 100644 --- a/astrai/model/components/mlp.py +++ b/astrai/model/components/mlp.py @@ -40,7 +40,7 @@ class DeepSeekMoE(nn.Module): n_layers: int = 1, moe_intermediate_size: Optional[int] = None, shared_expert_intermediate_size: Optional[int] = None, - norm_topk_prob: bool = False, + norm_topk_prob: bool = True, ): super().__init__() self.dim = dim @@ -50,8 +50,14 @@ class DeepSeekMoE(nn.Module): self.topk_method = topk_method self.norm_topk_prob = norm_topk_prob - expert_dim_ffn = moe_intermediate_size if moe_intermediate_size is not None else dim_ffn - shared_dim_ffn = shared_expert_intermediate_size if shared_expert_intermediate_size is not None else dim_ffn + expert_dim_ffn = ( + moe_intermediate_size if moe_intermediate_size is not None else dim_ffn + ) + shared_dim_ffn = ( + shared_expert_intermediate_size + if shared_expert_intermediate_size is not None + else dim_ffn + ) self.router = Linear(dim, n_routed_experts, bias=False) moe_scale = 1 / max(n_shared_experts, 1) + 1 / n_activated_experts diff --git a/tests/module/test_forward_configs.py b/tests/module/test_forward_configs.py index a6f20d4..1213c48 100644 --- a/tests/module/test_forward_configs.py +++ b/tests/module/test_forward_configs.py @@ -246,3 +246,36 @@ def test_moe_custom_intermediate_shape(): assert expert.up.weight.shape[0] == 20 assert expert.gate.weight.shape[0] == 20 assert expert.down.weight.shape[1] == 20 + + +def test_moe_defaults_preserve_normalized_routing(): + from astrai.config.model_config import AutoRegressiveLMConfig + + config = AutoRegressiveLMConfig( + **TINY_CONFIG, + ffn_type="moe", + n_routed_experts=4, + n_shared_experts=1, + n_activated_experts=2, + topk_method="greedy", + ) + model = AutoRegressiveLM(config) + + assert config.norm_topk_prob is True + assert model.layers[0].mlp.norm_topk_prob is True + + +@pytest.mark.parametrize("decoder_sparse_step", [0, -1]) +def test_moe_rejects_invalid_decoder_sparse_step(decoder_sparse_step): + from pydantic import ValidationError + + from astrai.config.model_config import AutoRegressiveLMConfig + + with pytest.raises(ValidationError, match="decoder_sparse_step must be at least 1"): + AutoRegressiveLMConfig( + **TINY_CONFIG, + ffn_type="moe", + n_routed_experts=4, + n_activated_experts=2, + decoder_sparse_step=decoder_sparse_step, + )