YuE2-3B-FP8 / patch_fast.py
DKmode22's picture
YuE2-3B AR NVFP4 (fp8-dynamic) — llm-compressor requant, 2026-09-13
2c0400f verified
Raw History Blame Contribute Delete
2.18 kB
#!/usr/bin/env python3
"""Idempotent patch for the installed yue2_infer `fast.py` (vLLM backend).
Adds ONE lever: if the env var YUE2_AR_CHECKPOINT names a directory, the vLLM
worker serves THAT checkpoint instead of the bf16 AR checkpoint it derives from
model.safetensors. That is how an NVFP4 (compressed-tensors) requant of the
derived Qwen3-shaped AR checkpoint is put on the real generation path. Nothing
else changes (dtype, KV sizing from config.json, logits processor, max_num_seqs=1).
Usage: patch_fast.py <site-packages>/yue2/fast.py [--check]
"""
from __future__ import annotations
import re
import sys
MARK = "# s5-patch: YUE2_AR_CHECKPOINT override (services/yue2-nvfp4/patch_fast.py)"
OLD = 'derived = derive_ar_checkpoint(setup["model_dir"])'
NEW = (
MARK + "\n"
' _override = os.environ.get("YUE2_AR_CHECKPOINT")\n'
' derived = Path(_override) if _override else derive_ar_checkpoint(setup["model_dir"])\n'
' if _override and not (derived / "config.json").exists():\n'
' raise FileNotFoundError(f"YUE2_AR_CHECKPOINT has no config.json: {derived}")\n'
' print(f"s5-patch: AR checkpoint = {derived} (override={bool(_override)})", file=sys.stderr, flush=True)'
)
def main() -> int:
path = sys.argv[1]
check = "--check" in sys.argv
src = open(path).read()
if MARK in src:
print("already patched")
return 0
if check:
print("NOT patched")
return 1
if src.count(OLD) != 1:
print(f"expected exactly one occurrence of {OLD!r}, found {src.count(OLD)}")
return 2
# the derive call is indented 4 spaces inside _worker_main
new_src = re.sub(r"^(\s+)" + re.escape(OLD) + r"$",
lambda m: m.group(1) + NEW.replace("\n ", "\n" + m.group(1)), src, count=1, flags=re.M)
if new_src == src:
print("substitution failed")
return 3
if "from pathlib import Path" not in new_src and "import Path" not in new_src:
print("fast.py has no Path import; refusing")
return 4
open(path, "w").write(new_src)
print("patched", path)
return 0
if __name__ == "__main__":
sys.exit(main())