TinyDecide / python /tests /test_conformance.py
TheREZOR's picture
Python, Rust and ESP32 engines, shared conformance set, promo video
f2878d0 verified
Raw History Blame Contribute Delete
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()