File size: 4,832 Bytes
3599318
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
#!/usr/bin/env python3
"""Apply the Agnes patch to an sglang python package directory (the one that
contains `srt/`).  Idempotent.

  1. srt/configs/agnes.py                    new file
  2. srt/utils/hf_transformers/common.py     register AgnesConfig in _CONFIG_REGISTRY
  3. srt/models/qwen3_5.py                   load_weights: translate the Agnes checkpoint
                                              (tensor prefixes, parallel-FFN fold)

Usage: apply_patch.py <path/to/sglang>            e.g. .../site-packages/sglang
"""

import os
import shutil
import sys

HERE = os.path.dirname(os.path.abspath(__file__))
MARK = "# === agnes ==="

TRANSLATE = '''
# === agnes ===
# Agnes 3.0 Flash checkpoints (config model_type "agnes") use their own tensor
# prefixes and carry a parallel FFN branch per layer.  This generator sits at
# the top of the weight stream and turns it into what the implementation below
# expects: delta_attn -> linear_attn, global_attn -> self_attn, and the branch
# concatenated onto the main gate / up (dim 0) and down (dim 1) projections,
# matching the widened intermediate_size set by sglang.srt.configs.agnes.
import json as _agnes_json
import os as _agnes_os
import re as _agnes_re

_AGNES_MLP_RE = _agnes_re.compile(r"^(.*\\.layers\\.\\d+\\.mlp\\.)(gate_proj|up_proj|down_proj)\\.weight$")


class _AgnesBranchReader:
    def __init__(self, model_path):
        from safetensors import safe_open

        self._open = safe_open
        self.path = model_path
        index = _agnes_os.path.join(model_path, "model.safetensors.index.json")
        self.weight_map = _agnes_json.load(open(index))["weight_map"]
        self.handles = {}

    def get(self, key):
        fn = self.weight_map[key]
        if fn not in self.handles:
            self.handles[fn] = self._open(_agnes_os.path.join(self.path, fn), framework="pt", device="cpu")
        return self.handles[fn].get_tensor(key)


def _agnes_translate_weights(model, weights):
    cfg = getattr(model.config, "text_config", None) or model.config
    width = int(getattr(cfg, "agnes_parallel_ffn_intermediate_size", 0) or 0)
    if width <= 0:
        yield from weights
        return
    model_path = (
        getattr(cfg, "agnes_model_path", None)
        or getattr(model.config, "agnes_model_path", None)
        or _agnes_os.environ.get("AGNES_MODEL_PATH")
        or getattr(model.config, "_name_or_path", None)
    )
    if not model_path or not _agnes_os.path.isdir(model_path):
        raise RuntimeError(
            f"agnes: cannot locate the checkpoint directory (got {model_path!r}); "
            "set AGNES_MODEL_PATH to the model directory"
        )
    reader = _AgnesBranchReader(model_path)
    for name, w in weights:
        if ".mlp.parallel_ffn." in name:
            continue
        m = _AGNES_MLP_RE.match(name)
        if m and "visual" not in name and not name.startswith("mtp"):
            extra = reader.get(f"{m.group(1)}parallel_ffn.{m.group(2)}.weight")
            dim = 1 if m.group(2) == "down_proj" else 0
            w = torch.cat([w, extra.to(device=w.device, dtype=w.dtype)], dim=dim)
        name = name.replace(".delta_attn.", ".linear_attn.").replace(".global_attn.", ".self_attn.")
        yield name, w
# === /agnes ===

'''

REGISTER = '''
# === agnes ===
from sglang.srt.configs.agnes import AgnesConfig as _AgnesConfig

_CONFIG_REGISTRY[_AgnesConfig.model_type] = _AgnesConfig
'''


def patch_file(path, edit):
    src = open(path, encoding="utf-8").read()
    if MARK in src:
        return "already patched"
    out = edit(src)
    if out is None:
        raise SystemExit(f"anchor not found in {path}")
    open(path, "w", encoding="utf-8").write(out)
    return "patched"


def edit_model_file(src):
    anchor = "QWEN3_5_KV_SCALE_MAPPER = WeightsMapper("
    hook = "        weights = QWEN3_5_KV_SCALE_MAPPER.apply(weights)\n"
    if anchor not in src or src.count(hook) < 1:
        return None
    src = src.replace(anchor, TRANSLATE + anchor, 1)
    src = src.replace(hook, "        weights = _agnes_translate_weights(self, weights)\n" + hook)
    return src


def main():
    if len(sys.argv) != 2:
        sys.exit(__doc__)
    pkg = os.path.abspath(sys.argv[1])
    srt = os.path.join(pkg, "srt")
    if not os.path.isdir(srt):
        sys.exit(f"{pkg} does not contain srt/")
    dst = os.path.join(srt, "configs", "agnes.py")
    shutil.copy2(os.path.join(HERE, "agnes_sglang_config.py"), dst)
    print(f"configs/agnes.py: installed")
    print("utils/hf_transformers/common.py:", patch_file(
        os.path.join(srt, "utils", "hf_transformers", "common.py"), lambda s: s.rstrip("\n") + "\n" + REGISTER))
    print("models/qwen3_5.py:", patch_file(os.path.join(srt, "models", "qwen3_5.py"), edit_model_file))
    print("APPLY_PATCH_OK")


if __name__ == "__main__":
    main()