# Tweaktron: Omni-Mythos — flat 3-layer composite # Layer 0: Comba (prelude) # Layer 1: MoM (body, gated_deltanet backend, returns router_logits) # Layer 2: Raven (coda) # # FLA layer internals are imported untouched. Block wiring mirrors the # upstream FLA modeling files (prenorm + fused RMSNorm residual + GatedMLP). from __future__ import annotations from typing import TYPE_CHECKING, Optional import torch import torch.nn as nn from transformers.modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast from transformers.modeling_utils import PreTrainedModel from transformers.utils import logging from fla.layers import MomAttention from fla.layers.comba import Comba from fla.layers.raven import Raven from fla.models.mom.modeling_mom import load_balancing_loss_func from fla.models.utils import Cache, FLAGenerationMixin from fla.modules import FusedCrossEntropyLoss, FusedLinearCrossEntropyLoss, RMSNorm from fla.modules import GatedMLP try: from transformers.modeling_layers import GradientCheckpointingLayer except ImportError: from fla.models.modeling_layers import GradientCheckpointingLayer from configuration_tweaktron import TweaktronOmniConfig if TYPE_CHECKING: from transformers.processing_utils import Unpack logger = logging.get_logger(__name__) class TweaktronOmniBlock(GradientCheckpointingLayer): """One prenorm block wrapping any FLA token mixer + GatedMLP. Handles both 3-tuple (Comba, Raven) and 4-tuple (MomAttention) mixer returns; router_logits is None for non-MoM layers. """ def __init__(self, config: TweaktronOmniConfig, layer_idx: int): super().__init__() self.layer_idx = layer_idx norm_cls = RMSNorm if config.fuse_norm else nn.RMSNorm self.attn_norm = norm_cls(config.hidden_size, eps=config.norm_eps) if layer_idx == 0: # ---- Comba prelude (FLA defaults) ---- self.attn = Comba( mode='chunk', hidden_size=config.hidden_size, expand_v=config.comba_expand_v, head_dim=config.comba_head_dim, num_heads=config.comba_num_heads, use_output_gate=config.comba_use_output_gate, use_output_correction=config.comba_use_output_correction, use_inner_decay=config.comba_use_inner_decay, correction_factor=config.comba_correction_factor, use_short_conv=True, conv_size=config.comba_conv_size, norm_eps=config.norm_eps, layer_idx=layer_idx, ) elif layer_idx == 1: # ---- MoM body (FLA defaults per project constraints) ---- self.attn = MomAttention( mode='chunk', hidden_size=config.hidden_size, expand_v=config.mom_expand_v, head_dim=config.mom_head_dim, num_heads=config.mom_num_heads, use_output_gate=config.mom_use_output_gate, use_short_conv=True, conv_size=config.mom_conv_size, norm_eps=config.norm_eps, layer_idx=layer_idx, num_memories=config.mom_num_memories, topk=config.mom_topk, capacity=config.mom_capacity, shared_mem=config.mom_shared_mem, single_kv_proj=config.mom_single_kv_proj, ) elif layer_idx == 2: # ---- Raven coda (FLA defaults) ---- self.attn = Raven( mode='chunk', hidden_size=config.hidden_size, expand_k=config.raven_expand_k, expand_v=config.raven_expand_v, num_heads=config.raven_num_heads, num_kv_heads=config.raven_num_kv_heads, num_slots=config.raven_num_slots, norm_eps=config.norm_eps, gate_logit_normalizer=config.raven_gate_logit_normalizer, feature_map=config.raven_feature_map, use_output_gate=config.raven_use_output_gate, decay_type=config.raven_decay_type, topk=config.raven_topk, bias_rmm=config.raven_bias_rmm, add_gumbel_noise=config.raven_add_gumbel_noise, router_score=config.raven_router_score, router_type=config.raven_router_type, max_position_embeddings=config.max_position_embeddings, fuse_norm=config.fuse_norm, layer_idx=layer_idx, ) else: raise ValueError(f"TweaktronOmni is a fixed 3-layer stack; got layer_idx={layer_idx}") self.mlp_norm = norm_cls(config.hidden_size, eps=config.norm_eps) self.mlp = GatedMLP( hidden_size=config.hidden_size, hidden_ratio=config.hidden_ratio, intermediate_size=config.intermediate_size, hidden_act=config.hidden_act, fuse_swiglu=config.fuse_swiglu, ) def forward( self, hidden_states: torch.Tensor, attention_mask: torch.Tensor | None = None, past_key_values: Cache | list[torch.FloatTensor] | None = None, use_cache: bool | None = False, output_attentions: bool | None = False, **kwargs: Unpack[dict], ): residual = hidden_states hidden_states = self.attn_norm(hidden_states) attn_out = self.attn( hidden_states=hidden_states, attention_mask=attention_mask, past_key_values=past_key_values, use_cache=use_cache, output_attentions=output_attentions, **kwargs, ) # MomAttention -> (o, attn, cache, router_logits); Comba/Raven -> (o, attn, cache) if len(attn_out) == 4: hidden_states, attentions, past_key_values, router_logits = attn_out else: hidden_states, attentions, past_key_values = attn_out router_logits = None if isinstance(self.mlp_norm, RMSNorm): hidden_states, residual = self.mlp_norm(hidden_states, residual, True) else: hidden_states = residual + hidden_states residual = hidden_states hidden_states = self.mlp_norm(hidden_states) hidden_states = self.mlp(hidden_states, **kwargs) hidden_states = residual + hidden_states return hidden_states, attentions, past_key_values, router_logits class TweaktronOmniPreTrainedModel(PreTrainedModel): config_class = TweaktronOmniConfig base_model_prefix = 'model' supports_gradient_checkpointing = True _no_split_modules = ['TweaktronOmniBlock'] _supports_cache_class = True def _init_weights(self, module: nn.Module): std = self.config.initializer_range if isinstance(module, (nn.Linear, nn.Conv1d)): nn.init.normal_(module.weight, mean=0.0, std=std) if module.bias is not None: nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): nn.init.normal_(module.weight, mean=0.0, std=std) elif hasattr(module, 'reset_parameters'): module.reset_parameters() class TweaktronOmniModel(TweaktronOmniPreTrainedModel): def __init__(self, config: TweaktronOmniConfig): super().__init__(config) self.padding_idx = config.pad_token_id self.vocab_size = config.vocab_size self.embeddings = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx) self.layers = nn.ModuleList([TweaktronOmniBlock(config, i) for i in range(3)]) self.norm = (RMSNorm if config.fuse_norm else nn.RMSNorm)(config.hidden_size, eps=config.norm_eps) self.gradient_checkpointing = False self.post_init() def get_input_embeddings(self): return self.embeddings def set_input_embeddings(self, value): self.embeddings = value def forward( self, input_ids: torch.LongTensor | None = None, attention_mask: torch.Tensor | None = None, inputs_embeds: torch.Tensor | None = None, past_key_values: Cache | list[torch.FloatTensor] | None = None, use_cache: bool | None = None, output_attentions: bool | None = None, output_hidden_states: bool | None = None, return_dict: bool | None = None, **kwargs: Unpack[dict], ): output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions output_hidden_states = ( output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states ) use_cache = use_cache if use_cache is not None else (self.config.use_cache if not self.training else False) return_dict = return_dict if return_dict is not None else self.config.use_return_dict if (input_ids is None) == (inputs_embeds is None): raise ValueError("Provide exactly one of input_ids or inputs_embeds") if inputs_embeds is None: inputs_embeds = self.embeddings(input_ids) hidden_states = inputs_embeds if use_cache and not isinstance(past_key_values, Cache): past_key_values = Cache.from_legacy_cache(past_key_values) if self.gradient_checkpointing and self.training and use_cache: logger.warning_once("`use_cache=True` is incompatible with gradient checkpointing; setting False") use_cache = False all_hidden_states = () if output_hidden_states else None all_attns = () if output_attentions else None all_router_logits = () for layer in self.layers: if output_hidden_states: all_hidden_states += (hidden_states,) hidden_states, attentions, past_key_values, router_logits = layer( hidden_states, attention_mask=attention_mask, past_key_values=past_key_values, use_cache=use_cache, output_attentions=output_attentions, **kwargs, ) if router_logits is not None: all_router_logits += (router_logits,) if output_attentions: all_attns += (attentions,) hidden_states = self.norm(hidden_states) if output_hidden_states: all_hidden_states += (hidden_states,) if not return_dict: return tuple(x for x in [hidden_states, past_key_values, all_hidden_states, all_attns] if x is not None) out = BaseModelOutputWithPast( last_hidden_state=hidden_states, past_key_values=past_key_values, hidden_states=all_hidden_states, attentions=all_attns, ) out.router_logits = all_router_logits return out class TweaktronOmniForCausalLM(TweaktronOmniPreTrainedModel, FLAGenerationMixin): _tied_weights_keys = [] # tie_word_embeddings must stay False def __init__(self, config: TweaktronOmniConfig): super().__init__(config) self.model = TweaktronOmniModel(config) self.vocab_size = config.vocab_size self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) self.num_memories = config.mom_num_memories self.topk = config.mom_topk self.aux_loss_scale = config.aux_loss_scale self.post_init() def get_input_embeddings(self): return self.model.embeddings def set_input_embeddings(self, value): self.model.embeddings = value def get_output_embeddings(self): return self.lm_head def set_output_embeddings(self, new_embeddings): self.lm_head = new_embeddings def get_decoder(self): return self.model def set_decoder(self, decoder): self.model = decoder def forward( self, input_ids: torch.LongTensor = None, attention_mask: torch.Tensor | None = None, inputs_embeds: torch.Tensor | None = None, past_key_values: Cache | list[torch.FloatTensor] | None = None, labels: torch.LongTensor | None = None, use_cache: bool | None = None, output_attentions: bool | None = None, output_hidden_states: bool | None = None, return_dict: bool | None = None, num_logits_to_keep: int | None = 0, **kwargs: Unpack[dict], ) -> tuple | CausalLMOutputWithPast: return_dict = return_dict if return_dict is not None else self.config.use_return_dict outputs = self.model( input_ids=input_ids, attention_mask=attention_mask, inputs_embeds=inputs_embeds, past_key_values=past_key_values, use_cache=use_cache, output_attentions=output_attentions, output_hidden_states=output_hidden_states, return_dict=True, **kwargs, ) hidden_states = outputs.last_hidden_state fuse_linear_and_cross_entropy = self.config.fuse_linear_cross_entropy and self.training logits = None if fuse_linear_and_cross_entropy else self.lm_head(hidden_states[:, -num_logits_to_keep:]) loss, aux_loss = None, None if labels is not None: if self.config.fuse_cross_entropy: loss_fct = FusedCrossEntropyLoss(inplace_backward=True) elif fuse_linear_and_cross_entropy: loss_fct = FusedLinearCrossEntropyLoss() else: loss_fct = nn.CrossEntropyLoss() labels = labels.to(hidden_states.device) labels = torch.cat((labels[..., 1:], torch.full_like(labels[:, :1], loss_fct.ignore_index)), 1) if fuse_linear_and_cross_entropy: loss = loss_fct( hidden_states.view(-1, self.config.hidden_size), labels.view(-1), self.lm_head.weight, self.lm_head.bias, ) else: loss = loss_fct(logits.view(-1, self.config.vocab_size), labels.view(-1)) router_logits = getattr(outputs, 'router_logits', ()) if router_logits: aux_loss = load_balancing_loss_func( router_logits, self.num_memories, self.topk, attention_mask, ) loss = loss + aux_loss.to(loss.device) * self.aux_loss_scale if not return_dict: output = (logits, outputs.past_key_values, outputs.hidden_states, outputs.attentions) return (loss,) + output if loss is not None else output return CausalLMOutputWithPast( loss=loss, logits=logits, past_key_values=outputs.past_key_values, hidden_states=outputs.hidden_states, attentions=outputs.attentions, )