TypeSafeAI
/

Qyvos / training /infer_qyvos.py
Manusagents's picture
Duplicate from SHSLab/Qyvos
6e78515
Raw History Blame Contribute Delete
4.17 kB
#!/usr/bin/env python3
"""Qyvos inference CLI (uses the official julia engine for release parity).
Examples:
python3 scripts/infer_qyvos.py --demo
python3 scripts/infer_qyvos.py --eval-test 1500
python3 scripts/infer_qyvos.py --row '{"state": "...", "question": "...",
"options": ["a", "b"], "type": "choice"}'
"""
import argparse
import json
import os
import sys
import time
from pathlib import Path
BASE = Path(os.environ.get("QYVOS_HOME", "/home/z/my-project/download/qyvos"))
sys.path.insert(0, str(BASE / "Julia-1"))
sys.path.insert(0, str(BASE / "scripts"))
MODEL_DIR = BASE / "Qyvos"
DATA = BASE / "data" / "data" / "release-v2-redistributable"
def load_engine():
import julia
# compat: disable julia's optional fast-path on transformers >=5.17
try:
import julia.router.encoder as _jre
_jre.specialize_decision_encoder = lambda model: False
except Exception:
pass
return julia.load_model(str(MODEL_DIR), device="cpu", max_length=1024, head_length=512)
def demo(engine, per_kind: int = 2) -> None:
import pyarrow.parquet as pq
pf = pq.ParquetFile(DATA / "test-00000-of-00001.parquet")
seen = {"choice": 0, "score": 0, "noul": 0}
for rg in range(pf.metadata.num_row_groups):
if all(v >= per_kind for v in seen.values()):
break
rows = pf.read_row_group(rg, columns=["kind", "question", "options", "target", "state_json"]).to_pylist()
for r in rows:
k = r["kind"]
if seen[k] >= per_kind:
continue
state = r["state_json"]
if isinstance(state, str):
state = json.loads(state)
req = {"state": state, "question": r["question"], "options": list(r["options"]), "type": k}
tgt = [float(x) for x in r["target"]]
pred = engine.predict([req])[0]
probs = [round(p, 4) for p in pred["probabilities"]]
gold = int(max(range(len(tgt)), key=lambda i: tgt[i]))
mark = "OK " if pred["index"] == gold else "MISS"
print(f"[{mark}] kind={k:6s} q={r['question'][:70]!r}")
print(f" options={list(r['options'])[:4]}")
print(f" pred index={pred['index']} probs={probs}")
print(f" gold index={gold} target={[round(t, 3) for t in tgt]}")
seen[k] += 1
if all(v >= per_kind for v in seen.values()):
break
def eval_test(n_rows: int) -> None:
import train_qyvos as T
from transformers import AutoTokenizer
from julia.model import JuliaDecisionModel
model = JuliaDecisionModel.from_pretrained(MODEL_DIR)
tok = AutoTokenizer.from_pretrained(MODEL_DIR / "tokenizer")
class Cfg:
eval_rows = n_rows
eval_batch = 8
max_length = 1024
head_length = 512
t0 = time.time()
acc, ce, detail = T.evaluate(model, tok, Cfg, split="test")
print(f"honest test eval (shuffled mixture, n={n_rows}): acc={acc:.4f} softCE={ce:.4f}")
for k, v in detail.items():
print(f" {k:7s} acc={v[0]:.4f} (n={v[1]})")
print(f"[{time.time()-t0:.0f}s]")
def main() -> int:
ap = argparse.ArgumentParser()
ap.add_argument("--demo", action="store_true")
ap.add_argument("--eval-test", type=int, default=0)
ap.add_argument("--row", type=str, default=None)
ap.add_argument("--per-kind", type=int, default=2)
args = ap.parse_args()
if not MODEL_DIR.exists():
print(f"model dir not found: {MODEL_DIR} (run build_qyvos.py first)", flush=True)
return 1
if args.eval_test:
eval_test(args.eval_test)
return 0
engine = load_engine()
if args.row:
req = json.loads(args.row)
pred = engine.predict([req])[0]
print(json.dumps({"index": pred["index"],
"probabilities": pred["probabilities"],
"selected": req["options"][pred["index"]]}, indent=2))
elif args.demo:
demo(engine, args.per_kind)
else:
ap.print_help()
return 0
if __name__ == "__main__":
sys.exit(main())