slmoe-test / configuration_slmoe.py
Banaxi-Tech's picture
Publish sequence-routed SLMoE architecture
bd0ade7 verified
Raw
History Blame Contribute Delete
4.73 kB
"""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"]