File size: 18,079 Bytes
fdc6474
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
#!/usr/bin/env python3
"""Per-expert routing stats on the STOCK NVFP4 GLM-5.2, coding/agentic corpus.

Runs /tmp/glm52-dl/nvfp4-full (all experts NVFP4, near-lossless routing)
across 8 GPUs with a forced modular MoE backend so RoutedExperts sees
topk_ids (VLLM_EXPERT_STATS_DIR hook in routed_experts.py).

Corpus (200 prompts, up to ~7.9k tokens each):
  30% raw code continuation       (vLLM py / cu / cpp sources)
  12% bash / automation corpus    (real CI scripts, bash completions,
                                   Dockerfiles, init scripts)
  34% multi-turn tool-calling     (GLM chat template + tool schemas; tool
                                   results are real logs, diffs, file
                                   contents, tracebacks; coding + devops +
                                   desktop-automation tasks)
  10% coding chat                 (debug/review/explain, embedded code)
  11% medical                     (MedQA textbook continuation + clinical
                                   knowledge Q&A -- generic knowledge only)
  3%  general chat

Output: /data/glm52-expert-stats-v2.npz  (counts [n_layers, 256])
"""
import glob
import json
import os
import random
import subprocess

import numpy as np

MODEL = "/tmp/glm52-dl/nvfp4-full"
STATS_DIR = "/data/glm52-expert-stats-v2"
OUT = "/data/glm52-expert-stats-v2.npz"
N_EXPERTS = 256
MAX_TOKENS = 7900
ROOT = "/home/coder/git/glm52"

TOOLS = [
    {"type": "function", "function": {
        "name": "bash",
        "description": "Run a shell command in the repo and return stdout/stderr.",
        "parameters": {"type": "object", "properties": {
            "command": {"type": "string", "description": "Command to run"},
            "timeout": {"type": "integer", "description": "Seconds"}},
            "required": ["command"]}}},
    {"type": "function", "function": {
        "name": "read_file",
        "description": "Read a file, optionally a line range.",
        "parameters": {"type": "object", "properties": {
            "path": {"type": "string"},
            "start_line": {"type": "integer"},
            "end_line": {"type": "integer"}}, "required": ["path"]}}},
    {"type": "function", "function": {
        "name": "edit_file",
        "description": "Replace old_string with new_string in a file.",
        "parameters": {"type": "object", "properties": {
            "path": {"type": "string"},
            "old_string": {"type": "string"},
            "new_string": {"type": "string"}},
            "required": ["path", "old_string", "new_string"]}}},
    {"type": "function", "function": {
        "name": "python",
        "description": "Execute python code and return the output.",
        "parameters": {"type": "object", "properties": {
            "code": {"type": "string"}}, "required": ["code"]}}},
    {"type": "function", "function": {
        "name": "search_code",
        "description": "Regex search across the repository.",
        "parameters": {"type": "object", "properties": {
            "pattern": {"type": "string"},
            "glob": {"type": "string"}}, "required": ["pattern"]}}},
]


def _read(path, limit=20000):
    try:
        return open(path, errors="ignore").read()[:limit]
    except OSError:
        return ""


def _sh(cmd):
    try:
        return subprocess.run(
            cmd, shell=True, capture_output=True, text=True, timeout=20,
            cwd=ROOT).stdout[:20000]
    except Exception:
        return ""


def gather_artifacts(rng):
    """Real terminal/log/diff/file content to use as tool results."""
    art = {"logs": [], "code": [], "diffs": [], "listings": [], "json": []}
    for lf in glob.glob("/tmp/glm52-*.log") + glob.glob("/data/*.log"):
        txt = _read(lf, 400000)
        for i in range(0, len(txt) - 4000, 40000):
            art["logs"].append(txt[i:i + rng.randint(1500, 5000)])
    code_files = (glob.glob(f"{ROOT}/vllm/vllm/**/*.py", recursive=True)
                  + glob.glob(f"{ROOT}/vllm/csrc/**/*.cu", recursive=True)
                  + glob.glob(f"{ROOT}/vllm/csrc/**/*.cuh", recursive=True)
                  + glob.glob(f"{ROOT}/tools/*.py"))
    rng.shuffle(code_files)
    art["code_files"] = code_files
    for p in code_files[:400]:
        t = _read(p)
        if len(t) > 500:
            art["code"].append((p, t))
    d = _sh("git -C vllm diff --stat; git -C vllm diff | head -800")
    if d:
        for i in range(0, len(d), 4000):
            art["diffs"].append(d[i:i + 4000])
    art["listings"].append(_sh("ls -la vllm/vllm/model_executor/layers/quantization/ | head -40"))
    art["listings"].append(_sh("find vllm/csrc/quantization -type f | head -30"))
    art["json"].append(_read("/data/glm52-v3/config.json", 4000))
    art["json"].append(_read(f"{ROOT}/hybrid_plan.json", 4000))
    return art


