Files
6krrt/plans/local-decision-classifier-prototype.py

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()