File size: 2,176 Bytes
2c0400f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
#!/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())