Record retrieval and Gemma evaluation results
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user