#!/usr/bin/env python """Graft a Qwen MoE MTP head into an MXFP4 (compressed-tensors) trunk. Produces this repo: an MXFP4-quantized Ornith-1.0-35B with an MTP draft head transplanted in, so vLLM does lossless self-speculative decoding (`--speculative-config '{"method":"mtp","num_speculative_tokens":3}'`). Generalized over head shape — it copies EVERY ``mtp.*`` tensor from the donor and derives the unquantized-ignore list from the rank-2 (matmul) tensors, so the same script handles a small dense head (a handful of Linears) and a MoE head (the `fc`, attention projections, AND hundreds of per-expert MLP projections) with no per-model tensor list. It does two things: 1. Copies all ``mtp.*`` head tensors (BF16) from a donor checkpoint into a new ``model-mtp.safetensors`` shard and patches ``model.safetensors.index.json``. Big trunk shards are hard-linked (no multi-GB copy); small files are copied. 2. Adds each ``mtp.*`` Linear module to ``quantization_config.ignore`` so the compressed-tensors loader keeps the BF16 head as-is instead of expecting MXFP4 weight-scales. A tensor is treated as a Linear (and ignored) iff it is rank-2 and ends in ``.weight``; rank-1 norms are left out. This is the one mixed-precision gotcha — without it the model won't load. Inputs: --donor a checkpoint that SHIPS the mtp.* head (here: the Qwen MoE thinking distill, which keeps its trained MoE MTP head in BF16). --target an MXFP4 compressed-tensors trunk of Ornith-1.0-35B (lm_head/embed_tokens/vision BF16). --out output checkpoint dir. Credits: trunk — DeepReinforce (Ornith-1.0-35B, MIT); MTP head — nerkyor's Qwen3.6-35B-A3B-DSV4Pro-Thinking distill (Apache-2.0), a DeepSeek-V4-Pro-Thinking distill of Qwen3.6-35B-A3B; architecture — Qwen MoE. Serving — vLLM + compressed-tensors. RDNA4 base image — Rob Smith / tcclaviger. License of the combined work: Apache-2.0. Usage: python recipe_graft_mxfp4_35b.py \ --donor Capicua25x/Qwen3.6-35B-A3B-DSV4Pro-Thinking-Distill-MXFP4-Vision \ --target ./Ornith-1.0-35B-MXFP4 \ --out ./Ornith-1.0-35B-MXFP4-Vision-MTP """ import argparse, glob, json, os, shutil from safetensors import safe_open from safetensors.torch import save_file MTP_SHARD = "model-mtp.safetensors" def resolve(path_or_repo: str) -> str: """Local dir with weights -> itself; else download the HF snapshot.""" if os.path.isdir(path_or_repo) and glob.glob(os.path.join(path_or_repo, "*.safetensors")): return path_or_repo from huggingface_hub import snapshot_download return snapshot_download(path_or_repo) def main() -> int: ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) ap.add_argument("--donor", required=True, help="checkpoint that ships the mtp.* head") ap.add_argument("--target", required=True, help="MXFP4 compressed-tensors trunk to graft into") ap.add_argument("--out", required=True, help="output checkpoint dir") ap.add_argument("--force", action="store_true", help="overwrite --out if it exists") args = ap.parse_args() donor, target = resolve(args.donor), resolve(args.target) # 1. pull EVERY mtp.* tensor from the donor; record the Linear (rank-2) modules to ignore head, ignore = {}, [] for shard in glob.glob(os.path.join(donor, "*.safetensors")): with safe_open(shard, framework="pt") as f: for k in f.keys(): if not k.startswith("mtp."): continue t = f.get_tensor(k).contiguous() head[k] = t if k.endswith(".weight") and t.dim() == 2: # a Linear weight -> mark unquantized ignore.append(k[: -len(".weight")]) assert head, "donor ships no mtp.* tensors" ignore = sorted(set(ignore)) print(f"head: {len(head)} tensors, {len(ignore)} Linear modules to ignore " f"(all BF16: {all(str(t.dtype) == 'torch.bfloat16' for t in head.values())})") # 2. validate the target can host the head idx = json.load(open(os.path.join(target, "model.safetensors.index.json"))) wm = idx["weight_map"] assert "lm_head.weight" in wm, "target missing lm_head.weight" assert any(k.endswith("embed_tokens.weight") for k in wm), "target missing embed_tokens" assert not any(k.startswith("mtp.") for k in wm), "target already has an mtp head" # 3. build output: hard-link big shards, copy small files if os.path.exists(args.out): assert args.force, f"{args.out} exists (use --force)" shutil.rmtree(args.out) os.makedirs(args.out) for fn in os.listdir(target): src = os.path.realpath(os.path.join(target, fn)) if not os.path.isfile(src): continue dst = os.path.join(args.out, fn) if fn.endswith(".safetensors"): try: os.link(src, dst) except OSError: shutil.copy2(src, dst) else: shutil.copy2(src, dst) # 4. write the mtp shard + patch the index save_file(head, os.path.join(args.out, MTP_SHARD), metadata={"format": "pt"}) added = 0 for k, t in head.items(): wm[k] = MTP_SHARD added += t.numel() * t.element_size() idx.setdefault("metadata", {}) if "total_size" in idx["metadata"]: idx["metadata"]["total_size"] += added json.dump(idx, open(os.path.join(args.out, "model.safetensors.index.json"), "w"), indent=2) # 5. mark the head unquantized in the compressed-tensors config cfg = json.load(open(os.path.join(args.out, "config.json"))) ig = cfg["quantization_config"]["ignore"] for m in ignore: if m not in ig: ig.append(m) json.dump(cfg, open(os.path.join(args.out, "config.json"), "w"), indent=2) print(f"OK -> {args.out} (+{MTP_SHARD} {added/1e6:.1f} MB, +{len(ignore)} ignore entries)") print("serve: vllm serve --speculative-config '{\"method\":\"mtp\",\"num_speculative_tokens\":3}'") return 0 if __name__ == "__main__": raise SystemExit(main())