Qwen3.5-4B-A3B-Student-v2 / research_code /convert_2b_to_4b_a3b.py
sepsy070716's picture
Release Qwen3.5-4B-A3B Student v2 with evaluation evidence
40bc205 verified
Raw History Blame Contribute Delete
9.39 kB
#!/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()