Akahsizrr's picture
Publish compatibility source compatibility/fuse2_model.py
645c18f verified
Raw
History Blame Contribute Delete
18.9 kB
"""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