Jarvis-Titan-M4-Activated / modeling_jarvis_titan_moe.py
dhanesh-hf's picture
Sync modeling_jarvis_titan_moe.py from Phase 0 base
0e449ea verified
Raw History Blame Contribute Delete
6.35 kB
"""
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()