File size: 7,024 Bytes
697ef33 | 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 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 | #!/usr/bin/env python3
# SPDX-License-Identifier: Apache-2.0
"""Register Glm5vForConditionalGeneration as a DSA-capable architecture in SGLang.
WHY THIS PATCH IS NECESSARY (not an environment workaround):
SGLang gates the DeepSeek-Sparse-Attention path on a hardcoded architecture
tuple in ``is_deepseek_dsa`` (``srt/configs/model_config.py``). Upstream's own
DSA vision-language model (Mistral Large 3) is supported by adding its VLM
arch ``PixtralForConditionalGeneration`` to exactly this tuple in-tree — that
is the upstream mechanism; no env-var/out-of-tree hook exists for it (unlike
models/processors/mm-arch, which we register via the shipped
``SGLANG_EXTERNAL_*`` hooks). Our GLM5V wrapper needs the same one-line entry
so the DSA attention backend (and the ``server_args`` DSA backend defaults)
recognize the model whose text backbone is GlmMoeDsa.
Three insertions, each at a place where upstream registered its own DSA/MLA
VLM archs (KimiK25 / Pixtral precedents):
1. ``model_config.py``: the ``is_deepseek_dsa`` arch tuple — gates the DSA
attention backend;
2. ``model_config.py``: the MLA attention-arch condition — without it the
KV pool is built as MHA and the DSA backend crashes on
``dsa_kv_cache_store_fp8``;
3. ``server_args.py``: the DSA model-adjustments gate list — sets the DSA
dense-attn kv-len threshold (= model index_topk) and the prefill/decode
DSA backend defaults (trtllm on Blackwell); without it decode dies with
``Unsupported self.dsa_decode_impl = None``.
A 4th insertion (added later) enables MTP/NEXTN speculative decoding: without
it the draft worker rebuilds the entire ~90GB multimodal model as the "draft"
and OOMs (our wrapper arch is in none of SGLang's draft-model rewrite lists).
See the ``DRAFT_ANCHOR`` block below. This is only needed when serving with
``--speculative-algorithm NEXTN``; it is a no-op otherwise. (Measured: MTP gives
~1.1-1.16x tok/s at low/moderate batch + ~2.0 accept-len, greedy outputs are
token-identical to non-spec, but it is graph-capped at batch 256 so plain
high-batch TP8 wins for throughput-bound RL rollout — MTP is a single-stream
LATENCY lever, not a throughput one.)
Other in-tree arch lists mentioning GlmMoeDsa/KimiK25 are deliberately NOT
patched: they gate features we do not use (encoder disaggregation) or pure
performance defaults we set explicitly (flashinfer allreduce fusion,
attention-backend fallback).
The same entries will be needed wherever slime launches SGLang rollout engines
for glm5v (Milestone 1) — candidate for an upstream sglang PR.
Idempotent; run before ``sglang.launch_server`` (serve.sh does this).
"""
import pathlib
import subprocess
import sys
ARCH = "Glm5vForConditionalGeneration"
DSA_ANCHOR = '"GlmMoeDsaForCausalLM",'
DSA_INSERT = (
DSA_ANCHOR
+ f'\n "{ARCH}", # glm5v: GLM-5.2 (DSA) text + MoonViT vision (sglang_glm5v)'
)
MLA_ANCHOR = 'or "KimiK25ForConditionalGeneration" in self.hf_config.architectures'
MLA_INSERT = (
MLA_ANCHOR
+ f'\n or "{ARCH}" in self.hf_config.architectures # glm5v (sglang_glm5v)'
)
# server_args.py DSA adjustments gate: the list containing KimiK25 is unique
# to that block (other GlmMoeDsa lists in the file lack it).
SA_ANCHOR = (
'"KimiK25ForConditionalGeneration",\n'
' "MistralLarge3ForCausalLM",\n'
' "PixtralForConditionalGeneration",\n'
' "GlmMoeDsaForCausalLM",'
)
SA_INSERT = SA_ANCHOR + f'\n "{ARCH}", # glm5v (sglang_glm5v)'
# --- MTP / speculative NEXTN draft-model support (4th insertion) ---
# WHY: with --speculative-algorithm NEXTN, SGLang builds a SEPARATE draft model
# (is_draft_model=True) whose arch it rewrites to a single-layer NextN module in
# ``_config_draft_model``. That method maps the *text* backbone
# ``GlmMoeDsaForCausalLM`` -> ``DeepseekV3ForCausalLMNextN`` (correct: our MTP
# weights live at text layer 78: eh_proj/enorm/hnorm + one block). But our
# TOP-LEVEL arch is the multimodal wrapper ``Glm5vForConditionalGeneration``,
# which is in NONE of the rewrite lists -> the draft worker instead rebuilds the
# ENTIRE multimodal model (MoonViT + full 78-layer 90GB text model) as the
# "draft" and OOMs (a second ~90GB copy on top of base weights + KV + graphs).
#
# FIX: mirror the shipped ``Step3p7ForConditionalGeneration`` precedent (a VLM
# wrapper whose draft is its text MTP): when is_draft_model and arch is our
# wrapper, swap hf_config -> hf_text_config and set arch to GlmMoeDsaForCausalLM.
# Inserted at the TOP of _config_draft_model so the existing
# GlmMoeDsaForCausalLM -> DeepseekV3ForCausalLMNextN rule then fires on it.
DRAFT_ANCHOR = (
" def _config_draft_model(self):\n"
" is_draft_model = self.is_draft_model\n"
)
DRAFT_INSERT = (
DRAFT_ANCHOR
+ "\n"
+ " # glm5v (sglang_glm5v): map the multimodal wrapper draft to its\n"
+ " # text backbone so the GlmMoeDsa->NextN rewrite below applies.\n"
+ f' if is_draft_model and self.hf_config.architectures[0] == "{ARCH}":\n'
+ " self.hf_config = self.hf_text_config\n"
+ ' self.hf_config.architectures = ["GlmMoeDsaForCausalLM"]\n'
)
def find_module(mod: str) -> pathlib.Path:
out = subprocess.run(
[sys.executable, "-c", f"import {mod} as m; print(m.__file__)"],
capture_output=True,
text=True,
check=True,
)
return pathlib.Path(out.stdout.strip())
def apply(path: pathlib.Path, anchor: str, insert: str, marker: str, label: str) -> None:
src = path.read_text()
if marker in src:
print(f"{label}: already patched")
return
assert anchor in src, f"{label}: anchor not found in {path} (sglang version changed?)"
path.write_text(src.replace(anchor, insert, 1))
print(f"{label}: patched ({path})")
def main() -> None:
mc = find_module("sglang.srt.configs.model_config")
apply(mc, DSA_ANCHOR, DSA_INSERT, f'"{ARCH}", # glm5v', "DSA tuple")
apply(mc, MLA_ANCHOR, MLA_INSERT, f'or "{ARCH}" in self.hf_config.architectures', "MLA condition")
apply(
mc,
DRAFT_ANCHOR,
DRAFT_INSERT,
f'if is_draft_model and self.hf_config.architectures[0] == "{ARCH}"',
"NEXTN draft-model map",
)
sa = find_module("sglang.srt.server_args")
apply(sa, SA_ANCHOR, SA_INSERT, f'"{ARCH}", # glm5v', "server_args DSA gate")
# verify the DSA gate end-to-end
check = subprocess.run(
[
sys.executable,
"-c",
(
"from sglang.srt.configs.model_config import is_deepseek_dsa;"
f"print(is_deepseek_dsa({{'architectures': ['{ARCH}'], 'index_topk': 2048}}))"
),
],
capture_output=True,
text=True,
check=True,
)
assert check.stdout.strip() == "True", check.stdout + check.stderr
print("verified: is_deepseek_dsa recognizes the arch")
if __name__ == "__main__":
main()
|