from transformers.configuration_utils import PretrainedConfig from transformers.utils import logging logger = logging.get_logger(__name__) class ScrapeGoatConfig(PretrainedConfig): model_type = "scrapegoat" keys_to_ignore_at_inference = ["past_key_values"] def __init__( self, vocab_size=32000, hidden_size=4096, num_hidden_layers=81, hidden_act="silu", max_position_embeddings=262144, initializer_range=0.02, rms_norm_eps=1e-06, use_cache=True, pad_token_id=0, bos_token_id=1, eos_token_id=2, tie_word_embeddings=False, rope_theta=10000.0, rope_scaling=None, attention_bias=False, attention_dropout=0.0, # Track A (Ornith) track_a_num_attention_heads=32, track_a_num_key_value_heads=2, track_a_head_dim=256, track_a_num_experts=512, track_a_moe_intermediate_size=1024, # Track B (Hy3) track_b_num_attention_heads=64, track_b_num_key_value_heads=8, track_b_head_dim=128, track_b_num_experts=192, num_experts=704, track_b_moe_intermediate_size=1536, track_b_intermediate_size=13312, # MoE shared num_experts_per_tok=8, output_router_logits=False, router_aux_loss_coef=0.001, # KDA (Kimi Delta Attention) kda_head_dim=256, kda_conv_kernel=3, kda_gqa_layers=None, # Quantile Balancing quantile_balancing=True, qb_iterations=5, # Attention Residuals attn_residual=True, attn_res_blocks=8, # StableMoE stable_moe_stage=1, # MoM (Mixture-of-Memories) mom_enabled=False, mom_num_memories=4, mom_active_memories=2, mom_shared_memory=True, mom_load_balancing=True, # R3 Routing Replay stable_moe_r3=False, stable_moe_r3_cache=True, # DSpark dspark_block_size=6, dspark_noise_token_id=0, dspark_target_layer_ids=(), dspark_markov_rank=256, num_attention_heads=32, num_key_value_heads=8, **kwargs, ): self.vocab_size = vocab_size self.hidden_size = hidden_size self.num_hidden_layers = num_hidden_layers self.hidden_act = hidden_act self.max_position_embeddings = max_position_embeddings self.initializer_range = initializer_range self.rms_norm_eps = rms_norm_eps self.use_cache = use_cache self.return_dict = kwargs.pop("return_dict", True) self.rope_theta = rope_theta self.rope_scaling = rope_scaling self.attention_bias = attention_bias self.attention_dropout = attention_dropout self.track_a_num_attention_heads = track_a_num_attention_heads self.track_a_num_key_value_heads = track_a_num_key_value_heads self.track_a_head_dim = track_a_head_dim self.track_a_num_experts = track_a_num_experts self.track_a_moe_intermediate_size = track_a_moe_intermediate_size self.track_b_num_attention_heads = track_b_num_attention_heads self.track_b_num_key_value_heads = track_b_num_key_value_heads self.track_b_head_dim = track_b_head_dim self.track_b_num_experts = track_b_num_experts self.num_experts = num_experts self.track_b_moe_intermediate_size = track_b_moe_intermediate_size self.track_b_intermediate_size = track_b_intermediate_size self.num_experts_per_tok = num_experts_per_tok self.output_router_logits = output_router_logits self.router_aux_loss_coef = router_aux_loss_coef # KDA self.kda_head_dim = kda_head_dim self.kda_conv_kernel = kda_conv_kernel if kda_gqa_layers is None: self.kda_gqa_layers = list(range(0, 81, 4)) else: self.kda_gqa_layers = kda_gqa_layers # Quantile Balancing self.quantile_balancing = quantile_balancing self.qb_iterations = qb_iterations # Attention Residuals self.attn_residual = attn_residual self.attn_res_blocks = attn_res_blocks # StableMoE self.stable_moe_stage = stable_moe_stage self.stable_moe_r3 = stable_moe_r3 self.stable_moe_r3_cache = stable_moe_r3_cache # MoM self.mom_enabled = mom_enabled self.mom_num_memories = mom_num_memories self.mom_active_memories = mom_active_memories self.mom_shared_memory = mom_shared_memory self.mom_load_balancing = mom_load_balancing self.dspark_block_size = dspark_block_size self.dspark_noise_token_id = dspark_noise_token_id self.dspark_target_layer_ids = dspark_target_layer_ids self.dspark_markov_rank = dspark_markov_rank self.num_attention_heads = num_attention_heads self.num_key_value_heads = num_key_value_heads 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, )