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