USER_TASKS = [
    "The test suite is failing after my last change. Find the failing tests, "
    "diagnose the root cause, and fix it.",
    "Profile why model loading is slow and optimize the hot path.",
    "Add per-layer logging of expert routing entropy to the MoE layer, with "
    "an env var to enable it.",
    "There's a CUDA illegal memory access in the new kernel under batch>64. "
    "Find and fix it.",
    "Refactor the quantization config to support per-layer overrides without "
    "breaking existing checkpoints.",
    "Review the recent diff for correctness and potential race conditions.",
    "The server OOMs at long context. Figure out the memory budget and fix "
    "the allocation.",
    "Write a benchmark comparing the two MoE dequant paths and report results.",
    "Investigate why throughput dropped 20% between these two log files.",
    "Implement graceful fallback when the JIT kernel build fails.",
    # devops / desktop-automation / scripting tasks
    "Write a bash script that watches a downloads directory and sorts new "
    "files into folders by type, then install it as a systemd user service.",
    "Set up a cron job that backs up my projects directory to a remote host "
    "nightly with rsync, keeping 7 rotated snapshots and logging failures.",
    "My disk is nearly full. Find what's consuming space and clean it up "
    "safely, keeping anything modified in the last week.",
    "Automate my morning setup: a script that starts tmux with three named "
    "windows, activates the venv, tails the server log, and opens the todo "
    "file in the editor.",
    "Batch-rename all photos in this directory tree from IMG_xxxx to "
    "date-based names using their EXIF timestamps, handling collisions.",
    "Write a script to monitor GPU memory and temperature and alert me via "
    "desktop notification when either crosses a threshold.",
    "Migrate all my dotfiles into a git-managed directory with symlinks and "
    "an install script that works on a fresh machine.",
    "Bulk-convert these CSV exports to parquet, validate row counts match, "
    "and produce a summary report of schema differences.",
]

THINK_SNIPPETS = [
    "Let me start by looking at the failing area.",
    "I need to inspect the code first.",
    "Let me check the logs for the actual error.",
    "First I'll search for where this is defined.",
    "Let me look at the diff to understand what changed.",
    "Now let me verify the fix compiles and passes tests.",
    "The error suggests a shape mismatch; checking the tensor layout.",
    "That output confirms the hypothesis. Applying a fix.",
    "Tests pass now. Let me summarize the change.",
]


