#!/usr/bin/env python3 """Expand Qwen3.5-2B into a text-only 4B-total/3B-active sparse MoE. The initial checkpoint preserves the dense 2B MLP exactly: the shared path and the selected routed expert each contribute one half of the original output. Extra neurons start output-neutral but trainable. """ from __future__ import annotations import argparse import json import shutil from pathlib import Path import torch from safetensors import safe_open from safetensors.torch import save_file from transformers import Qwen3_5MoeForCausalLM, Qwen3_5MoeTextConfig NUM_EXPERTS = 2 EXPERTS_PER_TOKEN = 1 SHARED_WIDTH = 6912 ROUTED_WIDTH = 6784 SEED = 35 def source_tensor(source: Path, weight_map: dict[str, str], name: str) -> torch.Tensor: with safe_open(source / weight_map[name], framework="pt", device="cpu") as handle: return handle.get_tensor(name) def save_shard( output: Path, filename: str, tensors: dict[str, torch.Tensor], weight_map: dict[str, str], ) -> int: tensors = {name: tensor.contiguous() for name, tensor in tensors.items()} save_file(tensors, output / filename, metadata={"format": "pt"}) weight_map.update({name: filename for name in tensors}) return sum(tensor.numel() * tensor.element_size() for tensor in tensors.values()) def build_config(source_config: dict) -> dict: text = dict(source_config["text_config"]) dense_width = int(text.pop("intermediate_size")) if text["hidden_size"] != 2048 or text["num_hidden_layers"] != 24 or dense_width != 6144: raise ValueError("converter is intentionally pinned to Qwen3.5-2B") text.update({ "architectures": ["Qwen3_5MoeForCausalLM"], "model_type": "qwen3_5_moe_text", "num_experts": NUM_EXPERTS, "num_experts_per_tok": EXPERTS_PER_TOKEN, "moe_intermediate_size": ROUTED_WIDTH, "shared_expert_intermediate_size": SHARED_WIDTH, "router_aux_loss_coef": 1e-3, "output_router_logits": False, }) return text def main() -> None: parser = argparse.ArgumentParser() parser.add_argument("--source", type=Path, required=True) parser.add_argument("--output", type=Path, required=True) args = parser.parse_args() if args.output.exists() and any(args.output.iterdir()): raise SystemExit(f"Refusing non-empty output directory: {args.output}") args.output.mkdir(parents=True, exist_ok=True) source_config = json.loads((args.source / "config.json").read_text()) config_dict = build_config(source_config) config = Qwen3_5MoeTextConfig(**config_dict) with torch.device("meta"): target = Qwen3_5MoeForCausalLM(config) target_keys = set(target.state_dict()) total_parameters = sum(parameter.numel() for parameter in target.parameters()) inactive = ( config.num_hidden_layers * (NUM_EXPERTS - EXPERTS_PER_TOKEN) * 3 * config.hidden_size * ROUTED_WIDTH ) active_parameters = total_parameters - inactive source_index = json.loads((args.source / "model.safetensors.index.json").read_text()) source_map: dict[str, str] = source_index["weight_map"] output_map: dict[str, str] = {} total_bytes = 0 # Copy the text backbone while dropping the vision tower and dense MLPs. common: dict[str, torch.Tensor] = {} prefix = "model.language_model." for shard in sorted(set(source_map.values())): with safe_open(args.source / shard, framework="pt", device="cpu") as handle: for source_name in handle.keys(): if not source_name.startswith(prefix) or ".mlp." in source_name: continue target_name = "model." + source_name[len(prefix):] if target_name in target_keys: common[target_name] = handle.get_tensor(source_name) total_bytes += save_shard(args.output, "model-common.safetensors", common, output_map) generator = torch.Generator(device="cpu").manual_seed(SEED) dense_width = 6144 hidden = config.hidden_size init_std = float(config.initializer_range) for layer in range(config.num_hidden_layers): src = f"model.language_model.layers.{layer}.mlp" dst = f"model.layers.{layer}.mlp" dense_gate = source_tensor(args.source, source_map, f"{src}.gate_proj.weight") dense_up = source_tensor(args.source, source_map, f"{src}.up_proj.weight") dense_down = source_tensor(args.source, source_map, f"{src}.down_proj.weight") dtype = dense_gate.dtype shared_gate = torch.randn(SHARED_WIDTH, hidden, generator=generator, dtype=torch.float32) shared_up = torch.randn(SHARED_WIDTH, hidden, generator=generator, dtype=torch.float32) shared_down = torch.zeros(hidden, SHARED_WIDTH, dtype=dtype) shared_gate.mul_(init_std).to(dtype=dtype) shared_up.mul_(init_std).to(dtype=dtype) shared_gate = shared_gate.to(dtype) shared_up = shared_up.to(dtype) shared_gate[:dense_width] = dense_gate shared_up[:dense_width] = dense_up # sigmoid(shared_expert_gate=0) gives the shared path a 0.5 multiplier. shared_down[:, :dense_width] = dense_down routed_gate_up = torch.empty( NUM_EXPERTS, 2 * ROUTED_WIDTH, hidden, dtype=dtype ) routed_down = torch.zeros(NUM_EXPERTS, hidden, ROUTED_WIDTH, dtype=dtype) for expert in range(NUM_EXPERTS): extra_gate = torch.randn( ROUTED_WIDTH, hidden, generator=generator, dtype=torch.float32 ).mul_(init_std).to(dtype) extra_up = torch.randn( ROUTED_WIDTH, hidden, generator=generator, dtype=torch.float32 ).mul_(init_std).to(dtype) extra_gate[:dense_width] = dense_gate extra_up[:dense_width] = dense_up routed_gate_up[expert, :ROUTED_WIDTH] = extra_gate routed_gate_up[expert, ROUTED_WIDTH:] = extra_up # Routed path supplies the other half of the original dense output. routed_down[expert, :, :dense_width] = dense_down * 0.5 tensors = { f"{dst}.gate.weight": torch.zeros(NUM_EXPERTS, hidden, dtype=dtype), f"{dst}.experts.gate_up_proj": routed_gate_up, f"{dst}.experts.down_proj": routed_down, f"{dst}.shared_expert.gate_proj.weight": shared_gate, f"{dst}.shared_expert.up_proj.weight": shared_up, f"{dst}.shared_expert.down_proj.weight": shared_down, f"{dst}.shared_expert_gate.weight": torch.zeros(1, hidden, dtype=dtype), } total_bytes += save_shard( args.output, f"model-layer-{layer:02d}.safetensors", tensors, output_map ) print(f"converted layer {layer + 1}/{config.num_hidden_layers}", flush=True) missing = sorted(target_keys - set(output_map) - {"lm_head.weight"}) unexpected = sorted(set(output_map) - target_keys) if missing or unexpected: raise RuntimeError(f"key audit failed: missing={missing[:20]} unexpected={unexpected[:20]}") (args.output / "model.safetensors.index.json").write_text(json.dumps({ "metadata": {"total_size": total_bytes}, "weight_map": dict(sorted(output_map.items())), }, indent=2)) (args.output / "config.json").write_text(config.to_json_string()) (args.output / "conversion_manifest.json").write_text(json.dumps({ "source": str(args.source), "initial_function": "Qwen3.5-2B text model (dense MLP split 50/50)", "total_parameters": total_parameters, "active_parameters": active_parameters, "num_experts": NUM_EXPERTS, "experts_per_token": EXPERTS_PER_TOKEN, "shared_intermediate_size": SHARED_WIDTH, "routed_intermediate_size": ROUTED_WIDTH, "vision_included": False, "seed": SEED, }, indent=2)) for filename in ("chat_template.jinja", "merges.txt", "tokenizer.json", "tokenizer_config.json", "vocab.json", "LICENSE"): source_file = args.source / filename if source_file.exists(): shutil.copy2(source_file, args.output / filename) (args.output / "README.md").write_text(f"""--- license: apache-2.0 base_model: - Qwen/Qwen3.5-2B - Qwen/Qwen3.5-4B library_name: transformers pipeline_tag: text-generation tags: [qwen3_5_moe, moe, upcycled, research] --- # Qwen3.5-4B-A3B-Student-v2 Text-only sparse-MoE research checkpoint initialized to preserve the text generation function of Qwen3.5-2B. It is intended for distillation from Qwen3.5-4B and is not yet claimed to match the 4B teacher. | Property | Value | |---|---:| | Total parameters | {total_parameters:,} | | Active parameters/token | {active_parameters:,} | | Experts / selected | {NUM_EXPERTS} / {EXPERTS_PER_TOKEN} | | Shared / routed width | {SHARED_WIDTH} / {ROUTED_WIDTH} | | Vision | No | The shared path and selected routed path initially contribute half of the dense Qwen3.5-2B MLP each. Extra neurons are output-neutral at initialization but can learn during distillation. See `conversion_manifest.json` for exact metadata. """) print(json.dumps({ "total_parameters": total_parameters, "active_parameters": active_parameters, "total_bytes": total_bytes, }, indent=2)) if __name__ == "__main__": main()