File size: 4,224 Bytes
48c8658
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Score a trained model on a labelled JSONL file: accuracy, macro-F1, confusion and latency.

    python evaluate.py --model my-router --task task.example.json --data data/example.jsonl
    python evaluate.py --model my-router --onnx my-router/onnx/model-int8-blockwise.onnx ...

Accuracy counts rows whose annotators all agreed (a single gold label). Always test on data the
model never trained or validated on. The majority-label baseline is printed for comparison.
"""

from __future__ import annotations

import argparse
import json
import os
import time

os.environ.setdefault("USE_TF", "0")

import numpy as np

from common import answer_probs, load_rows, load_task, row_state


def load_model(path: str, onnx_path: str | None, device: str | None, max_tokens: int):
    if onnx_path:
        from laya.onnx_agent import ONNXAgent

        model = ONNXAgent(path, onnx_path=onnx_path)
    else:
        import laya

        model = laya.Agent(path, device=device)
    model.cfg["max_len"] = max_tokens
    return model


def main() -> None:
    ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
    ap.add_argument("--model", required=True, help="checkpoint dir or Hub repo id")
    ap.add_argument("--task", required=True)
    ap.add_argument("--data", required=True, help="labelled JSONL the model has not seen")
    ap.add_argument("--onnx", help="score this ONNX file (with the checkpoint's config/tokenizer) instead")
    ap.add_argument("--device", default=None)
    ap.add_argument("--max-tokens", type=int, default=512)
    ap.add_argument("--out", help="write per-row predictions to this JSONL")
    args = ap.parse_args()

    task = load_task(args.task)
    labels, questions = task["labels"], task["questions"]
    rows = [r for r in load_rows(args.data, labels) if r["gold"] is not None]
    model = load_model(args.model, args.onnx, args.device, args.max_tokens)

    preds = {qi: [] for qi in range(len(questions))}
    latencies = []
    out = open(args.out, "w", encoding="utf-8") if args.out else None
    for row in rows:
        record = {"id": row["id"], "gold": row["gold"], "predictions": {}}
        for qi, q in enumerate(questions):
            t = time.perf_counter()
            answer = model.system_one(row_state(row), {"q": q})["answers"]["q"]
            latencies.append((time.perf_counter() - t) * 1000)
            probs = answer_probs(answer, q, labels)
            preds[qi].append(int(np.argmax(probs)))
            record["predictions"][qi] = dict(zip(labels, [round(p, 4) for p in probs]))
        if out:
            out.write(json.dumps(record, ensure_ascii=False) + "\n")
    if out:
        out.close()

    golds = [labels.index(r["gold"]) for r in rows]
    majority = max(labels, key=lambda l: sum(r["gold"] == l for r in rows))
    print(f"{len(rows)} rows with a gold label; always answering {majority!r} scores "
          f"{sum(r['gold'] == majority for r in rows) / len(rows):.1%}")
    for qi, q in enumerate(questions):
        p = preds[qi]
        acc = np.mean([a == b for a, b in zip(p, golds)])
        f1s = []
        for k in range(len(labels)):
            tp = sum(a == k == b for a, b in zip(p, golds))
            fp = sum(a == k != b for a, b in zip(p, golds))
            fn = sum(b == k != a for a, b in zip(p, golds))
            f1s.append(2 * tp / (2 * tp + fp + fn) if tp else 0.0)
        confusion = [[sum(g == i and a == j for a, g in zip(p, golds)) for j in range(len(labels))]
                     for i in range(len(labels))]
        print(f"\nquestion {qi} ({q['type']}): {q['instructions'][:70]!r}")
        print(f"  accuracy {acc:.1%}   macro-F1 {np.mean(f1s):.3f}")
        print("  confusion (rows = gold, columns = predicted):")
        width = max(len(l) for l in labels)
        print("  " + " " * width + "  " + "  ".join(f"{l:>{width}}" for l in labels))
        for label, line in zip(labels, confusion):
            print(f"  {label:>{width}}  " + "  ".join(f"{n:>{width}}" for n in line))
    lat = np.array(latencies)
    print(f"\nlatency per decision: p50 {np.percentile(lat, 50):.0f} ms, p95 {np.percentile(lat, 95):.0f} ms")


if __name__ == "__main__":
    main()