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()