test / tracker.py
Nanny7's picture
Claude Opus 4.7 (1M context)
Add model prediction tracker with 4-tab Gradio UI
6933ba4
Raw History Blame Contribute Delete
5.1 kB
from __future__ import annotations
import hashlib
import io
import json
import threading
import uuid
from collections import defaultdict
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Iterator, Literal
InputKind = Literal["image", "text"]
def hash_text(text: str) -> tuple[str, dict[str, Any]]:
raw = text.encode("utf-8")
digest = hashlib.sha256(raw).hexdigest()
return f"sha256:{digest}", {"len": len(text)}
def hash_image(image) -> tuple[str, dict[str, Any]]:
buf = io.BytesIO()
image.save(buf, format="PNG")
data = buf.getvalue()
digest = hashlib.sha256(data).hexdigest()
width, height = image.size
channels = len(image.getbands())
return f"sha256:{digest}", {
"shape": [height, width, channels],
"bytes": len(data),
}
def _percentile(values: list[float], pct: float) -> float | None:
if not values:
return None
values = sorted(values)
if len(values) == 1:
return values[0]
rank = (pct / 100.0) * (len(values) - 1)
lo = int(rank)
hi = min(lo + 1, len(values) - 1)
frac = rank - lo
return values[lo] + (values[hi] - values[lo]) * frac
class RunTracker:
def __init__(self, log_path: Path):
self._log_path = Path(log_path)
self._lock = threading.Lock()
self._log_path.parent.mkdir(parents=True, exist_ok=True)
@property
def log_path(self) -> Path:
return self._log_path
def record(
self,
*,
model_version: str,
input_kind: InputKind,
input_hash: str,
input_meta: dict,
output: object,
latency_ms: float,
client_meta: dict | None = None,
) -> str:
run_id = str(uuid.uuid4())
row = {
"run_id": run_id,
"timestamp": datetime.now(timezone.utc).isoformat().replace("+00:00", "Z"),
"model_version": model_version,
"input_kind": input_kind,
"input_hash": input_hash,
"input_meta": input_meta,
"output": output,
"latency_ms": float(latency_ms),
"client_meta": client_meta or {},
}
line = json.dumps(row, default=str, ensure_ascii=False) + "\n"
with self._lock:
with self._log_path.open("a", encoding="utf-8") as fh:
fh.write(line)
return run_id
def _read_all(self) -> list[dict]:
if not self._log_path.exists():
return []
rows: list[dict] = []
with self._log_path.open("r", encoding="utf-8") as fh:
for line in fh:
line = line.strip()
if not line:
continue
try:
rows.append(json.loads(line))
except json.JSONDecodeError:
continue
return rows
def iter_runs(
self,
*,
model_version: str | None = None,
since: datetime | None = None,
until: datetime | None = None,
limit: int | None = None,
) -> Iterator[dict]:
rows = self._read_all()
rows.sort(key=lambda r: r.get("timestamp", ""), reverse=True)
count = 0
for row in rows:
if model_version and row.get("model_version") != model_version:
continue
ts_str = row.get("timestamp", "")
ts: datetime | None = None
if ts_str:
try:
ts = datetime.fromisoformat(ts_str.replace("Z", "+00:00"))
except ValueError:
ts = None
if since and ts and ts < since:
continue
if until and ts and ts > until:
continue
yield row
count += 1
if limit is not None and count >= limit:
break
def stats(self) -> dict:
rows = self._read_all()
per_model_latencies: dict[str, list[float]] = defaultdict(list)
per_day: dict[str, int] = defaultdict(int)
for row in rows:
model = row.get("model_version", "unknown")
latency = row.get("latency_ms")
if isinstance(latency, (int, float)):
per_model_latencies[model].append(float(latency))
ts_str = row.get("timestamp", "")
if ts_str:
day = ts_str[:10]
per_day[day] += 1
per_model = {
model: {
"count": len(latencies),
"p50_ms": _percentile(latencies, 50),
"p95_ms": _percentile(latencies, 95),
}
for model, latencies in per_model_latencies.items()
}
return {"per_model": per_model, "per_day": dict(sorted(per_day.items()))}
def clear(self) -> int:
with self._lock:
if not self._log_path.exists():
return 0
count = sum(1 for _ in self._log_path.open("r", encoding="utf-8"))
self._log_path.write_text("", encoding="utf-8")
return count