import json, sys, time import numpy as np, pandas as pd sys.path.insert(0, sys.argv[3]) from decision2_coreml import Decision2CoreML, encode as r_encode, question_to_row as r_row from transformers import AutoTokenizer import pack, q64ctx ref_root, refpath, rel = sys.argv[1], sys.argv[2], sys.argv[3] tok = AutoTokenizer.from_pretrained(ref_root) m = Decision2CoreML(rel) d = pd.read_parquet("data/typed/all/test-00000-of-00001.parquet") ref = {j["id"]: j["answers"] for j in map(json.loads, open(refpath))} def lab(a): return ("true" if a["noul"] >= .5 else "false") if a["type"] == "noul" else max(a["probabilities"], key=a["probabilities"].get) from collections import defaultdict acc = defaultdict(lambda: [0, 0, 0]); margins = [] tokdiff = flips = n = 0; worst = 0; ts = [] for r in d.itertuples(): st, qs = json.loads(r.state), json.loads(r.questions) for qid, q in qs.items(): # tokenization vs upstream encoder up = pack.rows(tok, st, {qid: q})[0][2]; mine = r_encode(r_row(st, q), m.tokenizer) tokdiff += up["ids"] != mine["ids"] or up["candidate_positions"] != mine["candidate_positions"] t = time.time(); a = m.system_one(state=st, questions=qs)["answers"]; ts.append((time.time() - t) * 1000) for q in qs: rq, aq = ref[r.id][q], a[q]; n += 1; flips += lab(rq) != lab(aq) pr = rq.get("probabilities", {"true": rq.get("noul")}); pa = aq.get("probabilities", {"true": aq.get("noul")}) g = json.loads(r.gold)[q]["label"]; s_ = acc[rq["type"]]; s_[0] += 1; s_[1] += lab(rq) == g; s_[2] += lab(aq) == g if lab(rq) != lab(aq): ps = sorted(pr.values(), reverse=True) if len(pr) > 1 else sorted([pr["true"], 1 - pr["true"]], reverse=True) margins.append(round(ps[0] - ps[1], 4)) worst = max(worst, max(abs(pr[k] - pa[k]) for k in pr)) for f in ("confidence", "score"): if f in rq: assert abs(rq[f] - aq[f]) < 0.05, (r.id, q, f) assert set(rq) == set(aq), (set(rq) ^ set(aq)) print("flip margins (upstream top-2)", margins) for k, (c, h1, h2) in sorted(acc.items()): print(f" {k:6s} n={c} acc upstream {h1/c:.3f} coreml {h2/c:.3f}") print(f"typed: token mismatches {tokdiff}, {n} decisions, flips {flips}, max|dp| {worst:.4f}, request p50 {np.median(ts):.1f} ms") t = time.time(); a = m.system_one(state=q64ctx.state, questions=q64ctx.Q); m.system_one(state=q64ctx.state, questions=q64ctx.Q) ts = [] for _ in range(5): t = time.time(); m.system_one(state=q64ctx.state, questions=q64ctx.Q); ts.append((time.time() - t) * 1000) print(f"64 questions: {len(m._calls([r_encode(r_row(q64ctx.state, q), m.tokenizer) for q in q64ctx.Q.values()]))} calls, p50 {np.median(ts):.0f} ms end to end")