Scrapegoat-Tiny-Coder / configuration_scrapegoat.py
Nathan9's picture
Upload 9 files
5216a17 verified
Raw
History Blame Contribute Delete
5.23 kB
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,
)