File size: 9,365 Bytes
8c867d9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
"""Benchmarks for MM-Jev (run in the Colab kernel; `jev`, `DATA` in globals).

Accuracy:  typed-decisions (Laya notebook scoring), BTZSC AG News / DAIR Emotion / Banking77 (jev-benchmarks pilot-v1
protocol), MASSIVE scenario en/ko + intent en, XNLI en, JevBench public (easy / standard / hard), and our multimodal
held-out sets. Latency: batch-1 CUDA-graph path, p50 / p95, end-to-end (towers included for raw media).
"""
import json, math, time
import numpy as np, torch
from mmjev import FastConfig, Seg, options_of, set_ffn_width, QTYPES

LOGFILE = "/content/eval.log"


def say(*a):
    with open(LOGFILE, "a") as f:
        print(time.strftime("%H:%M:%S"), *a, file=f)


def ece(conf, corr, bins=10):
    conf, corr = np.asarray(conf), np.asarray(corr)
    e, edges = 0.0, np.linspace(0, 1, bins + 1)
    for lo, hi in zip(edges[:-1], edges[1:]):
        m = (conf > lo) & (conf <= hi)
        if m.any():
            e += m.mean() * abs(conf[m].mean() - corr[m].mean())
    return float(e)


def _T(temps, t, k):
    from mmjev import k_bucket
    temps = temps or {}
    return temps.get(f"{t}:{k_bucket(k)}", temps.get(t, 1.0))


@torch.no_grad()
def predict(jev, recs, fc, width, temps=None, bs=8):
    """-> list (per record) of list (per question) of calibrated probability vectors."""
    set_ffn_width(jev.lm, width)
    jev.eval()
    out = []
    order = sorted(range(len(recs)), key=lambda i: sum(len(options_of(q)[0]) for q in recs[i]["qs"]))
    res = [None] * len(recs)
    def _go(idx):
        batch = [recs[i] for i in idx]
        try:
            plans = [jev.plan(jev.encode_state(r["state"], fc), r["qs"], fc) for r in batch]
            lg = jev.run(jev.pack(plans), fc)
        except torch.OutOfMemoryError:
            torch.cuda.empty_cache()
            if len(idx) == 1:
                raise
            h = len(idx) // 2
            _go(idx[:h]); _go(idx[h:]); return
        for i, r, per_q in zip(idx, batch, lg):
            res[i] = [torch.softmax(l.float() / _T(temps, q["type"], len(l)), -1).cpu().numpy()
                      for q, l in zip(r["qs"], per_q)]

    for s in range(0, len(order), bs):
        _go(order[s:s + bs])
    set_ffn_width(jev.lm, None)
    return res


def basic_metrics(recs, probs):
    acc, conf, corr, brier, nll = [], [], [], [], []
    for r, pq in zip(recs, probs):
        for q, p, y in zip(r["qs"], pq, r["ys"]):
            c = float(int(p.argmax()) == y)
            acc.append(c); conf.append(float(p.max())); corr.append(c)
            oh = np.zeros_like(p); oh[y] = 1
            brier.append(float(((p - oh) ** 2).sum())); nll.append(float(-math.log(max(p[y], 1e-12))))
    return dict(n=len(acc), acc=float(np.mean(acc)), brier=float(np.mean(brier)), nll=float(np.mean(nll)),
                ece=ece(conf, corr))


def typed_metrics(recs, probs):
    """Exactly the Laya fine-tuning notebook scoring (accuracy, soft acc, Brier, ECE, score MAE in levels)."""
    accs, soft, brier, maes, confs, corrs = [], [], [], [], [], []
    by_type, by_wf = {}, {}
    for r, pq in zip(recs, probs):
        for qid, q, p in zip(r["qnames"], r["qs"], pq):
            g, t = r["gold"][qid], q["type"]
            if t == "choice":
                keys = list(q["criteria"])
                gp = np.array([g["probabilities"].get(k, 1e-6) for k in keys]); gp /= gp.sum()
                c = float(keys[int(p.argmax())] == str(g["label"]))
                soft.append(float((p * gp).sum())); brier.append(float(((p - gp) ** 2).sum()))
                confs.append(float(p.max()))
            elif t == "noul":
                pv = float(p[1]); gv = float(g.get("noul", g.get("probabilities", {}).get("true", 0.5)))
                c = float(("true" if pv >= 0.5 else "false") == str(g["label"]).lower())
                pd, gd = np.array([1 - pv, pv]), np.array([1 - gv, gv])
                soft.append(float((pd * gd).sum())); brier.append(float(((pd - gd) ** 2).sum()))
                confs.append(max(pv, 1 - pv))
            else:
                ev = float((p * np.arange(len(p))).sum())
                maes.append(abs(ev - float(g.get("score", 0.0))))
                c = float(int(p.argmax()) == int(g.get("label", round(g.get("score", 0)))))
                confs.append(float(p.max()))
            accs.append(c); corrs.append(c)
            by_type.setdefault(t, []).append(c); by_wf.setdefault(r["workflow"], []).append(c)
    return dict(n=len(accs), acc=float(np.mean(accs)), soft_acc=float(np.mean(soft)), brier=float(np.mean(brier)),
                ece=ece(confs, corrs, 15), score_mae=float(np.mean(maes)),
                by_type={k: round(float(np.mean(v)), 3) for k, v in by_type.items()},
                by_workflow={k: round(float(np.mean(v)), 3) for k, v in by_wf.items()})


