"""Configuration for the sequence-routed SLMoE causal language model.""" from __future__ import annotations import math from transformers import PretrainedConfig class SLMoEConfig(PretrainedConfig): model_type = "slmoe" keys_to_ignore_at_inference = ["router_aux_loss", "router_z_loss"] def __init__( self, vocab_size: int = 8192, hidden_size: int = 256, num_hidden_layers: int = 8, num_attention_heads: int = 8, num_key_value_heads: int = 2, head_dim: int = 32, num_experts: int = 64, num_experts_per_sequence: int = 13, expert_intermediate_size: int = 56, router_prefix_length: int = 32, router_jitter_noise: float = 0.01, router_aux_loss_coeff: float = 0.01, router_z_loss_coeff: float = 1e-3, expert_output_scale: float | None = None, max_position_embeddings: int = 4096, rope_theta: float = 100000.0, rms_norm_eps: float = 1e-6, initializer_range: float = 0.02, tie_word_embeddings: bool = True, use_cache: bool = True, bos_token_id: int = 1, eos_token_id: int = 2, pad_token_id: int = 0, unk_token_id: int = 3, **kwargs, ): if hidden_size != num_attention_heads * head_dim: raise ValueError("hidden_size must equal num_attention_heads * head_dim") if num_attention_heads % num_key_value_heads: raise ValueError("num_attention_heads must be divisible by num_key_value_heads") if not 0 < num_experts_per_sequence <= num_experts: raise ValueError("num_experts_per_sequence must be in [1, num_experts]") if router_prefix_length < 1: raise ValueError("router_prefix_length must be positive") self.vocab_size = vocab_size self.hidden_size = hidden_size self.num_hidden_layers = num_hidden_layers self.num_attention_heads = num_attention_heads self.num_key_value_heads = num_key_value_heads self.head_dim = head_dim self.num_experts = num_experts self.num_experts_per_sequence = num_experts_per_sequence self.expert_intermediate_size = expert_intermediate_size self.router_prefix_length = router_prefix_length self.router_jitter_noise = router_jitter_noise self.router_aux_loss_coeff = router_aux_loss_coeff self.router_z_loss_coeff = router_z_loss_coeff self.expert_output_scale = ( math.sqrt(num_experts_per_sequence) if expert_output_scale is None else expert_output_scale ) self.max_position_embeddings = max_position_embeddings self.rope_theta = rope_theta self.rms_norm_eps = rms_norm_eps self.initializer_range = initializer_range self.use_cache = use_cache super().__init__( tie_word_embeddings=tie_word_embeddings, bos_token_id=bos_token_id, eos_token_id=eos_token_id, pad_token_id=pad_token_id, unk_token_id=unk_token_id, **kwargs, ) def parameter_counts(self) -> dict[str, int]: """Return analytical total and single-sequence active parameter counts.""" hidden = self.hidden_size query_width = self.num_attention_heads * self.head_dim kv_width = self.num_key_value_heads * self.head_dim embeddings = self.vocab_size * hidden attention = ( hidden * query_width + 2 * hidden * kv_width + query_width * hidden ) attention_norms = 2 * self.head_dim block_norms = 2 * hidden one_expert = 3 * hidden * self.expert_intermediate_size all_experts = self.num_experts * one_expert active_experts = self.num_experts_per_sequence * one_expert router = hidden + hidden * self.num_experts final_norm = hidden output_head = 0 if self.tie_word_embeddings else embeddings shared_per_layer = attention + attention_norms + block_norms total = ( embeddings + self.num_hidden_layers * (shared_per_layer + all_experts) + router + final_norm + output_head ) active = ( embeddings + self.num_hidden_layers * (shared_per_layer + active_experts) + router + final_norm + output_head ) return { "total": total, "active_per_sequence": active, "one_expert_path": self.num_hidden_layers * one_expert, } SLMoEConfig.register_for_auto_class("AutoConfig") __all__ = ["SLMoEConfig"]