""" J.A.R.V.I.S. TITAN 14.8B MoE — Custom Model Architecture (DeepSeekMoE Style) Hyper-Optimized with Group-Disjoint Constrained Routing (0.00% Twin Collisions) """ import torch import torch.nn as nn import torch.nn.functional as F from transformers.models.qwen2.configuration_qwen2 import Qwen2Config from transformers.models.qwen2.modeling_qwen2 import ( Qwen2ForCausalLM, Qwen2Model, Qwen2DecoderLayer ) class JarvisTitanMoEConfig(Qwen2Config): model_type = "jarvis_titan_moe" keys_to_ignore_at_loading = ["rotary_emb.inv_freq"] def __init__( self, vocab_size=152064, hidden_size=3584, intermediate_size=18944, num_hidden_layers=28, num_attention_heads=28, num_key_value_heads=4, hidden_act="silu", max_position_embeddings=32768, initializer_range=0.02, rms_norm_eps=1e-6, use_cache=True, rope_theta=1000000.0, rope_scaling=None, attention_dropout=0.0, num_routed_experts=8, num_shared_experts=1, num_experts_per_tok=2, expert_intermediate_size=4736, routed_scaling_factor=1.5, group_disjoint_routing=True, **kwargs ): self.num_routed_experts = num_routed_experts self.num_shared_experts = num_shared_experts self.num_experts_per_tok = num_experts_per_tok self.expert_intermediate_size = expert_intermediate_size self.routed_scaling_factor = routed_scaling_factor self.group_disjoint_routing = group_disjoint_routing super().__init__( vocab_size=vocab_size, hidden_size=hidden_size, intermediate_size=intermediate_size, num_hidden_layers=num_hidden_layers, num_attention_heads=num_attention_heads, num_key_value_heads=num_key_value_heads, hidden_act=hidden_act, max_position_embeddings=max_position_embeddings, initializer_range=initializer_range, rms_norm_eps=rms_norm_eps, use_cache=use_cache, rope_theta=rope_theta, rope_scaling=rope_scaling, attention_dropout=attention_dropout, **kwargs ) class DeepSeekMoEMLP(nn.Module): def __init__(self, config: JarvisTitanMoEConfig): super().__init__() self.hidden_size = config.hidden_size self.expert_size = config.expert_intermediate_size self.num_routed = config.num_routed_experts self.top_k = config.num_experts_per_tok self.act_fn = nn.SiLU() self.group_disjoint = getattr(config, "group_disjoint_routing", True) # 1. Shared Expert self.shared_gate = nn.Linear(self.hidden_size, self.expert_size * config.num_shared_experts, bias=False) self.shared_up = nn.Linear(self.hidden_size, self.expert_size * config.num_shared_experts, bias=False) self.shared_down = nn.Linear(self.expert_size * config.num_shared_experts, self.hidden_size, bias=False) # 2. Routed Experts self.routed_gate = nn.Parameter(torch.empty(self.num_routed, self.expert_size, self.hidden_size)) self.routed_up = nn.Parameter(torch.empty(self.num_routed, self.expert_size, self.hidden_size)) self.routed_down = nn.Parameter(torch.empty(self.num_routed, self.hidden_size, self.expert_size)) # 3. Router Gate self.router = nn.Linear(self.hidden_size, self.num_routed, bias=False) def forward(self, x): batch_size, seq_len, hidden_dim = x.shape x_flat = x.view(-1, hidden_dim) # Shared Expert forward shared_out = self.shared_down(self.act_fn(self.shared_gate(x_flat)) * self.shared_up(x_flat)) # Router Logits router_logits = self.router(x_flat) routing_weights = F.softmax(router_logits, dim=-1) if self.group_disjoint: # Group-Disjoint Constrained Selection (0.00% Twin Collisions!) e1 = torch.argmax(routing_weights, dim=-1) # [N] p1 = torch.gather(routing_weights, 1, e1.unsqueeze(1)) masked = routing_weights.clone() for b in range(x_flat.shape[0]): grp = e1[b].item() % 3 for e in range(self.num_routed): if e % 3 == grp: masked[b, e] = -1.0 # Mask out all twin copies in the same chunk group e2 = torch.argmax(masked, dim=-1) p2 = torch.gather(routing_weights, 1, e2.unsqueeze(1)) topk_indices = torch.stack([e1, e2], dim=-1) topk_weights = torch.cat([p1, p2], dim=-1) topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True) else: topk_weights, topk_indices = torch.topk(routing_weights, self.top_k, dim=-1) topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True) routed_out = torch.zeros_like(x_flat) for k in range(self.top_k): expert_idx = topk_indices[:, k] weight = topk_weights[:, k].unsqueeze(-1) for e in range(self.num_routed): mask = (expert_idx == e) if mask.any(): tokens = x_flat[mask] h = self.act_fn(F.linear(tokens, self.routed_gate[e])) * F.linear(tokens, self.routed_up[e]) out = F.linear(h, self.routed_down[e]) routed_out[mask] += weight[mask] * out return (shared_out + routed_out).view(batch_size, seq_len, hidden_dim) class JarvisTitanMoEDecoderLayer(Qwen2DecoderLayer): def __init__(self, config: JarvisTitanMoEConfig, layer_idx: int): super().__init__(config, layer_idx) self.mlp = DeepSeekMoEMLP(config) class JarvisTitanMoEModel(Qwen2Model): config_class = JarvisTitanMoEConfig def __init__(self, config: JarvisTitanMoEConfig): super().__init__(config) self.layers = nn.ModuleList([ JarvisTitanMoEDecoderLayer(config, i) for i in range(config.num_hidden_layers) ]) self.post_init() class JarvisTitanMoEForCausalLM(Qwen2ForCausalLM): config_class = JarvisTitanMoEConfig def __init__(self, config: JarvisTitanMoEConfig): super().__init__(config) self.model = JarvisTitanMoEModel(config) self.post_init()