#!/usr/bin/env python3
"""Hardened llama.cpp speculative-decoding benchmark for DGX Spark (GB10).

All 27B configs share ONE binary; the only variable is the --spec-type flags.
N reps per (config, prompt); reports median + min/max; records output hashes
so greedy-divergence vs the no-spec control can be quantified.
"""
import json, subprocess, time, urllib.request, hashlib, statistics, os, signal, sys

SP    = os.environ.get("OUT_DIR", ".")
DIST  = os.environ.get("LLAMA_DIST", "./dist")
MODELS = os.environ.get("MODELS_DIR", "./models")
M     = MODELS
QWEN  = f"{M}/Qwen3.8-27B-GGUF/models--unsloth--Qwen3.8-27B-GGUF/snapshots/f1bfb127c64f7072bdd2cad55f258b9c8b2910fe/Qwen3.8-27B-UD-Q4_K_XL.gguf"
DS    = f"{M}/DeepSeek-V4-Flash-GGUF/UD-Q2_K_XL/DeepSeek-V4-Flash-UD-Q2_K_XL-00001-of-00003.gguf"
D2B   = f"{M}/models--empero-ai--Qwen3.8-2B-Distill-GGUF/snapshots/f4f73582d0b149595450c719b9a7521a03894f9c/Qwen3.8-2B-Q4_K_M.gguf"
DFL   = f"{M}/Qwen3.8-27B-DFlash2-Q8_0.gguf"
DSP   = f"{M}/Qwen3.8-27B-DSpark-Q8_0.gguf"
PORT  = 8090
REPS  = int(os.environ.get("REPS", "5"))

BASE27 = ["-m", QWEN, "--no-mmproj", "-ngl", "999", "-c", "32768", "--parallel", "1"]

CONFIGS = [
  ("A. baseline (no spec)",   BASE27),
  ("B. ngram-mod",            BASE27 + ["--spec-type", "ngram-mod"]),
  ("C. draft 2B (tuned)",     BASE27 + ["--spec-type", "draft-simple", "-md", D2B, "-ngld", "999",
                                        "--spec-draft-n-max", "32", "--spec-draft-n-min", "1",
                                        "--spec-draft-p-min", "0.75"]),
  ("D. DSpark",               BASE27 + ["--spec-type", "draft-dspark", "-md", DSP, "-ngld", "999"]),
  ("E. DFlash2",              BASE27 + ["--spec-type", "draft-dflash", "-md", DFL, "-ngld", "999"]),
  ("F. DeepSeek-V4-Flash MoE",["-m", DS, "-ngl", "999", "-c", "16384", "--parallel", "1", "-fa", "on"]),
]

_long_ctx = ("在磁共振成像中，弛豫时间是描述质子在射频脉冲激发后恢复到平衡态的特征时间常数。"
             "T1 反映纵向磁化恢复，T2 反映横向磁化衰减，二者由组织的分子环境决定。") * 45

PROMPTS = [
 ("zh-prose", "用中文写一段约300字的说明，介绍磁共振成像中T1加权像和T2加权像的物理原理差异。直接开始写。", 256),
 ("en-prose", "Explain in about 300 words why memory bandwidth, not FLOPS, limits token generation "
              "for dense transformer inference. Start directly.", 256),
 ("code-gen", "写一个 Python 函数，实现带路径压缩和按秩合并的并查集，包含完整 type hints 和 docstring。只输出代码。", 256),
 ("code-edit","重写下面的函数，加上 type hints、docstring 和输入校验，只输出代码：\n\n"
              "def process_scan(path, normalize=True, denoise=False):\n"
              "    img = load_nifti(path)\n    if normalize:\n"
              "        img = (img - img.mean()) / img.std()\n    if denoise:\n"
              "        img = gaussian_filter(img, sigma=1.0)\n    return img", 256),
 ("json-out", "输出一个 JSON 数组，包含 8 个对象，每个描述一种 MRI 序列，字段为 name/contrast/typical_TR_ms/"
              "typical_TE_ms/clinical_use。只输出 JSON。", 256),
 ("reasoning","一个水箱有两个进水管和一个出水管。A 管单独注满需 6 小时，B 管单独注满需 4 小时，"
              "出水管单独排空需 12 小时。三管齐开需要多久注满？请分步计算。", 256),
 ("translate","把下面这段话翻译成英文，保持技术术语准确：\n\n"
              "梯度回波序列通过施加反向梯度而非180度重聚脉冲来产生回波，因此对磁场不均匀性更敏感，"
              "其信号衰减遵循T2*而非T2。", 256),
 ("long-ctx", _long_ctx + "\n\n请用三句话总结上面这段话的核心概念。", 192),
]

