132 lines
5.3 KiB
Python
132 lines
5.3 KiB
Python
"""
|
|
Shared scoring / reporting for agent decision evals.
|
|
|
|
A concrete eval (stall, assignment, ...) supplies two callables:
|
|
- to_context(case) -> str : turn a dataset case into the model prompt context
|
|
- decide(context) -> obj : the decision function (returns an object with
|
|
.action / .reasoning / .confidence, or None)
|
|
|
|
and calls run_cli(...). Scoring treats a decision as correct if the majority
|
|
action across --runs samples is in the case's ``acceptable`` set (headline
|
|
"acceptable-rate"); matching the single ``ideal`` is the secondary "exact-rate".
|
|
"""
|
|
import argparse
|
|
import asyncio
|
|
import json
|
|
import sys
|
|
from collections import Counter
|
|
from datetime import datetime, timezone
|
|
from pathlib import Path
|
|
from statistics import mean
|
|
|
|
|
|
def load_cases(path: Path):
|
|
cases = []
|
|
for lineno, raw in enumerate(path.read_text(encoding="utf-8").splitlines(), 1):
|
|
line = raw.strip()
|
|
if not line or line.startswith("#"):
|
|
continue
|
|
try:
|
|
cases.append(json.loads(line))
|
|
except json.JSONDecodeError as e:
|
|
print(f"skipping malformed case on line {lineno}: {e}", file=sys.stderr)
|
|
return cases
|
|
|
|
|
|
def case_now(case):
|
|
"""Fixed date so only hour-of-day varies — keeps contexts reproducible."""
|
|
return datetime(2024, 1, 1, int(case.get("hour_of_day_utc", 12)), 0, tzinfo=timezone.utc)
|
|
|
|
|
|
async def evaluate(cases, to_context, decide, runs: int):
|
|
rows = []
|
|
for case in cases:
|
|
ctx = to_context(case)
|
|
decisions = []
|
|
for _ in range(runs):
|
|
d = await decide(ctx)
|
|
if d is not None:
|
|
decisions.append(d)
|
|
errors = runs - len(decisions)
|
|
|
|
if decisions:
|
|
counts = Counter(d.action for d in decisions)
|
|
majority, agree = counts.most_common(1)[0]
|
|
avg_conf = mean(d.confidence for d in decisions)
|
|
reasoning = next((d.reasoning for d in decisions if d.action == majority), "")
|
|
else:
|
|
majority, agree, avg_conf, reasoning = None, 0, 0.0, ""
|
|
|
|
rows.append({
|
|
"id": case["id"], "ideal": case.get("ideal"), "acceptable_set": case["acceptable"],
|
|
"got": majority, "agree": agree, "runs": runs, "errors": errors,
|
|
"avg_conf": avg_conf,
|
|
"acceptable": majority in case["acceptable"] if majority else False,
|
|
"exact": majority == case.get("ideal") if majority else False,
|
|
"reasoning": reasoning,
|
|
})
|
|
return rows
|
|
|
|
|
|
def print_report(rows):
|
|
print(f"\n{'case':<28} {'ideal':<14} {'got':<14} {'agree':<7} {'conf':<6} result")
|
|
print("-" * 82)
|
|
for r in rows:
|
|
if r["got"] is None:
|
|
result = "ERROR (no valid decision)"
|
|
elif r["exact"]:
|
|
result = "EXACT"
|
|
elif r["acceptable"]:
|
|
result = "ok (acceptable)"
|
|
else:
|
|
result = f"MISS (allowed: {', '.join(r['acceptable_set'])})"
|
|
print(f"{r['id']:<28} {str(r['ideal']):<14} {str(r['got']):<14} "
|
|
f"{r['agree']}/{r['runs']:<5} {r['avg_conf']:<6.2f} {result}")
|
|
if r["got"] is not None and not r["acceptable"]:
|
|
print(f"{'':<28} └─ why: {r['reasoning'][:110]}")
|
|
|
|
scored = [r for r in rows if r["got"] is not None]
|
|
n_scored = len(scored)
|
|
acc = sum(r["acceptable"] for r in scored)
|
|
exact = sum(r["exact"] for r in scored)
|
|
print("-" * 82)
|
|
print(f"cases: {len(rows)} scored: {n_scored} model-errors: {sum(r['errors'] for r in rows)}")
|
|
if n_scored:
|
|
print(f"acceptable-rate: {acc}/{n_scored} = {acc / n_scored:.0%}")
|
|
print(f"exact-rate: {exact}/{n_scored} = {exact / n_scored:.0%}")
|
|
return (acc / n_scored) if n_scored else 0.0
|
|
|
|
|
|
def run_cli(default_cases: Path, to_context, decide, llm_module):
|
|
"""argparse + dry-run + evaluate + report, shared across eval scripts.
|
|
``llm_module`` is core.llm so --model can override LLM_MODEL for the run."""
|
|
ap = argparse.ArgumentParser(description="Eval an agent decision against labelled cases.")
|
|
ap.add_argument("--cases", type=Path, default=default_cases, help="JSONL dataset")
|
|
ap.add_argument("--runs", type=int, default=1, help="samples per case (majority vote)")
|
|
ap.add_argument("--model", help="override LLM_MODEL for this run")
|
|
ap.add_argument("--dry-run", action="store_true", help="print contexts and labels; no API calls")
|
|
ap.add_argument("--min-pass-rate", type=float, default=0.0, help="exit non-zero if acceptable-rate below this")
|
|
args = ap.parse_args()
|
|
|
|
if args.model:
|
|
llm_module.LLM_MODEL = args.model
|
|
|
|
cases = load_cases(args.cases)
|
|
if not cases:
|
|
print(f"no cases found in {args.cases}", file=sys.stderr)
|
|
sys.exit(2)
|
|
|
|
if args.dry_run:
|
|
for case in cases:
|
|
print(f"\n=== {case['id']} (ideal={case.get('ideal')}, allowed={case['acceptable']}) ===")
|
|
print(to_context(case))
|
|
print(f"\n[dry-run] {len(cases)} cases, no API calls made.")
|
|
return
|
|
|
|
print(f"model: {llm_module.LLM_MODEL} cases: {len(cases)} runs/case: {args.runs}")
|
|
rows = asyncio.run(evaluate(cases, to_context, decide, args.runs))
|
|
pass_rate = print_report(rows)
|
|
if pass_rate < args.min_pass_rate:
|
|
print(f"\nFAIL: acceptable-rate {pass_rate:.0%} < required {args.min_pass_rate:.0%}", file=sys.stderr)
|
|
sys.exit(1)
|