#!/usr/bin/env python3
"""Follow-up experiments 1&3: ngram cold/warm cache separation + greedy divergence rate per method vs no-spec control"""
import json, subprocess, time, urllib.request, hashlib, os, signal, statistics

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"
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
BASE=["-m",QWEN,"--no-mmproj","-ngl","999","-c","32768","--parallel","1"]
CFG={
 "baseline":  BASE,
 "ngram-mod": BASE+["--spec-type","ngram-mod"],
 "draft-2B":  BASE+["--spec-type","draft-simple","-md",D2B,"-ngld","999",
                    "--spec-draft-n-max","32","--spec-draft-n-min","1","--spec-draft-p-min","0.75"],
 "DSpark":    BASE+["--spec-type","draft-dspark","-md",DSP,"-ngld","999"],
 "DFlash2":   BASE+["--spec-type","draft-dflash","-md",DFL,"-ngld","999"],
}
_lc=("在磁共振成像中，弛豫时间是描述质子在射频脉冲激发后恢复到平衡态的特征时间常数。"
     "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\ndef 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",_lc+"\n\n请用三句话总结上面这段话的核心概念。",192),
]
def start(args):
    env=dict(os.environ, LD_LIBRARY_PATH=f"{DIST}/lib:"+os.environ.get("LD_LIBRARY_PATH",""))
    p=subprocess.Popen([f"{DIST}/bin/llama-server",*args,"--host","127.0.0.1","--port",str(PORT)],
        stdout=open(f"{SP}/fu_srv.log","w"),stderr=subprocess.STDOUT,env=env,preexec_fn=os.setsid)
    t0=time.time()
    while time.time()-t0<900:
        if p.poll() is not None: return None
        try:
            with urllib.request.urlopen(f"http://127.0.0.1:{PORT}/health",timeout=3) as r:
                if b'"ok"' in r.read(): return p
        except Exception: pass
        time.sleep(3)
    return None
def stop(p):
    try: os.killpg(os.getpgid(p.pid),signal.SIGTERM); p.wait(timeout=120)
    except Exception: pass
    time.sleep(4)
def ask(prompt,n):
    body=json.dumps({"messages":[{"role":"user","content":prompt}],"max_tokens":n,
        "temperature":0,"top_k":1,"stream":False,"cache_prompt":False}).encode()
    r=urllib.request.Request(f"http://127.0.0.1:{PORT}/v1/chat/completions",data=body,
        headers={"Content-Type":"application/json"})
    with urllib.request.urlopen(r,timeout=1800) as f: d=json.loads(f.read())
    c=d["choices"][0]["message"]["content"]
    return d["timings"]["predicted_per_second"], hashlib.sha256(c.encode()).hexdigest()[:16]

# ---------- Experiment 1: ngram cold vs warm ----------
print("="*74); print("Exp1  ngram-mod cache contamination: restart server each run (cold) vs same process (warm)"); print("="*74)
cold={}; warm={}
for pid,ptext,npred in PROMPTS:
    cs=[]
    for _ in range(3):                       # fresh process each time => cold cache
        p=start(CFG["ngram-mod"])
        if not p: break
        cs.append(ask(ptext,npred)[0]); stop(p)
    p=start(CFG["ngram-mod"])                # same process, 5 runs => cumulative cache
    ws=[ask(ptext,npred)[0] for _ in range(5)] if p else []
    if p: stop(p)
    cold[pid]=statistics.median(cs) if cs else None
    warm[pid]=ws
    print(f"  {pid:10s} cold={cold[pid]:6.2f} t/s   warm={[round(x,1) for x in ws]}", flush=True)

# ---------- Experiment 3: greedy divergence rate ----------
print(); print("="*74); print("Exp3  Greedy divergence rate (temperature=0, top_k=1) vs no-spec control"); print("="*74)
p=start(CFG["baseline"]); ref={}
for pid,ptext,npred in PROMPTS: ref[pid]=ask(ptext,npred)[1]
stop(p)
div={}
for name in ["ngram-mod","draft-2B","DSpark","DFlash2"]:
    p=start(CFG[name])
    if not p: continue
    same=0; detail=[]
    for pid,ptext,npred in PROMPTS:
        h=ask(ptext,npred)[1]; ok=(h==ref[pid]); same+=ok
        detail.append((pid,ok))
    stop(p)
    div[name]={"same":same,"total":len(PROMPTS),"detail":detail}
    bad=[d[0] for d in detail if not d[1]]
    print(f"  {name:10s} matches control {same}/{len(PROMPTS)}"
          + (f"   diverged: {', '.join(bad)}" if bad else "   all identical"), flush=True)

json.dump({"ngram_cold":cold,"ngram_warm":warm,"divergence":div},
          open(f"{SP}/followup.json","w"),ensure_ascii=False,indent=1)
print(f"\nSaved {SP}/followup.json")
