Download python/tests/test_conformance.py from TheREZOR/TinyDecide: direct link, hf CLI and curl.
- Browser
- Download file 6.49 kB
-
https://huggingface.co/TheREZOR/TinyDecide/resolve/main/python/tests/test_conformance.py
- Command line
-
hf download hf://TheREZOR/TinyDecide/python/tests/test_conformance.py
-
curl -L -o test_conformance.py https://huggingface.co/TheREZOR/TinyDecide/resolve/main/python/tests/test_conformance.py
6.49 kB
| """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() | |