File size: 4,463 Bytes
4397e12
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Agent eval on held-out workspace tasks (seeds 0..99_999; RL trains on 100_000..999_999 and
pretraining data uses >= 1_000_000). Reports per-kind pass@1 (sampled), pass@k, grounding,
turns, parse errors and generated tokens. This is the "pretrained enough for RL?" gate.

  source env.sh && $TA_PY scripts/eval_agent.py --ckpt $TA_DATA/runs/m1/final.pt --tasks 90 --k 8
"""
import argparse
import json
import random
from collections import defaultdict
from math import comb

import torch
from tokenizers import Tokenizer

from tiny_agent.checkpoint import load_model
from tiny_agent.rollout import make_roller
from tiny_agent.tasks import HELDOUT_KINDS, KINDS, make_task
from tiny_agent.text import DATA

EVAL_SEED0 = 0


def pass_at_k(n, c, k):
    return 1.0 if n - c < k else 1 - comb(n - c, k) / comb(n, k)


def eval_tasks(n_tasks, seed0=EVAL_SEED0):
    return [make_task(random.Random(seed0 + i), KINDS[i % len(KINDS)]) for i in range(n_tasks)]


def evaluate(model, tok, n_tasks=90, k=8, temperature=1.0, batch=256, max_len=4096, verbose=0,
             engine="continuous"):
    """batch: cache rows (continuous engine: slots, refilled as episodes end; lockstep: batch size)."""
    model.eval()
    tasks = eval_tasks(n_tasks)
    flat = [t for t in tasks for _ in range(k)]
    if engine == "lockstep":
        roller = make_roller("lockstep", model, tok, max_len=max_len, temperature=temperature)
        eps = []
        for i in range(0, len(flat), batch):
            eps += roller.run(flat[i:i + batch])
    else:
        roller = make_roller("continuous", model, tok, max_len=max_len, temperature=temperature,
                             slots=min(batch, len(flat)))
        eps = roller.run(flat)
        roller.pool.shutdown()
        del roller
    by_kind = defaultdict(list)
    for ti, t in enumerate(tasks):
        by_kind[t.kind].append(eps[ti * k:(ti + 1) * k])
    rep = {}
    for kind, groups in by_kind.items():
        n_ok = [sum(e.correct for e in g) for g in groups]
        rep[kind] = dict(pass1=round(sum(c / k for c in n_ok) / len(groups), 3),
                         passk=round(sum(pass_at_k(k, c, k) for c in n_ok) / len(groups), 3))
    for name, sel in (("_in_dist", lambda k: k not in HELDOUT_KINDS), ("_held_out", lambda k: k in HELDOUT_KINDS)):
        gs = [g for kd, groups in by_kind.items() if sel(kd) for g in groups]
        if gs:
            rep[name] = dict(pass1=round(sum(e.correct for g in gs for e in g) / (len(gs) * k), 3),
                             passk=round(sum(pass_at_k(k, sum(e.correct for e in g), k) for g in gs) / len(gs), 3))
    allg = [e for e in eps]
    rep["_all"] = dict(
        pass1=round(sum(e.correct for e in allg) / len(allg), 3),
        passk=round(sum(v["passk"] * len(by_kind[kd]) for kd, v in rep.items()) / len(tasks), 3),
        grounded_given_correct=round(sum(e.grounded for e in allg if e.correct) / max(1, sum(e.correct for e in allg)), 3),
        submitted=round(sum(e.ws.submitted is not None for e in allg) / len(allg), 3),
        invented=round(sum(e.invented for e in allg) / len(allg), 3),
        turns=round(sum(e.turns for e in allg) / len(allg), 2),
        parse_errors=round(sum(e.parse_errors for e in allg) / len(allg), 3),
        repeat_calls=round(sum(e.repeats for e in allg) / len(allg), 3),
        flagged_turns=round(sum(sum(e.flags) for e in allg) / max(1, sum(e.turns for e in allg)), 4),
        gen_tokens=round(sum(e.gen_tokens for e in allg) / len(allg), 1),
        truncated=round(sum(e.truncated for e in allg) / len(allg), 3),
    )
    if verbose:
        from tiny_agent.chat import render
        print(render(eps[0].messages)[-3000:])
    return rep


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--ckpt", required=True)
    ap.add_argument("--tasks", type=int, default=150)
    ap.add_argument("--k", type=int, default=8)
    ap.add_argument("--temperature", type=float, default=1.0)
    ap.add_argument("--batch", type=int, default=256)
    ap.add_argument("--engine", default="continuous", choices=["continuous", "lockstep"])
    ap.add_argument("--verbose", type=int, default=1)
    a = ap.parse_args()
    tok = Tokenizer.from_file(f"{DATA}/tokenizer.json")
    model = load_model(a.ckpt, dtype=torch.bfloat16)
    rep = evaluate(model, tok, a.tasks, a.k, a.temperature, a.batch, verbose=a.verbose, engine=a.engine)
    print(json.dumps(rep, indent=1))


if __name__ == "__main__":
    main()