# Tweaktron: Omni-Mythos — flat 3-layer FLA composite # Layer 0: Comba (prelude) | Layer 1: MoM (body) | Layer 2: Raven (coda) # FLA layer internals are NOT modified; all layers imported from fla as-is. from transformers.configuration_utils import PretrainedConfig class TweaktronOmniConfig(PretrainedConfig): model_type = 'tweaktron' keys_to_ignore_at_inference = ['past_key_values'] def __init__( self, hidden_size: int = 2048, num_heads: int = 22, # shared block settings hidden_ratio: int | None = 4, intermediate_size: int = 8192, hidden_act: str = "swish", norm_eps: float = 1e-6, max_position_embeddings: int = 4096, # --- Comba (layer 0) — FLA defaults --- comba_num_heads: int = 6, comba_head_dim: int = 256, comba_expand_v: float = 2.0, comba_use_output_gate: bool = True, comba_use_output_correction: bool = True, comba_use_inner_decay: bool = True, comba_correction_factor: float = 1.0, comba_conv_size: int = 4, # --- MoM (layer 1) — FLA defaults per project constraints --- mom_num_heads: int = 8, mom_head_dim: int = 256, mom_expand_v: float = 1.0, mom_num_memories: int = 8, mom_topk: int = 2, mom_capacity: float = 1.0, mom_shared_mem: bool = True, mom_single_kv_proj: bool = False, mom_use_output_gate: bool = True, mom_conv_size: int = 4, aux_loss_scale: float = 0.01, # --- Raven (layer 2) — FLA defaults --- raven_num_heads: int = 4, raven_num_kv_heads: int | None = None, raven_num_slots: int = 64, raven_expand_k: float = 1.0, raven_expand_v: float = 1.0, raven_feature_map: str = 'swish', raven_use_output_gate: bool = False, raven_decay_type: str = 'Mamba2', raven_topk: int = 32, raven_bias_rmm: bool = False, raven_add_gumbel_noise: bool = True, raven_router_score: str = 'sigmoid', raven_router_type: str = 'lin', raven_gate_logit_normalizer: int = 8, # tokenizer / heads vocab_size: int = 32016, pad_token_id: int | None = None, bos_token_id: int = 1, eos_token_id: int = 2, tie_word_embeddings: bool = False, # MUST stay False (transformers 5.x _tied_weights_keys bug) initializer_range: float = 0.02, use_cache: bool = True, # fusion flags fuse_norm: bool = True, fuse_swiglu: bool = True, fuse_cross_entropy: bool = True, fuse_linear_cross_entropy: bool = False, **kwargs, ): self.hidden_size = hidden_size self.num_hidden_layers = 3 self.num_heads = 22 self.hidden_ratio = hidden_ratio self.intermediate_size = intermediate_size self.hidden_act = hidden_act self.norm_eps = norm_eps self.max_position_embeddings = max_position_embeddings self.comba_num_heads = comba_num_heads self.comba_head_dim = comba_head_dim self.comba_expand_v = comba_expand_v self.comba_use_output_gate = comba_use_output_gate self.comba_use_output_correction = comba_use_output_correction self.comba_use_inner_decay = comba_use_inner_decay self.comba_correction_factor = comba_correction_factor self.comba_conv_size = comba_conv_size self.mom_num_heads = mom_num_heads self.mom_head_dim = mom_head_dim self.mom_expand_v = mom_expand_v self.mom_num_memories = mom_num_memories self.mom_topk = mom_topk self.mom_capacity = mom_capacity self.mom_shared_mem = mom_shared_mem self.mom_single_kv_proj = mom_single_kv_proj self.mom_use_output_gate = mom_use_output_gate self.mom_conv_size = mom_conv_size self.aux_loss_scale = aux_loss_scale self.raven_num_heads = raven_num_heads self.raven_num_kv_heads = raven_num_kv_heads self.raven_num_slots = raven_num_slots self.raven_expand_k = raven_expand_k self.raven_expand_v = raven_expand_v self.raven_feature_map = raven_feature_map self.raven_use_output_gate = raven_use_output_gate self.raven_decay_type = raven_decay_type self.raven_topk = raven_topk self.raven_bias_rmm = raven_bias_rmm self.raven_add_gumbel_noise = raven_add_gumbel_noise self.raven_router_score = raven_router_score self.raven_router_type = raven_router_type self.raven_gate_logit_normalizer = raven_gate_logit_normalizer self.vocab_size = vocab_size self.initializer_range = initializer_range self.use_cache = use_cache self.fuse_norm = fuse_norm self.fuse_swiglu = fuse_swiglu self.fuse_cross_entropy = fuse_cross_entropy self.fuse_linear_cross_entropy = fuse_linear_cross_entropy if fuse_cross_entropy and fuse_linear_cross_entropy: raise ValueError("`fuse_cross_entropy` and `fuse_linear_cross_entropy` cannot both be True.") 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, )