def make_agentic_session(rng, art, tok):
    msgs = [{"role": "system", "content":
             "You are a senior software engineer agent working in a large "
             "ML inference codebase. Use the available tools to investigate "
             "and make changes. Be precise and verify your work."}]
    msgs.append({"role": "user", "content": rng.choice(USER_TASKS)})
    n_calls = rng.randint(3, 7)
    cid = 0
    for _ in range(n_calls):
        cid += 1
        kind = rng.random()
        if kind < 0.35:
            cmd = rng.choice([
                "python -m pytest tests/kernels -x -q 2>&1 | tail -30",
                "grep -rn 'def apply' vllm/model_executor/layers/fused_moe/ | head",
                "git diff --stat", "nvidia-smi | head -20",
                "python tools/budget.py",
                "ninja -C build 2>&1 | tail -20",
            ])
            call = {"name": "bash", "arguments": {"command": cmd}}
            result = rng.choice(art["logs"] + art["diffs"] + art["listings"])
        elif kind < 0.6:
            p, t = rng.choice(art["code"])
            s = rng.randrange(0, max(1, len(t) - 4000))
            call = {"name": "read_file", "arguments": {
                "path": p.replace(ROOT + "/", ""),
                "start_line": s // 60, "end_line": s // 60 + 80}}
            result = t[s:s + rng.randint(1500, 4000)]
        elif kind < 0.75:
            call = {"name": "search_code", "arguments": {
                "pattern": rng.choice([
                    r"topk_ids", r"create_weights", r"kv_cache_dtype",
                    r"cudaMemcpy", r"register_parameter"]),
                "glob": "**/*.py"}}
            result = rng.choice(art["listings"] + art["diffs"])
        elif kind < 0.9:
            p, t = rng.choice(art["code"])
            s = rng.randrange(0, max(1, len(t) - 2000))
            old = t[s:s + rng.randint(80, 300)]
            call = {"name": "edit_file", "arguments": {
                "path": p.replace(ROOT + "/", ""), "old_string": old,
                "new_string": old + "  # reviewed"}}
            result = "OK: 1 replacement made."
        else:
            call = {"name": "python", "arguments": {"code":
                    "import json\ncfg=json.load(open('config.json'))\n"
                    "print(cfg['quantization_config']['quant_method'])"}}
            result = rng.choice(art["json"])
        msgs.append({"role": "assistant",
                     "content": rng.choice(THINK_SNIPPETS),
                     "tool_calls": [{"id": f"c{cid}", "type": "function",
                                     "function": call}]})
        msgs.append({"role": "tool", "tool_call_id": f"c{cid}",
                     "content": result})
    msgs.append({"role": "assistant", "content":
                 "Based on the investigation: the root cause is identified "
                 "and the fix is applied. Summary of changes follows."})
    return tok.apply_chat_template(msgs, tools=TOOLS, tokenize=False)


def make_coding_chat(rng, art, tok):
    p, t = rng.choice(art["code"])
    s = rng.randrange(0, max(1, len(t) - 8000))
    snippet = t[s:s + rng.randint(3000, 9000)]
    ask = rng.choice([
        "Review this code for bugs and suggest improvements:",
        "Explain what this code does and identify any performance issues:",
        "This code crashes intermittently. Find the race condition:",
        "Refactor this for readability and add type hints:",
        "Write unit tests covering the edge cases of this code:",
    ])
    msgs = [
        {"role": "system", "content": "You are an expert code reviewer."},
        {"role": "user", "content": f"{ask}\n\n```python\n{snippet}\n```"},
    ]
    return tok.apply_chat_template(
        msgs, tokenize=False, add_generation_prompt=True)


def _bash_corpus(rng):
    files = (glob.glob(f"{ROOT}/vllm/.buildkite/**/*.sh", recursive=True)
             + glob.glob("/usr/share/bash-completion/completions/*")
             + glob.glob("/etc/init.d/*")
             + glob.glob(f"{ROOT}/vllm/docker/Dockerfile*")
             + glob.glob(f"{ROOT}/vllm/**/*.sh", recursive=True)
             + glob.glob("/etc/profile.d/*.sh"))
    rng.shuffle(files)
    buf, out = "", []
    for p in files:
        t = _read(p)
        if len(t) < 200:
            continue
        buf += f"# ==== {os.path.basename(p)} ====\n{t}\n\n"
        while len(buf) >= MAX_TOKENS * 4:
            out.append(buf[: MAX_TOKENS * 4])
            buf = buf[MAX_TOKENS * 4:]
    if buf:
        out.append(buf)
    return out


def _medqa_chunks(rng, n):
    files = glob.glob("/tmp/medqa/**/*.jsonl", recursive=True)
    rng.shuffle(files)
    out, buf = [], ""
    for f in files:
        for line in open(f, errors="ignore"):
            try:
                buf += json.loads(line)["text"] + "\n\n"
            except (json.JSONDecodeError, KeyError):
                continue
            while len(buf) >= MAX_TOKENS * 4:
                out.append(buf[: MAX_TOKENS * 4])
                buf = buf[MAX_TOKENS * 4:]
                if len(out) >= n:
                    return out
    return out


def build_corpus(tok):
    rng = random.Random(42)
    art = gather_artifacts(rng)
    prompts = []

    # 30% raw code continuation
    buf = ""
    for p in art["code_files"]:
        buf += _read(p) + "\n\n"
        while len(buf) >= MAX_TOKENS * 4 and len(prompts) < 60:
            prompts.append(buf[: MAX_TOKENS * 4])
            buf = buf[MAX_TOKENS * 4:]
        if len(prompts) >= 60:
            break

    # 12% bash / automation scripts
    prompts += _bash_corpus(rng)[:24]

    # 34% agentic tool-calling sessions
    for _ in range(68):
        prompts.append(make_agentic_session(rng, art, tok))

    # 10% coding chat
    for _ in range(20):
        prompts.append(make_coding_chat(rng, art, tok))

    # medical textbook continuation (MedQA corpus, 7%)
    prompts += _medqa_chunks(rng, 14)

    # ~8% general knowledge, medical-weighted (generic clinical knowledge
    # only -- no patient-style data)
    medical = [
        "Explain the differential diagnosis approach for acute chest pain, "
        "covering cardiac, pulmonary, gastrointestinal, and musculoskeletal "
        "causes, and which investigations discriminate between them.",
        "Describe how to systematically read a chest X-ray, common findings "
        "(consolidation, pneumothorax, effusion, cardiomegaly, nodules), and "
        "typical pitfalls in interpretation.",
        "Explain the pharmacology of beta-blockers: mechanism, receptor "
        "selectivity, indications, contraindications, and interactions.",
        "Walk through the pathophysiology of type 2 diabetes from insulin "
        "resistance to complications, and the mechanism of each major drug "
        "class used to treat it.",
        "Explain how CT and MRI imaging work physically, when each is "
        "preferred clinically, and their contraindications.",
        "Describe the interpretation of a full blood count and common "
        "patterns: microcytic vs macrocytic anaemia, neutrophilia, "
        "lymphopenia, thrombocytopenia, and their differential causes.",
        "Explain the staging and grading of solid tumours, TNM notation, "
        "and how imaging and histopathology contribute to each.",
        "Describe the physiology of the cardiac cycle and how it maps to "
        "ECG waveforms, heart sounds, and common arrhythmia mechanisms.",
        "Explain sepsis: definitions, pathophysiology, early recognition "
        "criteria, and the evidence behind initial management bundles.",
        "Describe how radiology AI models are validated for clinical use: "
        "ground truth curation, reader studies, sensitivity/specificity "
        "tradeoffs, and regulatory considerations.",
        "Explain acid-base disturbances and how to interpret an arterial "
        "blood gas step by step, with compensated and mixed examples.",
        "Describe the mechanisms and comparative effectiveness of the major "
        "vaccine platforms: live attenuated, inactivated, subunit, mRNA, "
        "and viral vector.",
    ]
    general = ["the history of the transistor", "the economics of shipping",
               "how compilers optimize loops", "how DNS resolution works",
               "what causes inflation", "the water cycle"]
    for q in medical:
        msgs = [{"role": "user", "content": q + " Be thorough and detailed."}]
        prompts.append(tok.apply_chat_template(
            msgs, tokenize=False, add_generation_prompt=True))
    for t in general:
        msgs = [{"role": "user", "content":
                 f"Give me a long, detailed explanation of {t}, with "
                 "history, fundamentals, examples, and common misconceptions."}]
        prompts.append(tok.apply_chat_template(
            msgs, tokenize=False, add_generation_prompt=True))

    rng.shuffle(prompts)
    return prompts


def main():
    os.environ["VLLM_EXPERT_STATS_DIR"] = STATS_DIR
    from vllm import LLM, SamplingParams
    from vllm.inputs import TokensPrompt

    llm = LLM(
        model=MODEL,
        pipeline_parallel_size=8,
        gpu_memory_utilization=0.6,
        kv_cache_dtype="fp8_ds_mla",
        max_model_len=8192,
        max_num_seqs=4,
        enforce_eager=True,
        kernel_config={"moe_backend": "cutlass"},
    )
    tok = llm.get_tokenizer()
    texts = build_corpus(tok)
    print(f"{len(texts)} calibration prompts", flush=True)
    prompts = [
        TokensPrompt(prompt_token_ids=tok.encode(t)[:MAX_TOKENS])
        for t in texts
    ]
    sp = SamplingParams(max_tokens=1, temperature=0.0)
    outs = llm.generate(prompts, sp)
    total_tokens = sum(len(o.prompt_token_ids) for o in outs)
    del llm

    files = sorted(glob.glob(os.path.join(STATS_DIR, "layer_*.npy")))
    assert files, f"no stats dumped in {STATS_DIR}"
    layer_ids = sorted(int(os.path.basename(f)[6:-4]) for f in files)
    counts = np.stack(
        [np.load(os.path.join(STATS_DIR, f"layer_{li}.npy"))
         for li in layer_ids])
    np.savez(OUT, counts=counts, layer_ids=np.array(layer_ids),
             total_tokens=total_tokens)
    frac = counts / counts.sum(axis=1, keepdims=True).clip(min=1)
    sf = np.sort(frac, axis=1)[:, ::-1]
    print(f"tokens: {total_tokens}, layers: {len(layer_ids)}")
    for k in (16, 32, 64):
        m = sf[:, :k].sum(axis=1)
        print(f"top-{k} mass: mean {m.mean():.3f} min {m.min():.3f} "
              f"max {m.max():.3f}")
    print("saved", OUT)


if __name__ == "__main__":
    main()