Files
AI_engine/evals/_harness.py

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)