| """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_hidden_size: int = 4096, |
| expert_intermediate_size: int = 2048, |
| experts_per_layer: dict | None = None, |
| num_augmented_layers: int = 0, |
| top_k_experts: int = 2, |
| |
| bridge_rank: int = 7, |
| coding_enabled: bool = True, |
| fp8_enabled: bool = False, |
| |
| router_init_scale: float = -2.0, |
| 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) |
| |
| 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 |
| """ |
| |
| logits = self.gate(hidden_states) |
| scores = F.softplus(logits).sqrt() |
|
|
| |
| topk_weights, topk_indices = scores.topk(self.top_k, dim=-1) |
|
|
| |
| 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 |
|
|
| |
| self.attention_type = getattr(host_layer, "attention_type", "full_attention") |
|
|
| |
| self.bridge_in = nn.Linear(host_hidden, expert_hidden, bias=False) |
| self.bridge_out = nn.Linear(expert_hidden, host_hidden, bias=False) |
|
|
| |
| self.router = Fuse2Router( |
| expert_hidden, num_experts, top_k, router_init_scale |
| ) |
|
|
| |
| self.experts = nn.ModuleList([ |
| SwiGLUExpert(expert_hidden, expert_intermediate) |
| for _ in range(num_experts) |
| ]) |
|
|
| |
| self.repair_down = nn.Linear(host_hidden, bridge_rank, bias=False) |
| self.repair_up = nn.Linear(bridge_rank, host_hidden, bias=False) |
|
|
| |
| 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: |
| |
| 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 |
|
|
| |
| original_shape = hidden_states.shape |
| h_flat = hidden_states.reshape(-1, original_shape[-1]) |
| expert_input = self.bridge_in(h_flat) |
|
|
| |
| topk_weights, expert_indices, router_logits = self.router(expert_input) |
|
|
| |
| batch_tokens = h_flat.shape[0] |
| expert_output = torch.zeros_like(expert_input) |
|
|
| for k in range(self.top_k): |
| indices = expert_indices[:, k] |
| weights = topk_weights[:, k] |
|
|
| |
| 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 |
|
|
| |
| coding_delta = self.bridge_out(expert_output) |
|
|
| |
| repair_delta = self.repair_up(self.repair_down(h_flat)) |
|
|
| |
| 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 |
|
|
| |
| 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", "") |
| |
| 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 |
|
|