Ornith-1.0-35B-MXFP4-Vision-MTP / recipe_graft_mxfp4.py
Capicua25x's picture
Ornith-1.0-35B MXFP4 + grafted MoE MTP head (vision), for AMD RDNA4 / vLLM
71690bc verified
Raw History Blame Contribute Delete
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())