File size: 6,488 Bytes
f2878d0 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 | """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()
|