AliceAI-Foundation-80B-A3B-Base / configuration_alice_ai.py
fdrose's picture
Initial release
aeac2c6
Raw History Blame
4.87 kB
from transformers.configuration_utils import PretrainedConfig
class AliceAIConfig(PretrainedConfig):
model_type = "alice_ai"
def __init__(
self,
vocab_size: int = 129024,
hidden_size: int = 2048,
num_hidden_layers: int = 48,
num_attention_heads: int = 16,
num_key_value_heads: int = 2,
head_dim: int = 256,
linear_num_key_heads: int = 32,
linear_num_value_heads: int = 32,
linear_key_head_dim: int = 128,
linear_value_head_dim: int = 128,
linear_conv_kernel_dim: int = 4,
num_experts: int = 512,
num_experts_per_tok: int = 10,
moe_intermediate_size: int = 512,
shared_expert_intermediate_size: int = 512,
block_attn_res_block_size: int = 4,
router_score_function: str = "sigmoid",
router_bias_correction: bool = True,
kda_allow_negative_eigenvalues: bool = False,
max_position_embeddings: int = 262144,
rope_theta: float = 1_000_000.0,
partial_rotary_factor: float = 0.25,
rms_norm_eps: float = 1e-6,
hidden_act: str = "silu",
initializer_range: float = 0.02,
attention_dropout: float = 0.0,
use_cache: bool = True,
output_router_logits: bool = False,
layer_types: list[str] | None = None,
tie_word_embeddings: bool = False,
pad_token_id: int | None = None,
bos_token_id: int | None = None,
eos_token_id: int | list[int] | None = None,
**kwargs,
) -> None:
if layer_types is None:
layer_types = [
"full_attention" if (layer_idx + 1) % 4 == 0 else "linear_attention"
for layer_idx in range(num_hidden_layers)
]
super().__init__(
pad_token_id=pad_token_id,
bos_token_id=bos_token_id,
eos_token_id=eos_token_id,
tie_word_embeddings=tie_word_embeddings,
**kwargs,
)
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.linear_num_key_heads = linear_num_key_heads
self.linear_num_value_heads = linear_num_value_heads
self.linear_key_head_dim = linear_key_head_dim
self.linear_value_head_dim = linear_value_head_dim
self.linear_conv_kernel_dim = linear_conv_kernel_dim
self.num_experts = num_experts
self.num_experts_per_tok = num_experts_per_tok
self.moe_intermediate_size = moe_intermediate_size
self.shared_expert_intermediate_size = shared_expert_intermediate_size
self.block_attn_res_block_size = block_attn_res_block_size
self.router_score_function = router_score_function
self.router_bias_correction = router_bias_correction
self.kda_allow_negative_eigenvalues = kda_allow_negative_eigenvalues
self.max_position_embeddings = max_position_embeddings
self.rope_theta = rope_theta
self.partial_rotary_factor = partial_rotary_factor
self.rms_norm_eps = rms_norm_eps
self.hidden_act = hidden_act
self.initializer_range = initializer_range
self.attention_dropout = attention_dropout
self.use_cache = use_cache
self.output_router_logits = output_router_logits
self.layer_types = layer_types
self.number_of_conv_states = 3
self._validate_fields()
def _validate_fields(self) -> None:
if self.block_attn_res_block_size <= 0:
raise ValueError("block_attn_res_block_size must be positive")
if self.linear_conv_kernel_dim < 2:
raise ValueError("linear_conv_kernel_dim must be at least 2")
if self.router_score_function != "sigmoid":
raise ValueError("This architecture requires sigmoid routing")
if not 0 < self.num_experts_per_tok <= self.num_experts:
raise ValueError("num_experts_per_tok must be between 1 and num_experts")
if self.num_attention_heads % self.num_key_value_heads != 0:
raise ValueError(
"num_attention_heads must be divisible by num_key_value_heads"
)
if self.linear_num_value_heads % self.linear_num_key_heads != 0:
raise ValueError(
"linear_num_value_heads must be divisible by linear_num_key_heads"
)
if len(self.layer_types) != self.num_hidden_layers:
raise ValueError("layer_types must contain one entry per hidden layer")
unknown = set(self.layer_types) - {"linear_attention", "full_attention"}
if unknown:
raise ValueError(f"Unsupported layer types: {sorted(unknown)}")