def accuracy_suite(jev, DATA, fc, width, temps=None, tag="", tasks=None):
    ev = [r for v in list(DATA.values()) for r in v if r["split"] == "eval"]
    tasks = tasks or sorted(set(r["task"] for r in ev))
    out = {}
    for t in tasks:
        rs = [r for r in ev if r["task"] == t]
        t0 = time.time()
        pr = predict(jev, rs, fc, width, temps)
        m = typed_metrics(rs, pr) if t == "typed_decisions" else basic_metrics(rs, pr)
        m["modality"] = rs[0]["modality"]
        out[t] = m
        say(f"[{tag}] {t:22s} n={m['n']:5d} acc {m['acc']:.3f} ece {m['ece']:.3f} ({time.time() - t0:.0f}s)")
    return out


# ------------------------------------------------------------------------------------------ latency
def _time(f, n):
    ts = []
    for _ in range(n):
        torch.cuda.synchronize(); t0 = time.perf_counter(); f(); torch.cuda.synchronize()
        ts.append((time.perf_counter() - t0) * 1000)
    return ts


@torch.no_grad()
def latency_suite(jev, DATA, fc, width, n_rep=3, max_items=40):
    """End-to-end decide() latency, batch 1, CUDA-graph path. Raw media -> towers are included."""
    set_ffn_width(jev.lm, width); jev._graphs = {}
    ev = [r for v in DATA.values() for r in v if r["split"] == "eval"]
    groups = {}
    for r in ev:
        if r.get("raw") is not None:
            groups.setdefault(f"{r['modality']}:{r['task']}", []).append((r["raw"], r["qs"]))
    for t in ("typed_decisions", "btzsc_agnews", "btzsc_banking77", "jevbench_hard"):
        rs = [r for r in ev if r["task"] == t][:max_items]
        groups[f"text:{t}"] = [(r["state"], r["qs"]) for r in rs]
    res = {}
    for g, items in groups.items():
        items = items[:max_items]
        for st, qs in items:                    # warm-up: capture every bucket graph first
            jev.decide(st, qs, fc=fc, graph=True)
        ts = []
        for st, qs in items:
            ts += _time(lambda: jev.decide(st, qs, fc=fc, graph=True), n_rep)
        nq = float(np.mean([len(qs) for _, qs in items]))
        res[g] = dict(n=len(items), questions_per_call=nq, p50_ms=float(np.percentile(ts, 50)),
                      p95_ms=float(np.percentile(ts, 95)), per_question_ms=float(np.percentile(ts, 50) / nq))
        say(f"  {g:32s} p50 {res[g]['p50_ms']:7.1f} ms  p95 {res[g]['p95_ms']:7.1f} ms  ({nq:.1f} q/call)")
    set_ffn_width(jev.lm, None)
    return res


@torch.no_grad()
def questions_scaling(jev, fc, width, text=None):
    """Laya-style: 1, 5, 10, 50 questions on one state in one call."""
    set_ffn_width(jev.lm, width)
    text = text or ("Hi, we were billed twice for March. Please refund the duplicate today or we will cancel our plan.")
    base_qs = [{"type": "choice", "instructions": "Which department should handle this?",
                "criteria": {"billing": "invoices, payments, refunds", "technical": "bugs, outages",
                             "other": "everything else"}},
               {"type": "score", "instructions": "How urgent is this?", "criteria": ["not urgent", "soon", "blocking"]},
               {"type": "noul", "instructions": "Does the user threaten to cancel or leave?"},
               {"type": "noul", "instructions": "Does the user explicitly request a refund?"},
               {"type": "noul", "instructions": "Is the user angry?"}]
    out = {}
    for n in (1, 5, 10, 50):
        qs = [dict(base_qs[i % 5], instructions=base_qs[i % 5]["instructions"] + (f" (#{i})" if i >= 5 else ""))
              for i in range(n)]
        jev.decide(text, qs, fc=fc, graph=True)
        ts = _time(lambda: jev.decide(text, qs, fc=fc, graph=True), 10)
        out[n] = float(np.median(ts))
        say(f"  {n:3d} questions/call: {out[n]:6.1f} ms ({out[n] / n:5.1f} ms/q)")
    set_ffn_width(jev.lm, None)
    return out


@torch.no_grad()
def tree_vs_naive(jev, DATA, fc, width, n=30):
    """openjev-style one pass per option vs the packed option tree (same model, eager path)."""
    set_ffn_width(jev.lm, width)
    rs = [r for r in DATA["typed"] if r["split"] == "eval"][:n]
    rows = {}
    for name, fn in (("tree (1 pass)", lambda p: jev.run(jev.pack([p]), fc)),
                     ("naive (1 pass / option)", lambda p: [jev.run(jev.pack([q]), fc) for q in jev.isolate(p)])):
        ts = []
        for r in rs:
            p = jev.plan(jev.encode_state(r["state"], fc), r["qs"], fc)
            fn(p)
            ts += _time(lambda: fn(p), 2)
        rows[name] = float(np.median(ts))
        say(f"  {name:26s} {rows[name]:7.1f} ms per typed-decisions case")
    set_ffn_width(jev.lm, None)
    return rows