Benchmarking swaps models on the GPU production is serving from. Ollama evicts to make room, so the first run unpinned gemma4:e2b and left gemma4:e4b resident: the next voice turn would have paid a ~36s cold load, and only the monitoring noticing unexpected_models caught it. Snapshot residency and pinning before the run, then evict whatever the benchmark loaded and re-pin what was pinned before. The restore is wired to SIGTERM as well as the normal exit path. Python runs `finally` for SIGINT, which arrives as KeyboardInterrupt, but the default SIGTERM action terminates outright — so a `timeout`, a systemd stop or a plain `kill` skipped the guard entirely. That was not theoretical: the first SIGTERM after adding this bypassed it, and the pinned model survived only because the run had not reached the second model yet. Co-Authored-By: Claude <noreply@anthropic.com>
306 lines
13 KiB
Python
306 lines
13 KiB
Python
"""
|
|
Benchmark Steward routing quality against model and thinking settings.
|
|
|
|
Talks to Ollama directly. No Tatlock server, no agents, no tools, nothing is
|
|
executed — the mutating fixtures ("turn on the lights", "update the wiki") only
|
|
ever produce a routing decision. That makes this cheap and repeatable, and it
|
|
isolates the question: does the Steward still pick the right capabilities when
|
|
the model reasons less?
|
|
|
|
The request body is byte-identical to StewardAgent._call_ollama, plus the
|
|
`think` flag under test, so a cell labelled `unset` is exactly what production
|
|
sends today.
|
|
|
|
Three thinking settings, because "on vs off" hides the interesting case:
|
|
|
|
unset what production sends now. gemma4 reasons by default, and the
|
|
response carries no `thinking` field, so those tokens are generated
|
|
and discarded.
|
|
true reasoning requested explicitly and returned in `thinking`.
|
|
false reasoning suppressed.
|
|
|
|
Scoring is deliberately asymmetric. A missing capability under-routes and the
|
|
Butler answers without a tool it needed; a spurious one over-routes, and that is
|
|
a real agent call — a stray librarian is a multi-second web search on a query
|
|
that asked for arithmetic. Over-routing is the predicted failure when thinking
|
|
is off, so `forbid` violations are reported separately rather than folded into
|
|
one accuracy number.
|
|
|
|
Usage:
|
|
.venv/bin/python scripts/benchmark_routing.py
|
|
.venv/bin/python scripts/benchmark_routing.py --models gemma4:e2b
|
|
.venv/bin/python scripts/benchmark_routing.py --think false --repeats 3
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import json
|
|
import signal
|
|
import statistics
|
|
import sys
|
|
import time
|
|
from datetime import UTC, datetime
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
import httpx
|
|
|
|
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
|
sys.path.insert(0, str(PROJECT_ROOT))
|
|
|
|
from scripts.fixtures.routing_fixtures import FIXTURES # noqa: E402
|
|
from src.agents.steward.agent import build_steward_prompt # noqa: E402
|
|
from src.agents.steward.service import _DELEGATE_LINE_RE, _extract_capabilities # noqa: E402
|
|
from src.core.startup import register_household_members # noqa: E402
|
|
|
|
OLLAMA_URL = "http://localhost:11434"
|
|
DEFAULT_MODELS = ["gemma4:e2b", "gemma4:e4b"]
|
|
DEFAULT_THINK = ["unset", "true", "false"]
|
|
RESULTS_DIR = PROJECT_ROOT / "logs"
|
|
|
|
|
|
def build_body(model: str, prompt: str, think: str) -> dict[str, Any]:
|
|
"""Mirror StewardAgent._call_ollama exactly, then add the flag under test."""
|
|
body: dict[str, Any] = {
|
|
"model": model,
|
|
"prompt": prompt,
|
|
"stream": False,
|
|
"options": {
|
|
"temperature": 0.3, # Lower = more consistent
|
|
"top_p": 0.9,
|
|
},
|
|
}
|
|
if think != "unset":
|
|
body["think"] = think == "true"
|
|
return body
|
|
|
|
|
|
def call(client: httpx.Client, body: dict[str, Any]) -> dict[str, Any] | None:
|
|
try:
|
|
response = client.post(f"{OLLAMA_URL}/api/generate", json=body)
|
|
response.raise_for_status()
|
|
return response.json()
|
|
except Exception as exc: # noqa: BLE001 - a failed cell must not abort the run
|
|
print(f" ! {exc}", file=sys.stderr)
|
|
return None
|
|
|
|
|
|
PINNED_THRESHOLD_SECONDS = 86400
|
|
|
|
|
|
def snapshot_residency(client: httpx.Client) -> dict[str, bool]:
|
|
"""Which models are resident, and which of those are pinned.
|
|
|
|
Benchmarking swaps models on a GPU that production is using. Ollama evicts to
|
|
make room, so a run silently unpins the model Tatlock serves from and leaves
|
|
its own behind: the next voice turn then pays a ~36s cold load. Snapshot
|
|
before, restore after.
|
|
"""
|
|
try:
|
|
data = client.get(f"{OLLAMA_URL}/api/ps", timeout=10).json()
|
|
except Exception: # noqa: BLE001
|
|
return {}
|
|
resident: dict[str, bool] = {}
|
|
now = datetime.now(UTC)
|
|
for model in data.get("models", []):
|
|
expires = model.get("expires_at", "")
|
|
pinned = False
|
|
try:
|
|
delta = datetime.fromisoformat(expires.replace("Z", "+00:00")) - now
|
|
pinned = delta.total_seconds() > PINNED_THRESHOLD_SECONDS
|
|
except ValueError:
|
|
pass
|
|
resident[model["name"]] = pinned
|
|
return resident
|
|
|
|
|
|
def set_keep_alive(client: httpx.Client, model: str, keep_alive: Any) -> bool:
|
|
"""Load/unload/pin a model. Embedding models reject /api/generate."""
|
|
payload = {"model": model, "keep_alive": keep_alive}
|
|
for endpoint in ("generate", "embed"):
|
|
try:
|
|
response = client.post(f"{OLLAMA_URL}/api/{endpoint}", json=payload, timeout=180)
|
|
if response.status_code == 200:
|
|
return True
|
|
if response.status_code == 400 and "does not support generate" in response.text:
|
|
continue # embedding-only model; try /api/embed
|
|
return False
|
|
except Exception: # noqa: BLE001
|
|
return False
|
|
return False
|
|
|
|
|
|
def restore_residency(client: httpx.Client, before: dict[str, bool], used: list[str]) -> None:
|
|
"""Evict what the benchmark loaded, then re-pin what production had pinned."""
|
|
base = {name.split(":")[0] for name in before}
|
|
for model in used:
|
|
if model not in before and model.split(":")[0] not in base:
|
|
print(f" restoring: unloading benchmark model {model}")
|
|
set_keep_alive(client, model, 0)
|
|
for name, pinned in before.items():
|
|
if pinned:
|
|
ok = set_keep_alive(client, name, -1)
|
|
print(f" restoring: re-pinned {name}" if ok
|
|
else f" restoring: FAILED to re-pin {name} -- run warmup-ollama.sh")
|
|
|
|
|
|
def score(fixture: dict, found: list[str]) -> dict[str, Any]:
|
|
expected = set(fixture["expect"])
|
|
forbidden = set(fixture["forbid"])
|
|
got = set(found)
|
|
missing = sorted(expected - got)
|
|
spurious = sorted(got & forbidden)
|
|
return {
|
|
"found": found,
|
|
"missing": missing,
|
|
"spurious": spurious,
|
|
# Exact only when everything expected arrived and nothing forbidden did.
|
|
"exact": not missing and not spurious,
|
|
"under_routed": bool(missing),
|
|
"over_routed": bool(spurious),
|
|
}
|
|
|
|
|
|
def run_cell(client: httpx.Client, model: str, think: str, repeats: int) -> list[dict[str, Any]]:
|
|
rows: list[dict[str, Any]] = []
|
|
for fixture in FIXTURES:
|
|
prompt = build_steward_prompt(fixture["query"], [])
|
|
body = build_body(model, prompt, think)
|
|
for rep in range(repeats):
|
|
started = time.perf_counter()
|
|
data = call(client, body)
|
|
elapsed_ms = (time.perf_counter() - started) * 1000
|
|
if data is None:
|
|
rows.append({
|
|
"id": fixture["id"], "group": fixture["group"], "rep": rep,
|
|
"error": True, "exact": False, "under_routed": False, "over_routed": False,
|
|
})
|
|
continue
|
|
|
|
text = data.get("response", "") or ""
|
|
found = _extract_capabilities(text)
|
|
rows.append({
|
|
"id": fixture["id"],
|
|
"group": fixture["group"],
|
|
"rep": rep,
|
|
"error": False,
|
|
"latency_ms": round(elapsed_ms, 1),
|
|
"eval_tokens": data.get("eval_count"),
|
|
"prompt_tokens": data.get("prompt_eval_count"),
|
|
# Did the model obey the documented output shape at all?
|
|
"has_delegate_line": bool(_DELEGATE_LINE_RE.search(text)),
|
|
# Whether reasoning came back, as opposed to being generated and dropped.
|
|
"thinking_returned": bool(data.get("thinking")),
|
|
"response_chars": len(text),
|
|
**score(fixture, found),
|
|
})
|
|
return rows
|
|
|
|
|
|
def summarise(rows: list[dict[str, Any]]) -> dict[str, Any]:
|
|
ok = [r for r in rows if not r["error"]]
|
|
if not ok:
|
|
return {"n": 0, "errors": len(rows)}
|
|
latencies = [r["latency_ms"] for r in ok]
|
|
tokens = [r["eval_tokens"] for r in ok if r["eval_tokens"] is not None]
|
|
return {
|
|
"n": len(ok),
|
|
"errors": len(rows) - len(ok),
|
|
"exact_pct": round(100 * sum(r["exact"] for r in ok) / len(ok), 1),
|
|
"under_routed_pct": round(100 * sum(r["under_routed"] for r in ok) / len(ok), 1),
|
|
"over_routed_pct": round(100 * sum(r["over_routed"] for r in ok) / len(ok), 1),
|
|
"format_ok_pct": round(100 * sum(r["has_delegate_line"] for r in ok) / len(ok), 1),
|
|
"thinking_returned_pct": round(100 * sum(r["thinking_returned"] for r in ok) / len(ok), 1),
|
|
"latency_ms_median": round(statistics.median(latencies), 1),
|
|
"latency_ms_mean": round(statistics.fmean(latencies), 1),
|
|
"eval_tokens_median": round(statistics.median(tokens), 1) if tokens else None,
|
|
"eval_tokens_total": sum(tokens) if tokens else None,
|
|
}
|
|
|
|
|
|
def _raise_on_term(signum, _frame):
|
|
"""Turn SIGTERM into an exception so the restore guard actually runs.
|
|
|
|
Python runs `finally` for SIGINT, which arrives as KeyboardInterrupt, but the
|
|
default SIGTERM action terminates the process outright — a `timeout`, a
|
|
systemd stop or a plain `kill` would skip the restore and leave production
|
|
unpinned. Verified the hard way: an earlier SIGTERM here bypassed the guard
|
|
entirely and only luck kept the pinned model resident.
|
|
"""
|
|
raise KeyboardInterrupt(f"signal {signum}")
|
|
|
|
|
|
def main() -> int:
|
|
signal.signal(signal.SIGTERM, _raise_on_term)
|
|
|
|
parser = argparse.ArgumentParser(description=__doc__)
|
|
parser.add_argument("--models", default=",".join(DEFAULT_MODELS))
|
|
parser.add_argument("--think", default=",".join(DEFAULT_THINK),
|
|
help="comma-separated subset of unset,true,false")
|
|
parser.add_argument("--repeats", type=int, default=1)
|
|
parser.add_argument("--timeout", type=float, default=180.0)
|
|
args = parser.parse_args()
|
|
|
|
models = [m.strip() for m in args.models.split(",") if m.strip()]
|
|
think_modes = [t.strip() for t in args.think.split(",") if t.strip()]
|
|
|
|
# build_steward_prompt reads the registry, and the registry is populated at
|
|
# application startup. Without this the prompt lists no capabilities and every
|
|
# cell scores zero for reasons that have nothing to do with the model.
|
|
register_household_members()
|
|
|
|
print(f"{len(FIXTURES)} fixtures x {len(models)} models x {len(think_modes)} think "
|
|
f"x {args.repeats} repeats = {len(FIXTURES) * len(models) * len(think_modes) * args.repeats} calls\n")
|
|
|
|
cells: dict[str, Any] = {}
|
|
with httpx.Client(timeout=args.timeout) as client:
|
|
before = snapshot_residency(client)
|
|
pinned_before = [n for n, p in before.items() if p]
|
|
print(f"resident before: {sorted(before)}"
|
|
f"{f' (pinned: {pinned_before})' if pinned_before else ''}\n")
|
|
try:
|
|
for model in models:
|
|
# Absorb the cold load (~36s) outside the measurements.
|
|
print(f"warming {model} ...", flush=True)
|
|
call(client, build_body(model, "hi", "false"))
|
|
for think in think_modes:
|
|
key = f"{model}|think={think}"
|
|
print(f" {key} ...", end=" ", flush=True)
|
|
started = time.perf_counter()
|
|
rows = run_cell(client, model, think, args.repeats)
|
|
summary = summarise(rows)
|
|
cells[key] = {"summary": summary, "rows": rows}
|
|
print(f"exact={summary.get('exact_pct')}% "
|
|
f"over={summary.get('over_routed_pct')}% "
|
|
f"median={summary.get('latency_ms_median')}ms "
|
|
f"({time.perf_counter() - started:.0f}s)")
|
|
finally:
|
|
# finally, not a tidy exit path: a Ctrl-C or a failed cell must not
|
|
# leave production unpinned behind us.
|
|
print("\nrestoring GPU residency ...")
|
|
restore_residency(client, before, models)
|
|
|
|
RESULTS_DIR.mkdir(parents=True, exist_ok=True)
|
|
stamp = datetime.now(UTC).strftime("%Y%m%dT%H%M%SZ")
|
|
out = RESULTS_DIR / f"routing-bench-{stamp}.json"
|
|
out.write_text(json.dumps({
|
|
"generated_at": datetime.now(UTC).isoformat(),
|
|
"fixtures": len(FIXTURES),
|
|
"repeats": args.repeats,
|
|
"cells": cells,
|
|
}, indent=2))
|
|
|
|
print(f"\n{'cell':28} {'exact':>7} {'under':>7} {'over':>7} {'fmt':>6} {'tok':>7} {'ms':>8}")
|
|
print("-" * 76)
|
|
for key, cell in cells.items():
|
|
s = cell["summary"]
|
|
print(f"{key:28} {s.get('exact_pct'):>6}% {s.get('under_routed_pct'):>6}% "
|
|
f"{s.get('over_routed_pct'):>6}% {s.get('format_ok_pct'):>5}% "
|
|
f"{str(s.get('eval_tokens_median')):>7} {s.get('latency_ms_median'):>8}")
|
|
print(f"\nwritten to {out}")
|
|
return 0
|
|
|
|
|
|
if __name__ == "__main__":
|
|
sys.exit(main())
|