llm_known_answer.py 3.6 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546
  1. #!/usr/bin/env python3
  2. # -*- coding: utf-8 -*-
  3. """离线模型已知答案回放 (T-LLM): 用固定问题集打运行中的工作台 /api/ask, 核 (1) 出文 (2) 期望契约条被引用 (3) 服务端校闸全过 (4) 耗时预算.
  4. 用法: python scripts/llm_known_answer.py [--cases <json>] [--n 6] [--budget 300] [--port 18033] [--model qwen3:8b]
  5. 问题集默认取 outputs/rudong/windscada/_demo/offline_qa_test_*_v25_*.json 的 runs[].{q, expected_claim}. 输出 logs/known_answer_<ts>.json + 一屏结果; exit 0 = 全过."""
  6. from __future__ import annotations
  7. import argparse, glob, json, sys, time, urllib.request
  8. from pathlib import Path
  9. ROOT = Path(__file__).resolve().parents[1]
  10. def ask(port, q, model, budget):
  11. req = urllib.request.Request(f"http://127.0.0.1:{port}/api/ask", data=json.dumps(dict(q=q, model=model)).encode(), headers={"Content-Type": "application/json"})
  12. with urllib.request.urlopen(req, timeout=30) as r: r.read()
  13. t0 = time.time()
  14. while time.time() - t0 < budget:
  15. time.sleep(2)
  16. with urllib.request.urlopen(f"http://127.0.0.1:{port}/api/ask_status", timeout=30) as r: s = json.load(r)
  17. if s.get("state") in ("done", "error"): return s, round(time.time() - t0, 1)
  18. return dict(state="timeout"), round(time.time() - t0, 1)
  19. def main():
  20. ap = argparse.ArgumentParser(); ap.add_argument("--cases"); ap.add_argument("--n", type=int, default=6); ap.add_argument("--budget", type=int, default=300); ap.add_argument("--port", type=int, default=18033); ap.add_argument("--model", default="qwen3:8b"); a = ap.parse_args()
  21. src = a.cases or sorted(glob.glob(str(ROOT / "outputs/rudong/windscada/_demo/offline_qa_test_*_v25_*.json")))[-1]
  22. runs = json.loads(Path(src).read_text(encoding="utf-8"))["runs"][: a.n]; out = []; allok = True
  23. for r in runs:
  24. s, secs = ask(a.port, r["q"], a.model, a.budget); ans = s.get("answer") or ""; v = s.get("verify") or {}
  25. cited = bool(r.get("expected_claim")) and (r["expected_claim"] in ans or any(r["expected_claim"] in str(x) for x in (v.get("明细") or [])))
  26. ok = s.get("state") == "done" and len(ans) > 40 and cited and bool(v.get("全通过", True)) and secs <= a.budget
  27. allok &= ok; out.append(dict(q=r["q"], expected=r.get("expected_claim"), state=s.get("state"), secs=secs, chars=len(ans), cited=cited, verify_all_pass=v.get("全通过"), escalated=s.get("refined_from"), err=(s.get("err") or "")[:120], ok=ok))
  28. print(f" [{'OK' if ok else 'FAIL'}] {secs:6.1f}s {len(ans):4d}字 引用{'✓' if cited else '✗'} 校闸{'✓' if v.get('全通过', True) else '✗'} {'升档→' + str(s.get('refined_from')) if s.get('refined_from') else ''} | {r['q'][:40]}")
  29. rep = dict(ts=time.strftime("%Y-%m-%dT%H:%M:%S"), source=src, model=a.model, port=a.port, n=len(out), n_ok=sum(1 for x in out if x["ok"]), budget_s=a.budget, runs=out)
  30. (ROOT / "logs").mkdir(exist_ok=True); p = ROOT / "logs" / f"known_answer_{time.strftime('%Y%m%d_%H%M%S')}.json"; p.write_text(json.dumps(rep, ensure_ascii=False, indent=1), encoding="utf-8")
  31. print(f"已知答案回放 {rep['n_ok']}/{rep['n']} 过 → {p.name}"); return 0 if allok else 1
  32. if __name__ == "__main__":
  33. # 控制台可能是 GBK(中文 Windows 代码页 936): 正文里的 ✔ ✗ ✅ ⚠ 这类字符编不出来会抛
  34. # UnicodeEncodeError, 脚本干成了事却以退出码 1 结束(同类坑见 src/console.py)。降级为 '?' 而不是崩;
  35. # 不用 import 是为了兼顾 python -m 与直接当脚本跑两种启动方式。
  36. import sys as _sys
  37. for _s in (_sys.stdout, _sys.stderr):
  38. try: _s.reconfigure(errors='replace')
  39. except Exception: pass
  40. sys.exit(main())