Download recipe_graft_mxfp4.py from Capicua25x/Ornith-1.0-35B-MXFP4-Vision-MTP: direct link, hf CLI and curl.
- Browser
- Download file 6.14 kB
-
https://huggingface.co/Capicua25x/Ornith-1.0-35B-MXFP4-Vision-MTP/resolve/main/recipe_graft_mxfp4.py
- Command line
-
hf download hf://Capicua25x/Ornith-1.0-35B-MXFP4-Vision-MTP/recipe_graft_mxfp4.py
-
curl -L -o recipe_graft_mxfp4.py https://huggingface.co/Capicua25x/Ornith-1.0-35B-MXFP4-Vision-MTP/resolve/main/recipe_graft_mxfp4.py
6.14 kB
| #!/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 <out> --speculative-config '{\"method\":\"mtp\",\"num_speculative_tokens\":3}'") | |
| return 0 | |
| if __name__ == "__main__": | |
| raise SystemExit(main()) | |