"""Fuse-2 model: Qwen3 host + DeepSeek V4 Flash coding experts. Architecture: per-layer expert augmentation (Option B from the master plan). At each augmented host layer, coding experts from DeepSeek V4 Flash are added alongside the host's native FFN. A learned router decides which experts fire. Key design principles (from fuse1 lessons): - bridge_out zero-init → model starts as exact Qwen3-4B - repair_up zero-init → no residual correction initially - Router initialized to low activation → coding path fires rarely at first - Frozen experts, frozen host → only bridges + routers + repair train - use_cache = False initially (correctness first) """ from __future__ import annotations import math import os from copy import deepcopy from typing import Iterator import torch import torch.nn as nn import torch.nn.functional as F from transformers import Qwen3Config, Qwen3ForCausalLM, Qwen3Model class FP8Linear(nn.Module): def __init__(self, input_features: int, output_features: int, bias: bool = False): super().__init__() self.register_buffer( "weight", torch.zeros(output_features, input_features, dtype=torch.float8_e4m3fn) ) self.register_buffer("scale", torch.ones((), dtype=torch.float32)) if bias: self.register_buffer("bias", torch.zeros(output_features)) else: self.bias = None @classmethod def from_linear(cls, linear: nn.Linear) -> "FP8Linear": result = cls(linear.in_features, linear.out_features, linear.bias is not None) with torch.no_grad(): weight = linear.weight.detach().float() max_value = weight.abs().amax() scale = max_value / 448.0 scale = torch.where(scale > 0, scale, torch.ones_like(scale)) result.weight.copy_((weight / scale).clamp(-448.0, 448.0).to(result.weight.dtype)) result.scale.copy_(scale) if linear.bias is not None: result.bias.copy_(linear.bias.detach()) return result def forward(self, x: torch.Tensor) -> torch.Tensor: weight = (self.weight.float() * self.scale).to(dtype=x.dtype) return F.linear(x, weight, self.bias) def convert_linears_to_fp8(module: nn.Module) -> None: for name, child in list(module.named_children()): if isinstance(child, FP8Linear): continue if isinstance(child, nn.Linear): setattr(module, name, FP8Linear.from_linear(child)) else: convert_linears_to_fp8(child) class Fuse2Config(Qwen3Config): """Qwen3 config extended with Fuse-2 MoE coding expert parameters.""" model_type = "fuse2" def __init__( self, # Expert configuration expert_hidden_size: int = 4096, # DeepSeek V4 hidden expert_intermediate_size: int = 2048, # DeepSeek V4 expert intermediate experts_per_layer: dict | None = None, # layer_idx -> list of expert IDs num_augmented_layers: int = 0, top_k_experts: int = 2, # Bridge configuration bridge_rank: int = 7, coding_enabled: bool = True, fp8_enabled: bool = False, # Router configuration router_init_scale: float = -2.0, # low initial activation load_balance_coef: float = 0.01, **kwargs, ): super().__init__(**kwargs) self.expert_hidden_size = expert_hidden_size self.expert_intermediate_size = expert_intermediate_size self.experts_per_layer = experts_per_layer or {} self.num_augmented_layers = num_augmented_layers self.top_k_experts = top_k_experts self.bridge_rank = bridge_rank self.coding_enabled = coding_enabled self.fp8_enabled = fp8_enabled self.router_init_scale = router_init_scale self.load_balance_coef = load_balance_coef self.use_cache = True self.architectures = ["Fuse2ForCausalLM"] self.auto_map = { **getattr(self, "auto_map", {}), "AutoModel": "fuse2_model.Fuse2Model", "AutoModelForCausalLM": "fuse2_model.Fuse2ForCausalLM", } class SwiGLUExpert(nn.Module): """A single DeepSeek V4 Flash expert (SwiGLU FFN). gate_proj: (intermediate, hidden) up_proj: (intermediate, hidden) down_proj: (hidden, intermediate) """ def __init__(self, hidden_size: int, intermediate_size: int): super().__init__() self.gate_proj = nn.Linear(hidden_size, intermediate_size, bias=False) self.up_proj = nn.Linear(hidden_size, intermediate_size, bias=False) self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=False) self._fuse2_scale: float | None = None def forward(self, x: torch.Tensor) -> torch.Tensor: if self._fuse2_scale is None: weight = self.gate_proj.weight if hasattr(weight, "dequantize"): weight = weight.dequantize() std_val = weight.float().std().item() self._fuse2_scale = 0.025 / std_val if std_val > 1.0 else 1.0 scale = self._fuse2_scale value = F.silu(self.gate_proj(x) * scale) * self.up_proj(x) * scale value = torch.clamp(value, -10.0, 10.0) return self.down_proj(value) * scale class Fuse2Router(nn.Module): """Per-layer router for coding experts. Uses sqrtsoftplus scoring (matching DeepSeek V4's approach) with top-k selection and optional load balancing. """ def __init__( self, input_dim: int, num_experts: int, top_k: int = 2, init_scale: float = -2.0, ): super().__init__() self.num_experts = num_experts self.top_k = min(top_k, num_experts) self.gate = nn.Linear(input_dim, num_experts, bias=False) # Initialize to low activation so coding path fires rarely at start nn.init.normal_(self.gate.weight, mean=0.0, std=0.01) self.init_scale = init_scale def forward( self, hidden_states: torch.Tensor, ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """Route tokens to experts. Args: hidden_states: (batch*seq, expert_hidden) — already bridged Returns: router_weights: (batch*seq, top_k) — softmax weights for selected experts expert_indices: (batch*seq, top_k) — which experts were selected router_logits: (batch*seq, num_experts) — raw logits for load balancing """ # sqrtsoftplus scoring (from DeepSeek V4) logits = self.gate(hidden_states) # (tokens, num_experts) scores = F.softplus(logits).sqrt() # Top-k selection topk_weights, topk_indices = scores.topk(self.top_k, dim=-1) # Normalize weights topk_weights = topk_weights / (topk_weights.sum(dim=-1, keepdim=True) + 1e-8) return topk_weights, topk_indices, logits class Fuse2AugmentedLayer(nn.Module): """One Qwen3 layer augmented with DeepSeek V4 coding experts. Forward flow: 1. Standard Qwen3 attention + FFN (frozen) 2. bridge_in: host_hidden → expert_hidden 3. router: select top-k coding experts 4. experts: parallel SwiGLU computation 5. bridge_out: expert_hidden → host_hidden (zero-init) 6. repair: rank-r residual correction (zero-init) 7. hidden += coding_delta + repair_delta """ def __init__( self, host_layer: nn.Module, host_hidden: int, expert_hidden: int, expert_intermediate: int, num_experts: int, top_k: int = 2, bridge_rank: int = 7, router_init_scale: float = -2.0, coding_enabled: bool = True, ): super().__init__() self.host_layer = host_layer self.coding_enabled = coding_enabled self.num_experts = num_experts self.top_k = top_k # Expose host layer attributes needed by the Qwen3 model forward pass self.attention_type = getattr(host_layer, "attention_type", "full_attention") # Bridge: host space ↔ expert space self.bridge_in = nn.Linear(host_hidden, expert_hidden, bias=False) self.bridge_out = nn.Linear(expert_hidden, host_hidden, bias=False) # Router self.router = Fuse2Router( expert_hidden, num_experts, top_k, router_init_scale ) # Experts (frozen, loaded from DeepSeek V4 Flash) self.experts = nn.ModuleList([ SwiGLUExpert(expert_hidden, expert_intermediate) for _ in range(num_experts) ]) # Residual repair (low-rank) self.repair_down = nn.Linear(host_hidden, bridge_rank, bias=False) self.repair_up = nn.Linear(bridge_rank, host_hidden, bias=False) # Initialize for preservation: zero-init bridge_out and repair_up nn.init.normal_(self.bridge_in.weight, mean=0.0, std=0.02) nn.init.zeros_(self.bridge_out.weight) nn.init.normal_(self.repair_down.weight, mean=0.0, std=0.02) nn.init.zeros_(self.repair_up.weight) if os.getenv("FUSE2_VLLM_COMPAT") == "1": self.host_layer.bridge_in = self.bridge_in self.host_layer.bridge_out = self.bridge_out self.host_layer.router = self.router self.host_layer.experts = self.experts self.host_layer.repair_down = self.repair_down self.host_layer.repair_up = self.repair_up def forward( self, hidden_states: torch.Tensor, attention_mask: torch.Tensor | None = None, position_ids: torch.LongTensor | None = None, past_key_values=None, use_cache: bool | None = False, position_embeddings=None, **kwargs, ) -> torch.Tensor: # 1. Run the host layer (attention + FFN) hidden_states = self.host_layer( hidden_states=hidden_states, attention_mask=attention_mask, position_ids=position_ids, past_key_values=past_key_values, use_cache=use_cache, position_embeddings=position_embeddings, **kwargs, ) if not self.coding_enabled or self.num_experts == 0: return hidden_states # 2. Bridge to expert space original_shape = hidden_states.shape h_flat = hidden_states.reshape(-1, original_shape[-1]) expert_input = self.bridge_in(h_flat) # (tokens, expert_hidden) # 3. Route to experts topk_weights, expert_indices, router_logits = self.router(expert_input) # 4. Compute expert outputs (sparse — only selected experts) batch_tokens = h_flat.shape[0] expert_output = torch.zeros_like(expert_input) for k in range(self.top_k): indices = expert_indices[:, k] # (tokens,) weights = topk_weights[:, k] # (tokens,) # Group tokens by expert for efficient computation for eid in range(self.num_experts): mask = indices == eid if not mask.any(): continue expert_in = expert_input[mask] expert_out = self.experts[eid](expert_in) expert_output[mask] += weights[mask].unsqueeze(-1) * expert_out # 5. Bridge back to host space coding_delta = self.bridge_out(expert_output) # 6. Repair repair_delta = self.repair_up(self.repair_down(h_flat)) # 7. Residual addition result = h_flat + coding_delta + repair_delta return result.reshape(original_shape) def get_router_logits(self) -> torch.Tensor | None: """Return last router logits for load balancing loss.""" return getattr(self, "_last_router_logits", None) class Fuse2Model(Qwen3Model): """Qwen3 decoder with Fuse-2 coding expert augmentation.""" config_class = Fuse2Config def __init__(self, config: Fuse2Config): super().__init__(config) experts_per_layer = config.experts_per_layer or {} augmented_count = 0 for layer_idx_str, expert_ids in experts_per_layer.items(): layer_idx = int(layer_idx_str) if layer_idx >= len(self.layers): raise ValueError( f"Layer {layer_idx} out of range " f"(model has {len(self.layers)} layers)" ) num_experts = len(expert_ids) if num_experts == 0: continue original_layer = self.layers[layer_idx] self.layers[layer_idx] = Fuse2AugmentedLayer( host_layer=original_layer, host_hidden=config.hidden_size, expert_hidden=config.expert_hidden_size, expert_intermediate=config.expert_intermediate_size, num_experts=num_experts, top_k=min(config.top_k_experts, num_experts), bridge_rank=config.bridge_rank, router_init_scale=config.router_init_scale, coding_enabled=config.coding_enabled, ) augmented_count += 1 config.num_augmented_layers = augmented_count class Fuse2ForCausalLM(Qwen3ForCausalLM): """Qwen3-4B host + DeepSeek V4 Flash coding experts.""" config_class = Fuse2Config _no_split_modules = ["Qwen3DecoderLayer", "Fuse2AugmentedLayer"] def __init__(self, config: Fuse2Config): super().__init__(config) self.model = Fuse2Model(config) if config.fp8_enabled: convert_linears_to_fp8(self) def set_coding_enabled(self, enabled: bool) -> None: """Toggle the coding expert path.""" for layer in self.model.layers: if isinstance(layer, Fuse2AugmentedLayer): layer.coding_enabled = enabled def get_augmented_layers(self) -> list[tuple[int, Fuse2AugmentedLayer]]: """Return (index, layer) pairs for all augmented layers.""" return [ (i, layer) for i, layer in enumerate(self.model.layers) if isinstance(layer, Fuse2AugmentedLayer) ] def get_trainable_params(self) -> dict[str, nn.Parameter]: """Return only the trainable parameters (bridges, routers, repair).""" trainable = {} for name, param in self.named_parameters(): if any( key in name for key in ("bridge_in", "bridge_out", "router", "repair_down", "repair_up") ): trainable[name] = param return trainable def freeze_host_and_experts(self) -> None: """Freeze everything except bridges, routers, and repair.""" for name, param in self.named_parameters(): if any( key in name for key in ("bridge_in", "bridge_out", "router", "repair_down", "repair_up") ): param.requires_grad = True else: param.requires_grad = False def count_parameters(self) -> dict[str, int]: """Count parameters by category.""" counts = { "host": 0, "experts": 0, "bridges": 0, "routers": 0, "repair": 0, "total": 0, "trainable": 0, } for name, param in self.named_parameters(): n = param.numel() counts["total"] += n if param.requires_grad: counts["trainable"] += n if "bridge_in" in name or "bridge_out" in name: counts["bridges"] += n elif "router" in name: counts["routers"] += n elif "repair" in name: counts["repair"] += n elif "experts" in name: counts["experts"] += n else: counts["host"] += n return counts def forward( self, input_ids: torch.LongTensor | None = None, attention_mask: torch.Tensor | None = None, position_ids: torch.LongTensor | None = None, past_key_values=None, inputs_embeds: torch.FloatTensor | None = None, labels: torch.LongTensor | None = None, use_cache: bool | None = None, **kwargs, ): if use_cache is None: use_cache = self.config.use_cache return super().forward( input_ids=input_ids, attention_mask=attention_mask, position_ids=position_ids, past_key_values=past_key_values, inputs_embeds=inputs_embeds, labels=labels, use_cache=use_cache, **kwargs, ) def load_expert_weights( model: Fuse2ForCausalLM, expert_dir: str, expert_mapping: dict[int, list[int]], ) -> dict: """Load extracted DeepSeek V4 expert weights into the Fuse2 model. Args: model: Fuse2 model with augmented layers expert_dir: directory containing expert safetensors expert_mapping: layer_idx -> list of expert IDs (matching selection order) Returns: Manifest of loaded tensors with hash verification """ from safetensors.torch import load_file import glob # Load all shards shard_files = sorted(glob.glob(f"{expert_dir}/experts-*.safetensors")) if not shard_files: raise FileNotFoundError(f"No expert shards found in {expert_dir}") all_tensors = {} for shard in shard_files: all_tensors.update(load_file(shard)) loaded = {} for layer_idx, expert_ids in expert_mapping.items(): augmented = model.model.layers[layer_idx] if not isinstance(augmented, Fuse2AugmentedLayer): raise ValueError(f"Layer {layer_idx} is not augmented") for local_idx, global_eid in enumerate(expert_ids): prefix = f"layer{layer_idx:02d}_expert{global_eid:03d}" for pname in ("gate_proj.weight", "up_proj.weight", "down_proj.weight"): key = f"{prefix}.{pname}" if key not in all_tensors: raise KeyError(f"Missing expert tensor: {key}") tensor = all_tensors[key] target_name = pname.replace(".", "_").replace("_weight", "") # Map to expert module parts = pname.split(".") module = augmented.experts[local_idx] for part in parts[:-1]: module = getattr(module, part) param = getattr(module, parts[-1]) param.data.copy_(tensor.to(param.dtype)) loaded[key] = { "shape": list(tensor.shape), "destination": f"layers.{layer_idx}.experts.{local_idx}.{pname}", } return loaded