def wait_ready(proc, timeout=900):
    t0 = time.time()
    while time.time() - t0 < timeout:
        if proc.poll() is not None: return False
        try:
            with urllib.request.urlopen(f"http://127.0.0.1:{PORT}/health", timeout=3) as r:
                if b'"ok"' in r.read(): return True
        except Exception: pass
        time.sleep(3)
    return False

def ask(prompt, n_predict):
    body = json.dumps({"messages":[{"role":"user","content":prompt}], "max_tokens":n_predict,
                       "temperature":0, "top_k":1, "stream":False, "cache_prompt":False}).encode()
    req = urllib.request.Request(f"http://127.0.0.1:{PORT}/v1/chat/completions", data=body,
                                 headers={"Content-Type":"application/json"})
    with urllib.request.urlopen(req, timeout=1800) as r: d = json.loads(r.read())
    t = d["timings"]
    return {"tg": t["predicted_per_second"], "pp": t["prompt_per_second"],
            "n_gen": t["predicted_n"], "n_prompt": t["prompt_n"], "pp_ms": t["prompt_ms"],
            "draft_n": t.get("draft_n"), "draft_acc": t.get("draft_n_accepted"),
            "hash": hashlib.sha256(d["choices"][0]["message"]["content"].encode()).hexdigest()[:16]}

results = {}
for name, args in CONFIGS:
    print(f"\n{'='*70}\n{name}\n{'='*70}", flush=True)
    env = dict(os.environ, LD_LIBRARY_PATH=f"{DIST}/lib:" + os.environ.get("LD_LIBRARY_PATH",""))
    proc = subprocess.Popen([f"{DIST}/bin/llama-server", *args, "--host","127.0.0.1","--port",str(PORT)],
                            stdout=open(f"{SP}/harness_srv.log","w"), stderr=subprocess.STDOUT,
                            env=env, preexec_fn=os.setsid)
    if not wait_ready(proc):
        print(f"  !! startup failed, skipping"); results[name] = {"error":"startup failed"}
        try: os.killpg(os.getpgid(proc.pid), signal.SIGKILL)
        except Exception: pass
        continue
    try:
        ask(PROMPTS[0][1], 32)   # warmup, not counted
        cfg = {}
        for pid, ptext, npred in PROMPTS:
            runs = [ask(ptext, npred) for _ in range(REPS)]
            tgs = sorted(r["tg"] for r in runs)
            cfg[pid] = {"tg_median": statistics.median(tgs), "tg_min": tgs[0], "tg_max": tgs[-1],
                        "tg_all": [round(x,2) for x in tgs],
                        "pp_median": statistics.median(r["pp"] for r in runs),
                        "pp_ms_median": statistics.median(r["pp_ms"] for r in runs),
                        "n_prompt": runs[0]["n_prompt"], "n_gen_median": statistics.median(r["n_gen"] for r in runs),
                        "hashes": sorted({r["hash"] for r in runs}),
                        "draft_n": runs[0]["draft_n"], "draft_acc": runs[0]["draft_acc"]}
            spread = (tgs[-1]-tgs[0])/statistics.median(tgs)*100
            print(f"  {pid:10s} tg={statistics.median(tgs):6.2f} t/s "
                  f"[{tgs[0]:5.2f}–{tgs[-1]:5.2f}, ±{spread:.1f}%]  "
                  f"pp={statistics.median(r['pp'] for r in runs):7.1f}  "
                  f"deterministic={'yes' if len(cfg[pid]['hashes'])==1 else 'no('+str(len(cfg[pid]['hashes']))+' variants)'}", flush=True)
        results[name] = cfg
    finally:
        try: os.killpg(os.getpgid(proc.pid), signal.SIGTERM)
        except Exception: pass
        proc.wait(timeout=120); time.sleep(5)

json.dump(results, open(f"{SP}/results.json","w"), ensure_ascii=False, indent=1)
print(f"\nSaved {SP}/results.json")
