Record retrieval and Gemma evaluation results

This commit is contained in:
Hamza Ayed
2026-10-03 18:01:18 +03:00
parent b5a48c5a96
commit 1f8014f226
6 changed files with 363 additions and 15 deletions
+27 -5
View File
@@ -3,6 +3,7 @@
from __future__ import annotations
import argparse
import atexit
import json
import re
import sys
@@ -26,9 +27,18 @@ def load_cases(path: Path) -> list[dict]:
return cases
def post_json(url: str, payload: dict, timeout: float) -> dict:
def post_json(
url: str,
payload: dict,
timeout: float,
*,
token: str | None = None,
) -> dict:
body = json.dumps(payload, ensure_ascii=False).encode("utf-8")
request = Request(url, data=body, headers={"Content-Type": "application/json"}, method="POST")
headers = {"Content-Type": "application/json"}
if token:
headers["Authorization"] = f"Bearer {token}"
request = Request(url, data=body, headers=headers, method="POST")
with urlopen(request, timeout=timeout) as response:
return json.loads(response.read().decode("utf-8"))
@@ -70,6 +80,17 @@ def main() -> int:
missing = set(args.case_ids) - {case["id"] for case in cases}
if missing:
parser.error(f"Unknown case ID(s): {', '.join(sorted(missing))}")
base_url = args.base_url.rstrip("/")
session = post_json(base_url + "/v1/auth/local-session", {}, args.timeout)
token = session["access_token"]
def logout() -> None:
try:
post_json(base_url + "/v1/auth/logout", {}, args.timeout, token=token)
except Exception:
pass
atexit.register(logout)
warmup = None
if args.warmup:
warmup_payload = {
@@ -80,7 +101,7 @@ def main() -> int:
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)
post_json(base_url + "/v1/chat/completions", warmup_payload, args.timeout, token=token)
warmup = {"completed": True, "elapsed_seconds": round(time.perf_counter() - started, 2)}
print(f" warmup completed in {warmup['elapsed_seconds']}s")
results = []
@@ -90,12 +111,12 @@ def main() -> int:
}
if args.model:
payload["model"] = args.model
url = args.base_url.rstrip("/") + ("/v1/agent/run" if case["endpoint"] == "agent" else "/v1/chat/completions")
url = base_url + ("/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)
response = post_json(url, payload, args.timeout, token=token)
answer = extract_answer(case, response)
checks = []
for check in case.get("checks", []):
@@ -158,6 +179,7 @@ def main() -> int:
"results": results,
}
output.write_text(json.dumps(report, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
logout()
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)