Download code/scripts/eval_agent.py from darioooooo0o/tiny-agent-112m: direct link, hf CLI and curl.
- Browser
- Download file 4.46 kB
-
https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/scripts/eval_agent.py
- Command line
-
hf download hf://darioooooo0o/tiny-agent-112m/code/scripts/eval_agent.py
-
curl -L -o eval_agent.py https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/scripts/eval_agent.py
4.46 kB
| """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() | |