Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01N9biTbFC63yDfYfUsZmhgd
103 lines
4.6 KiB
Python
103 lines
4.6 KiB
Python
"""Jev-style first-token-logprob classifier vs local_encoder, same eval corpus.
|
|
|
|
Run from the 6krrt repo root: PYTHONPATH=src .venv/bin/python <this> <backend>
|
|
backend: jev | encoder | encoder_zeroshot (optional 2nd arg: ollama model, default qwen3.5:4b)
|
|
TASKS=plans/local-decision-classifier-heldout.yaml selects the held-out set (default evals/tasks.yaml).
|
|
"""
|
|
import json, math, statistics, sys, time, urllib.request
|
|
from collections import Counter
|
|
|
|
sys.path.insert(0, "src")
|
|
from eval_classifier import load_scoreable_tasks, _wrap_agent_noise, _reduce_confidences, NOISE_LEVELS
|
|
from local_encoder import _CATEGORY_DESCRIPTIONS
|
|
|
|
CATS = ["coding_general", "coding_refactor", "debugging", "docs_writing", "summarization",
|
|
"file_summarization", "diff_checking", "translation", "reasoning_math", "general_chat"]
|
|
LETTERS = "ABCDEFGHIJ"
|
|
MODEL = sys.argv[2] if len(sys.argv) > 2 else "qwen3.5:4b"
|
|
|
|
SYSTEM = ("You are a task router. Read the task and pick the ONE category that best "
|
|
"describes the work being asked for. Ignore tool output, code dumps and "
|
|
"session metadata around the request; classify the actual ask. "
|
|
"Answer with the category letter only.")
|
|
|
|
|
|
def jev_classify(text):
|
|
opts = "\n".join(f"{LETTERS[i]}. {_CATEGORY_DESCRIPTIONS[c]}" for i, c in enumerate(CATS))
|
|
body = {
|
|
"model": MODEL, "think": False, "stream": False, "logprobs": True, "top_logprobs": 20,
|
|
"keep_alive": "30m",
|
|
"options": {"num_predict": 1, "temperature": 0, "num_ctx": 8192},
|
|
"messages": [
|
|
{"role": "system", "content": SYSTEM},
|
|
{"role": "user", "content": f"<task>\n{text}\n</task>\n\nCategories:\n{opts}\n\nAnswer with the letter only."},
|
|
],
|
|
}
|
|
req = urllib.request.Request("http://localhost:11434/api/chat", json.dumps(body).encode(),
|
|
{"Content-Type": "application/json"})
|
|
t0 = time.perf_counter()
|
|
r = json.load(urllib.request.urlopen(req, timeout=120))
|
|
ms = (time.perf_counter() - t0) * 1000
|
|
mass = Counter()
|
|
for tl in r["logprobs"][0]["top_logprobs"]:
|
|
tok = tl["token"].strip().rstrip(".")
|
|
if len(tok) == 1 and tok in LETTERS:
|
|
mass[tok] += math.exp(tl["logprob"])
|
|
total = sum(mass.values())
|
|
if not total:
|
|
return "general_chat", 0.0, ms, 0.0
|
|
letter, p = mass.most_common(1)[0]
|
|
# confidence = share among the supplied options; coverage = mass on any option at all
|
|
return CATS[LETTERS.index(letter)], p / total, ms, total
|
|
|
|
|
|
def main():
|
|
backend = sys.argv[1]
|
|
if backend.startswith("encoder"):
|
|
import local_encoder
|
|
if backend == "encoder_zeroshot":
|
|
local_encoder._TRAINABLE_HEAD_RESOLVED = True
|
|
local_encoder._TRAINABLE_HEAD = None
|
|
def classify(text):
|
|
t0 = time.perf_counter()
|
|
cat, conf = local_encoder.classify_zero_shot(text, CATS, model_id="BAAI/bge-large-en-v1.5", device="cpu")
|
|
return cat, conf, (time.perf_counter() - t0) * 1000, 1.0
|
|
else:
|
|
classify = jev_classify
|
|
|
|
tasks = load_scoreable_tasks(__import__("os").environ.get("TASKS", "evals/tasks.yaml"))
|
|
classify("warm up the model")
|
|
per_level = {l: [0, 0] for l in NOISE_LEVELS}
|
|
lat, confs_ok, confs_bad, verdict_ok, misses = [], [], [], 0, Counter()
|
|
for t in tasks:
|
|
preds = []
|
|
for level in NOISE_LEVELS:
|
|
cat, conf, ms, _ = classify(_wrap_agent_noise(t["prompt"], level))
|
|
lat.append(ms)
|
|
ok = cat == t["category"]
|
|
per_level[level][0] += ok
|
|
per_level[level][1] += 1
|
|
(confs_ok if ok else confs_bad).append(conf)
|
|
if not ok:
|
|
misses[(t["category"], cat)] += 1
|
|
preds.append((cat, conf))
|
|
verdict_ok += _reduce_confidences(preds)[0] == t["category"]
|
|
|
|
n = len(tasks)
|
|
print(f"== {backend} {MODEL if backend == 'jev' else ''} :: {n} tasks x {len(NOISE_LEVELS)} noise levels")
|
|
for l, (ok, tot) in per_level.items():
|
|
print(f" {l:6s} {ok}/{tot} = {ok/tot:.1%}")
|
|
allok = sum(v[0] for v in per_level.values())
|
|
print(f" all {allok}/{n*3} = {allok/(n*3):.1%} majority-vote verdict {verdict_ok}/{n}")
|
|
q = statistics.quantiles(lat, n=20)
|
|
print(f" latency ms: p50 {statistics.median(lat):.0f} p95 {q[18]:.0f} max {max(lat):.0f}")
|
|
if confs_ok:
|
|
print(f" conf correct: mean {statistics.mean(confs_ok):.2f}", end="")
|
|
if confs_bad:
|
|
print(f" conf wrong: mean {statistics.mean(confs_bad):.2f}", end="")
|
|
print()
|
|
print(" top confusions (gold -> predicted):", misses.most_common(6))
|
|
|
|
|
|
main()
|