"""Conformance: the Python engine against the JavaScript reference (sdk/conformance). cd sdk/python && python tests/test_conformance.py [S768_dir] [full_dir] [conformance_dir] Pass criteria, for both builds: tokenizer ids exact; per request ids, truncated, tokens and error messages exact; pick, span tok and span text exact; max |dp| and max |dqvec| <= 1e-3; make_protos() reproduces the reference protos (vec, cnt, center within 1e-5, lam exactly). """ import json import sys from pathlib import Path PKG = Path(__file__).resolve().parent.parent # sdk/python sys.path.insert(0, str(PKG)) # never the repo-root training package import tinydecide # noqa: E402 from tinydecide import TinyDecide, WordPiece, make_protos # noqa: E402 assert Path(tinydecide.__file__).resolve().parent == PKG / "tinydecide", tinydecide.__file__ TOL = 1e-3 args = sys.argv[1:] HF = (PKG / "../model.bin").exists() # python/ inside the Hugging Face repo BUILDS = {"S768": Path(args[0]) if len(args) > 0 else PKG / (".." if HF else "../../web_S768"), "full": Path(args[1]) if len(args) > 1 else PKG / ("../full" if HF else "../../web_fc")} CONF = Path(args[2]) if len(args) > 2 else PKG / "../conformance" class Stats: def __init__(self): self.answers = self.pick_bad = self.span_bad = self.type_bad = 0 self.dp = self.dq = self.dz = 0.0 self.notes = [] def note(self, s): if len(self.notes) < 8: self.notes.append(s) def diff(a, b): return max((abs(x - y) for x, y in zip(a, b)), default=0.0) if len(a) == len(b) else float("inf") def compare_answers(name, exp, got, st: Stats): if len(exp) != len(got): st.type_bad += 1 st.note(f"{name}: {len(got)} answers, expected {len(exp)}") return for i, (a, b) in enumerate(zip(exp, got)): st.answers += 1 if a["type"] != b["type"]: st.type_bad += 1 st.note(f"{name}[{i}]: type {b['type']} != {a['type']}") continue if a["type"] in ("choice", "score"): st.dp = max(st.dp, diff(a["probs"], b["probs"]), abs(a["confidence"] - b["confidence"])) if a["type"] == "score": st.dp = max(st.dp, abs(a["score"] - b["score"])) st.dz = max(st.dz, diff(a["z0"], b["z0"])) st.dq = max(st.dq, diff(a["qvec"], b["qvec"])) if a["pick"] != b["pick"]: st.pick_bad += 1 st.note(f"{name}[{i}]: pick {b['pick']} != {a['pick']} ({a['probs']} vs {b['probs']})") elif a["type"] == "noul": st.dp = max(st.dp, abs(a["p"] - b["p"])) st.dz = max(st.dz, diff(a["z0"], b["z0"])) st.dq = max(st.dq, diff(a["qvec"], b["qvec"])) else: st.dp = max(st.dp, abs(a["p_present"] - b["p_present"]), abs(a["p_span"] - b["p_span"])) if list(a["tok"]) != list(b["tok"]) or a["text"] != b["text"]: st.span_bad += 1 st.note(f"{name}[{i}]: span {b['tok']} {b['text']!r} != {a['tok']} {a['text']!r}") def run_build(build, model_dir, cases): exp = json.loads((CONF / f"expected.{build}.json").read_text(encoding="utf-8")) model = TinyDecide.load(model_dir) tok = WordPiece(model.meta["tokenizer"]) ok = True tok_bad = 0 for i, s in enumerate(cases["tokenizer"]): ids = tok.encode(s)[0] if ids != exp["tokenizer"][i]: tok_bad += 1 if tok_bad <= 5: print(f" tokenizer differs on {s!r}:\n js {exp['tokenizer'][i]}\n py {ids}") ok &= tok_bad == 0 st = Stats() ids_bad = err_bad = 0 ms, n_ms = 0.0, 0 for r, e in zip(cases["requests"], exp["requests"]): try: got = model.answer(r["state"], r["questions"]) except ValueError as ex: if e.get("error") != str(ex): err_bad += 1 st.note(f"{r['name']}: raised {ex!s} (expected {e.get('error')})") continue if "error" in e: err_bad += 1 st.note(f"{r['name']}: no error, expected {e['error']}") continue ms += got["ms"] n_ms += 1 if got["ids"] != e["ids"] or got["truncated"] != e["truncated"] or got["tokens"] != e["tokens"]: ids_bad += 1 st.note(f"{r['name']}: ids/truncated/tokens differ") continue compare_answers(r["name"], e["answers"], got["answers"], st) dvec = 0.0 lam_bad = cnt_bad = 0 for c, e in zip(cases["protos"], exp["protos"]): p = make_protos(c["question"]["type"], e["examples"], model.meta["beta"]) ref = e["protos"] dvec = max(dvec, diff(p["vec"].tolist(), ref["vec"]), diff(p["center"].tolist(), ref["center"] or p["center"].tolist())) if p["cnt"] != ref["cnt"]: cnt_bad += 1 if p["lam"] != ref["lam"]: lam_bad += 1 st.note(f"{c['name']}: lam {p['lam']} != {ref['lam']}") if c["no_center"]: p["center"] = None got = model.answer(c["state"], [c["question"]], [p]) if got["ids"] != e["ids"]: ids_bad += 1 st.note(f"{c['name']}: ids differ") continue compare_answers(c["name"], e["answers"], got["answers"], st) ok &= (ids_bad == 0 and err_bad == 0 and st.pick_bad == 0 and st.span_bad == 0 and st.type_bad == 0 and st.dp <= TOL and st.dq <= TOL and st.dz <= TOL and dvec <= 1e-5 and lam_bad == 0 and cnt_bad == 0) for s in st.notes: print(" " + s) print(f"{build}: tokenizer {len(cases['tokenizer']) - tok_bad}/{len(cases['tokenizer'])} exact | " f"requests {len(cases['requests'])} (id mismatches {ids_bad}, error mismatches {err_bad}) | " f"answers {st.answers}: pick mismatches {st.pick_bad}, span mismatches {st.span_bad}, " f"max|dp| {st.dp:.2e}, max|dqvec| {st.dq:.2e}, max|dz0| {st.dz:.2e} | " f"protos {len(cases['protos'])}: max|dvec| {dvec:.2e}, lam mismatches {lam_bad} | " f"{ms / max(n_ms, 1):.0f} ms/request | {'PASS' if ok else 'FAIL'}") return ok def main(): cases = json.loads((CONF / "cases.json").read_text(encoding="utf-8")) ok = True for build, d in BUILDS.items(): ok &= run_build(build, d, cases) sys.exit(0 if ok else 1) if __name__ == "__main__": main()