raya / training /evaluate.py
cderinbogaz's picture
Add training kit: train your own System-1 model
48c8658 verified
Raw History Blame Contribute Delete
4.22 kB
"""Score a trained model on a labelled JSONL file: accuracy, macro-F1, confusion and latency.
python evaluate.py --model my-router --task task.example.json --data data/example.jsonl
python evaluate.py --model my-router --onnx my-router/onnx/model-int8-blockwise.onnx ...
Accuracy counts rows whose annotators all agreed (a single gold label). Always test on data the
model never trained or validated on. The majority-label baseline is printed for comparison.
"""
from __future__ import annotations
import argparse
import json
import os
import time
os.environ.setdefault("USE_TF", "0")
import numpy as np
from common import answer_probs, load_rows, load_task, row_state
def load_model(path: str, onnx_path: str | None, device: str | None, max_tokens: int):
if onnx_path:
from laya.onnx_agent import ONNXAgent
model = ONNXAgent(path, onnx_path=onnx_path)
else:
import laya
model = laya.Agent(path, device=device)
model.cfg["max_len"] = max_tokens
return model
def main() -> None:
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
ap.add_argument("--model", required=True, help="checkpoint dir or Hub repo id")
ap.add_argument("--task", required=True)
ap.add_argument("--data", required=True, help="labelled JSONL the model has not seen")
ap.add_argument("--onnx", help="score this ONNX file (with the checkpoint's config/tokenizer) instead")
ap.add_argument("--device", default=None)
ap.add_argument("--max-tokens", type=int, default=512)
ap.add_argument("--out", help="write per-row predictions to this JSONL")
args = ap.parse_args()
task = load_task(args.task)
labels, questions = task["labels"], task["questions"]
rows = [r for r in load_rows(args.data, labels) if r["gold"] is not None]
model = load_model(args.model, args.onnx, args.device, args.max_tokens)
preds = {qi: [] for qi in range(len(questions))}
latencies = []
out = open(args.out, "w", encoding="utf-8") if args.out else None
for row in rows:
record = {"id": row["id"], "gold": row["gold"], "predictions": {}}
for qi, q in enumerate(questions):
t = time.perf_counter()
answer = model.system_one(row_state(row), {"q": q})["answers"]["q"]
latencies.append((time.perf_counter() - t) * 1000)
probs = answer_probs(answer, q, labels)
preds[qi].append(int(np.argmax(probs)))
record["predictions"][qi] = dict(zip(labels, [round(p, 4) for p in probs]))
if out:
out.write(json.dumps(record, ensure_ascii=False) + "\n")
if out:
out.close()
golds = [labels.index(r["gold"]) for r in rows]
majority = max(labels, key=lambda l: sum(r["gold"] == l for r in rows))
print(f"{len(rows)} rows with a gold label; always answering {majority!r} scores "
f"{sum(r['gold'] == majority for r in rows) / len(rows):.1%}")
for qi, q in enumerate(questions):
p = preds[qi]
acc = np.mean([a == b for a, b in zip(p, golds)])
f1s = []
for k in range(len(labels)):
tp = sum(a == k == b for a, b in zip(p, golds))
fp = sum(a == k != b for a, b in zip(p, golds))
fn = sum(b == k != a for a, b in zip(p, golds))
f1s.append(2 * tp / (2 * tp + fp + fn) if tp else 0.0)
confusion = [[sum(g == i and a == j for a, g in zip(p, golds)) for j in range(len(labels))]
for i in range(len(labels))]
print(f"\nquestion {qi} ({q['type']}): {q['instructions'][:70]!r}")
print(f" accuracy {acc:.1%} macro-F1 {np.mean(f1s):.3f}")
print(" confusion (rows = gold, columns = predicted):")
width = max(len(l) for l in labels)
print(" " + " " * width + " " + " ".join(f"{l:>{width}}" for l in labels))
for label, line in zip(labels, confusion):
print(f" {label:>{width}} " + " ".join(f"{n:>{width}}" for n in line))
lat = np.array(latencies)
print(f"\nlatency per decision: p50 {np.percentile(lat, 50):.0f} ms, p95 {np.percentile(lat, 95):.0f} ms")
if __name__ == "__main__":
main()