Files
sovereign_ai/SovereignAI-Starter/scripts/run_model_eval.py
T

171 lines
8.0 KiB
Python

"""Run the local Arabic/programming baseline and save outputs for human review."""
from __future__ import annotations
import argparse
import json
import re
import sys
import time
import statistics
from datetime import datetime, timezone
from pathlib import Path
from urllib.error import HTTPError, URLError
from urllib.request import Request, urlopen
ROOT = Path(__file__).resolve().parents[1]
DATASET = ROOT / "evals" / "baseline_arabic_programming.jsonl"
def load_cases(path: Path) -> list[dict]:
cases = [json.loads(line) for line in path.read_text(encoding="utf-8").splitlines() if line.strip()]
ids = [case["id"] for case in cases]
if len(ids) != len(set(ids)):
raise ValueError("Dataset contains duplicate case IDs")
return cases
def post_json(url: str, payload: dict, timeout: float) -> dict:
body = json.dumps(payload, ensure_ascii=False).encode("utf-8")
request = Request(url, data=body, headers={"Content-Type": "application/json"}, method="POST")
with urlopen(request, timeout=timeout) as response:
return json.loads(response.read().decode("utf-8"))
def extract_answer(case: dict, response: dict) -> str:
if case["endpoint"] == "agent":
return str(response.get("result", response.get("answer", "")))
choices = response.get("choices", [])
if not choices:
raise ValueError("Chat API returned no choices")
return str(choices[0].get("message", {}).get("content", ""))
def run_check(check: dict, answer: str) -> bool:
if check["type"] == "contains":
return check["value"].casefold() in answer.casefold()
if check["type"] == "regex":
return re.search(check["value"], answer, re.IGNORECASE) is not None
if check["type"] == "fenced_code":
pattern = rf"```{re.escape(check['language'])}\b[\s\S]+?```"
return re.search(pattern, answer, re.IGNORECASE) is not None
raise ValueError(f"Unsupported automatic check: {check['type']}")
def main() -> int:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--base-url", default="http://127.0.0.1:8000", help="Base URL of the local FastAPI service")
parser.add_argument("--model", default=None, help="Model name; defaults to the API's configured model")
parser.add_argument("--timeout", type=float, default=180.0, help="Per-question request timeout in seconds")
parser.add_argument("--warmup", action="store_true", help="Warm the requested model once before measured cases")
parser.add_argument("--dataset", type=Path, default=DATASET)
parser.add_argument("--case", action="append", dest="case_ids", help="Run only this case ID; may be repeated")
parser.add_argument("--output", type=Path, default=None, help="Output JSON path; defaults to evals/results/<model>_<date>.json")
args = parser.parse_args()
cases = load_cases(args.dataset)
if args.case_ids:
cases = [case for case in cases if case["id"] in args.case_ids]
missing = set(args.case_ids) - {case["id"] for case in cases}
if missing:
parser.error(f"Unknown case ID(s): {', '.join(sorted(missing))}")
warmup = None
if args.warmup:
warmup_payload = {
"messages": [{"role": "user", "content": "أجب بكلمة جاهز."}],
"max_tokens": 8,
}
if args.model:
warmup_payload["model"] = args.model
print("Warming model before timing the evaluation...", flush=True)
started = time.perf_counter()
post_json(args.base_url.rstrip("/") + "/v1/chat/completions", warmup_payload, args.timeout)
warmup = {"completed": True, "elapsed_seconds": round(time.perf_counter() - started, 2)}
print(f" warmup completed in {warmup['elapsed_seconds']}s")
results = []
for index, case in enumerate(cases, start=1):
payload = {"task": case["input"]} if case["endpoint"] == "agent" else {
"messages": [{"role": "user", "content": case["input"]}]
}
if args.model:
payload["model"] = args.model
url = args.base_url.rstrip("/") + ("/v1/agent/run" if case["endpoint"] == "agent" else "/v1/chat/completions")
print(f"[{index}/{len(cases)}] {case['id']} ({case['endpoint']})", flush=True)
started = time.perf_counter()
result = {"id": case["id"], "category": case["category"], "rubric": case.get("rubric", [])}
try:
response = post_json(url, payload, args.timeout)
answer = extract_answer(case, response)
checks = []
for check in case.get("checks", []):
passed = (
str(response.get("tool", "")).casefold() == check["value"].casefold()
if check["type"] == "tool_is"
else run_check(check, answer)
)
checks.append({"check": check, "passed": passed})
result.update(
status="ok",
elapsed_seconds=round(time.perf_counter() - started, 2),
answer=answer,
tool=response.get("tool"),
checks=checks,
human_rating=None,
human_notes="",
)
print(f" completed in {result['elapsed_seconds']}s")
except HTTPError as error:
detail = error.read().decode("utf-8", errors="replace")
message = f"HTTP {error.code}: {detail[:1000]}"
result.update(
status="error",
elapsed_seconds=round(time.perf_counter() - started, 2),
error=message,
checks=[{"check": check, "passed": False} for check in case.get("checks", [])],
)
print(f" ERROR: {message}", file=sys.stderr, flush=True)
except (URLError, TimeoutError, ValueError, KeyError) as error:
result.update(
status="error",
elapsed_seconds=round(time.perf_counter() - started, 2),
error=str(error),
checks=[{"check": check, "passed": False} for check in case.get("checks", [])],
)
print(f" ERROR: {error}", file=sys.stderr, flush=True)
results.append(result)
model_slug = re.sub(r"[^A-Za-z0-9._-]+", "_", args.model or "default-model")
output = args.output or ROOT / "evals" / "results" / f"{model_slug}_{datetime.now().strftime('%Y-%m-%d_%H%M%S')}.json"
output.parent.mkdir(parents=True, exist_ok=True)
report = {
"created_at_utc": datetime.now(timezone.utc).isoformat(),
"base_url": args.base_url,
"requested_model": args.model,
"dataset": str(args.dataset.relative_to(ROOT) if args.dataset.is_relative_to(ROOT) else args.dataset),
"evaluation_note": "Automatic checks are narrow format/value checks. Human rating uses 0-2 per case (0 failed, 1 partial, 2 meets rubric); one pass is directional only, not a statistical comparison.",
"warmup": warmup,
"latency_summary_seconds": (
{
"mean": round(statistics.mean(item["elapsed_seconds"] for item in results if item["status"] == "ok"), 2),
"median": round(statistics.median(item["elapsed_seconds"] for item in results if item["status"] == "ok"), 2),
"min": round(min(item["elapsed_seconds"] for item in results if item["status"] == "ok"), 2),
"max": round(max(item["elapsed_seconds"] for item in results if item["status"] == "ok"), 2),
}
if any(item["status"] == "ok" for item in results)
else None
),
"results": results,
}
output.write_text(json.dumps(report, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
failed = sum(item["status"] != "ok" for item in results)
passed_checks = sum(check["passed"] for item in results for check in item.get("checks", []))
total_checks = sum(len(case.get("checks", [])) for case in cases)
print(f"Saved {len(results)} cases to {output}")
print(f"Requests failed: {failed}; automatic checks passed: {passed_checks}/{total_checks}")
return 1 if failed else 0
if __name__ == "__main__":
raise SystemExit(main())