File size: 2,865 Bytes
61b6fb9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Per-item outputs for the audit: every test question with the full probability vector.

    PYTHONPATH=<decisions worktree> python scripts/decisions/eval_items.py --checkpoint CKPT --name v3 \
        --test typed_en=PATH --test tasksource=PATH ... --out DIR [--engram-off]

Same reader and prompt as train_decisions.py (decisions.Reader, prefill num_logits=1, bf16), one
question per prompt. Writes DIR/<name>[-engram_off]__<test>.jsonl with one line per question:
    {test, id, qid, type, grupo, keys, probs, pred, gold, soft}
`gold` is the list of accepted keys, `soft` the teacher distribution when the data has one.
"""
from __future__ import annotations

import argparse
import contextlib
import json
import sys
import time
from pathlib import Path

import torch

sys.path.insert(0, str(Path(__file__).resolve().parent))
from decisions import Reader, read_jsonl  # noqa: E402


@torch.no_grad()
def run(reader, cases, test, fh):
    n = 0
    for c in cases:
        lang = c.get("lang", "es")
        for qid, q in c["questions"].items():
            ids, keys = reader.prompt(c["state"], q, lang)
            p = torch.softmax(reader.logits_eval(ids, len(keys)), -1).tolist()
            soft = (c.get("suave") or {}).get(qid)
            fh.write(json.dumps({"test": test, "id": c["id"], "qid": qid, "type": q["type"], "grupo": c.get("grupo"),
                                 "keys": keys, "probs": p, "pred": keys[max(range(len(p)), key=p.__getitem__)],
                                 "gold": c["gold"][qid], "soft": soft, "prompt_tokens": len(ids)},
                                ensure_ascii=False) + "\n")
            n += 1
    return n


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--checkpoint", type=Path, required=True)
    ap.add_argument("--name", required=True)
    ap.add_argument("--test", action="append", default=[], help="name=path.jsonl")
    ap.add_argument("--out", type=Path, required=True)
    ap.add_argument("--engram-off", action="store_true")
    a = ap.parse_args()
    from mini_v41.inference import InferenceModel
    im = InferenceModel.from_checkpoint(a.checkpoint, device="cuda:0", dtype="bf16")
    im.model.eval()
    reader = Reader(im)
    ctx = contextlib.nullcontext()
    if a.engram_off:
        from mini_v41.engram import engram_disabled
        ctx = engram_disabled(im.model)
    a.out.mkdir(parents=True, exist_ok=True)
    tag = a.name + ("-engram_off" if a.engram_off else "")
    with ctx:
        for spec in a.test:
            test, path = spec.split("=", 1)
            t0 = time.time()
            with open(a.out / ("%s__%s.jsonl" % (tag, test)), "w", encoding="utf-8") as fh:
                n = run(reader, read_jsonl(path), test, fh)
            print("%s %s %d questions %.0fs" % (tag, test, n, time.time() - t0), flush=True)


if __name__ == "__main__":
    main()