Files
6krrt/scripts/train_encoder_head.py

406 lines
15 KiB
Python

#!/usr/bin/env python3
"""Fit the trainable logistic-regression head for local_encoder (one-shot, offline).
Spec item D of ``plans/local-encoder-backbone-and-accuracy-measurement.md``:
nearest-centroid IS a linear classifier whose weights are hand-written
description strings; this script fits those weights instead, on the same
frozen encoder embeddings, and Platt-calibrates the result so the reported
confidence is P(correct) — spec item C falls out for free.
No raw task text is ever captured (``docs/operations.md``): the training
corpus is SYNTHETIC, generated from the committed ``evals/tasks.yaml`` task
set in agent-traffic shape — every prompt wrapped in the same structural
agent-session noise ``src/local_encoder.py::_isolate_task_text`` strips, at
the three noise levels the committed eval harness uses. The corpus and the
fitted coefficients are both committed, diffable artifacts:
evals/synthetic/encoder-training.jsonl one row per (task, noise)
evals/synthetic/encoder-head-coefficients.json coef/intercept + Platt a/b
What it does, in order:
1. Generate the corpus: ALL 57 tasks (including ``tool_use_agentic``,
labeled from the task's own ``category`` field) x 3 noise variants
(clean / short / long) via ``eval_classifier._wrap_agent_noise``.
2. Embed every row through local_encoder's FROZEN pipeline
(``local_encoder._embed_task`` — isolate, fit to the token window,
per-model query prefix, one forward pass, snapshot-configured pool,
L2-normalize), so the head is trained on exactly what serving feeds it.
3. Fit ``LogisticRegression(solver="lbfgs")`` — convex, so a seeded run is
reproducible — and Platt-calibrate with
``CalibratedClassifierCV(method="sigmoid", ensemble=False)``.
4. Verify the serialized artifact reproduces sklearn's own
``predict_proba`` in pure Python before writing it.
5. Write the coefficients JSON and report top-1 accuracy on the 46-task
scoreable eval set, nearest-centroid vs the fitted head.
sklearn is a TRAINING-time dependency only (``requirements-encoder.txt``);
the fitted artifact is served by ``local_encoder._TrainableHead`` in pure
Python, so the router needs no sklearn at runtime.
PYTHONPATH=src python scripts/train_encoder_head.py
PYTHONPATH=src python scripts/train_encoder_head.py --device cuda
"""
from __future__ import annotations
import argparse
import json
import math
import sys
from collections import Counter
from pathlib import Path
from typing import Any, Optional
_PROJECT_ROOT = Path(__file__).resolve().parent.parent
_SRC = str(_PROJECT_ROOT / "src")
if _SRC not in sys.path:
sys.path.insert(0, _SRC)
import yaml
from config import load_config
TASKS_PATH = _PROJECT_ROOT / "evals" / "tasks.yaml"
CORPUS_PATH = _PROJECT_ROOT / "evals" / "synthetic" / "encoder-training.jsonl"
COEFFICIENTS_PATH = (
_PROJECT_ROOT / "evals" / "synthetic" / "encoder-head-coefficients.json"
)
# The committed artifact is trained for the documented default backbone.
# config.local.yaml overlays are deliberately IGNORED here: the committed
# coefficients must not depend on which machine happens to run the script.
DEFAULT_MODEL_ID = "BAAI/bge-large-en-v1.5"
DEFAULT_DEVICE = "cpu"
# Coefficient rounding: embeddings are float32 (~7 significant digits), so
# 8 decimals on coef/intercept keeps the serialized artifact well under the
# embedding's own precision while staying small and diffable.
_COEF_DECIMALS = 8
# Platt calibrators are two scalars per class; full precision is free.
_CALIB_DECIMALS = 8
# Tolerance for the serialized-artifact replication check: the rounded
# coef/intercept move each decision value by at most ~1e-6 on unit-norm
# embeddings, so the reconstructed probabilities must match sklearn's with
# an order of magnitude of headroom beyond that.
_REPLICATION_TOLERANCE = 1e-4
# --- corpus generation ----------------------------------------------------
def build_corpus() -> list[dict]:
"""One row per (task, noise_level): text, category, noise_level.
All 57 tasks, including tool_use_agentic: the eval set excludes
tool-use because classifier.candidate_categories deliberately cannot
emit that label, but the head's class set is not bound by that policy —
the serving-side wiring (``classify_zero_shot``) restricts the head's
probabilities to the caller's candidate set at request time, so extra
classes are inert until a candidate set asks for them.
"""
from eval_classifier import NOISE_LEVELS, _wrap_agent_noise
tasks = (yaml.safe_load(TASKS_PATH.read_text()) or {}).get("tasks") or []
if not tasks:
raise SystemExit(f"no tasks found in {TASKS_PATH}")
rows: list[dict] = []
for task in tasks:
for level in NOISE_LEVELS:
rows.append({
"text": _wrap_agent_noise(task["prompt"], level),
"category": task["category"],
"noise_level": level,
})
return rows
def write_corpus(rows: list[dict]) -> None:
CORPUS_PATH.parent.mkdir(parents=True, exist_ok=True)
with CORPUS_PATH.open("w") as f:
for row in rows:
f.write(json.dumps(row, ensure_ascii=False) + "\n")
counts: Counter = Counter(r["category"] for r in rows)
print(f"wrote {len(rows)} rows -> {CORPUS_PATH}")
for cat in sorted(counts):
print(f" {cat:20s} {counts[cat]} rows")
# --- embedding extraction (the frozen pipeline) ----------------------------
def extract_embeddings(
rows: list[dict], model_id: str, device: str,
) -> tuple[list[list[float]], list[str]]:
"""X (one L2-normalized embedding per row) and y, via local_encoder.
Deliberately batch-of-1: that is exactly the shape the serving path
embeds (``classify_zero_shot`` -> ``_embed_task``), so training and
serving share identical preprocessing — the whole point of fitting the
head on frozen embeddings.
"""
import local_encoder
model, tokenizer = local_encoder._load_model_tokenizer(model_id, device)
_, tensor_device = local_encoder._resolve_device(device)
X: list[list[float]] = []
y: list[str] = []
for i, row in enumerate(rows):
emb = local_encoder._embed_task(
row["text"], model, tokenizer, model_id, tensor_device,
)
X.append(emb[0].tolist())
y.append(row["category"])
if (i + 1) % 25 == 0 or i + 1 == len(rows):
print(f" embedded {i + 1}/{len(rows)}")
return X, y
# --- fit + calibrate -------------------------------------------------------
def fit_calibrated_head(
X: list[list[float]], y: list[str], model_id: str, cv: int,
) -> tuple[dict, Any]:
"""Fit lbfgs logistic regression, Platt-calibrate it, return (payload, model).
``ensemble=False`` matters for the artifact: it fits the per-class
sigmoid calibrators on out-of-fold predictions, then refits ONE base
estimator on all the data — so the serialized artifact is a single
coef/intercept matrix plus one (a, b) pair per class, exactly what
``_TrainableHead`` serves. (The default ``ensemble=True`` would serialize
one model per fold.)
"""
from sklearn.calibration import CalibratedClassifierCV
from sklearn.linear_model import LogisticRegression
# lbfgs on the multinomial objective is sklearn's default for multiclass
# (the explicit multi_class= parameter was deprecated in 1.5 and removed
# in 1.8 — passing it here would raise on the installed version).
base = LogisticRegression(solver="lbfgs", max_iter=2000)
calibrated = CalibratedClassifierCV(
base, method="sigmoid", cv=cv, ensemble=False,
)
calibrated.fit(X, y)
fitted = calibrated.calibrated_classifiers_[0]
estimator = fitted.estimator
payload = {
"model_id": model_id,
"embedding_dim": len(X[0]),
"classes": [str(c) for c in estimator.classes_],
"coef": [
[round(float(v), _COEF_DECIMALS) for v in row]
for row in estimator.coef_
],
"intercept": [
round(float(b), _COEF_DECIMALS) for b in estimator.intercept_
],
"calibration": "platt",
"calibration_a": [
round(float(c.a_), _CALIB_DECIMALS) for c in fitted.calibrators
],
"calibration_b": [
round(float(c.b_), _CALIB_DECIMALS) for c in fitted.calibrators
],
}
return payload, calibrated
def _platt_sigmoid(a: float, b: float, decision: float) -> float:
"""expit(-(a*decision + b)) — the exact _SigmoidCalibration.predict form."""
z = a * decision + b
if z > 700.0:
return 0.0
if z < -700.0:
return 1.0
return 1.0 / (1.0 + math.exp(z))
def verify_serialized_payload(
payload: dict, calibrated: Any, X: list[list[float]],
) -> float:
"""Max probability drift between the serialized artifact and sklearn.
Recomputes predict_proba for the training rows from the ROUNDED payload
in pure Python (the exact arithmetic ``_TrainableHead`` runs at serving
time) and compares against sklearn's own calibrated predict_proba. A
drift beyond the tolerance means the artifact would not faithfully
reproduce the model it claims to serialize — refuse to write it.
"""
sklearn_proba = calibrated.predict_proba(X)
classes: list[str] = payload["classes"]
worst = 0.0
for i, row in enumerate(X):
# sklearn's calibrated predict_proba columns follow estimator.classes_,
# which is exactly the payload's classes order.
reconstructed: list[float] = []
for k, _cat in enumerate(classes):
decision = payload["intercept"][k]
coef_row = payload["coef"][k]
for j, xv in enumerate(row):
decision += coef_row[j] * xv
reconstructed.append(_platt_sigmoid(
payload["calibration_a"][k],
payload["calibration_b"][k],
decision,
))
total = sum(reconstructed)
if total <= 0.0:
reconstructed = [1.0 / len(classes)] * len(classes)
else:
reconstructed = [d / total for d in reconstructed]
for k in range(len(classes)):
worst = max(worst, abs(reconstructed[k] - sklearn_proba[i][k]))
return worst
# --- before/after accuracy on the scoreable eval set -----------------------
def _accuracy(rows: list) -> Optional[float]:
if not rows:
return None
return sum(1 for _tid, gold, pred, _c in rows if gold == pred) / len(rows)
def _print_accuracy(title: str, results: dict, noise_levels: tuple) -> None:
print(f"\n## {title}")
print(f"{'noise':8} {'correct':>8} {'total':>7} {'accuracy':>10}")
for level in noise_levels:
rows = results["gold_history"][level]
acc = _accuracy(rows)
if acc is None:
continue
correct = sum(1 for _t, g, p, _c in rows if g == p)
print(f"{level:8} {correct:>8} {len(rows):>7} {acc:10.3f}")
tv_rows = [
(_id, gold, pred, conf)
for _id, (gold, pred, conf) in sorted(results["task_verdicts"].items())
]
acc = _accuracy(tv_rows)
if acc is not None:
correct = sum(1 for _t, g, p, _c in tv_rows if g == p)
print(f"{'overall':8} {correct:>8} {len(tv_rows):>7} {acc:10.3f}")
def run_accuracy(
title: str,
tasks: list[dict],
categories: list[str],
noise_levels: tuple,
*,
model_id: str,
device: str,
head: Optional[Any],
) -> None:
import local_encoder
from eval_classifier import run_eval
# The head switch is the module global: None reproduces the pre-head
# classifier exactly; an installed head takes the fitted path.
local_encoder._TRAINABLE_HEAD = head
local_encoder._TRAINABLE_HEAD_RESOLVED = True
results = run_eval(
tasks, categories, model_id=model_id, device=device,
noise_levels=noise_levels,
)
_print_accuracy(title, results, noise_levels)
# --- CLI -------------------------------------------------------------------
def main() -> int:
ap = argparse.ArgumentParser(description=__doc__)
ap.add_argument("--model", help=f"encoder model id (default {DEFAULT_MODEL_ID})")
ap.add_argument("--device", help=f"encoder device (default {DEFAULT_DEVICE})")
ap.add_argument(
"--cv", type=int, default=3,
help="folds for CalibratedClassifierCV (default 3; 3 or 5 are sane)",
)
ap.add_argument(
"--skip-eval", action="store_true",
help="skip the before/after accuracy report (still trains + writes)",
)
args = ap.parse_args()
cfg = load_config(
str(_PROJECT_ROOT / "config" / "config.yaml"), include_overlay=False,
)
enc = cfg.classifier.encoder
model_id = args.model or (enc.model if enc is not None else DEFAULT_MODEL_ID)
device = args.device or (enc.device if enc is not None else DEFAULT_DEVICE)
print(f"model={model_id} device={device} cv={args.cv}")
rows = build_corpus()
write_corpus(rows)
print(f"\nembedding {len(rows)} rows through the frozen pipeline...")
X, y = extract_embeddings(rows, model_id, device)
payload, calibrated = fit_calibrated_head(X, y, model_id, cv=args.cv)
# sklearn's honest generalization estimate for the fit, before any
# artifact talk: 5-fold stratified in-corpus CV of the uncalibrated LR.
from sklearn.linear_model import LogisticRegression
from sklearn.model_selection import StratifiedKFold, cross_val_score
scores = cross_val_score(
LogisticRegression(solver="lbfgs", max_iter=2000), X, y,
cv=StratifiedKFold(n_splits=5, shuffle=True, random_state=0),
)
print(
f"in-corpus 5-fold CV accuracy (uncalibrated LR): "
f"{scores.mean():.3f} +/- {scores.std():.3f}"
)
worst = verify_serialized_payload(payload, calibrated, X)
if worst > _REPLICATION_TOLERANCE:
print(
f"serialized artifact drifts {worst:.3e} from sklearn predict_proba "
f"(tolerance {_REPLICATION_TOLERANCE:.1e}) — refusing to write",
file=sys.stderr,
)
return 1
print(f"serialized replication check: max prob drift {worst:.2e} (ok)")
COEFFICIENTS_PATH.parent.mkdir(parents=True, exist_ok=True)
COEFFICIENTS_PATH.write_text(json.dumps(payload, indent=1) + "\n")
print(
f"wrote {COEFFICIENTS_PATH} "
f"({len(payload['classes'])} classes x {payload['embedding_dim']} dims)"
)
if args.skip_eval:
return 0
from eval_classifier import NOISE_LEVELS, load_scoreable_tasks
tasks = load_scoreable_tasks(str(TASKS_PATH))
categories = list(cfg.classifier_candidate_categories)
unknown = {t["category"] for t in tasks} - set(categories)
if unknown:
print(
f"tasks reference categories missing from the candidate set: "
f"{sorted(unknown)}",
file=sys.stderr,
)
return 1
import local_encoder
head = local_encoder._TrainableHead(payload, str(COEFFICIENTS_PATH))
run_accuracy(
"Nearest-centroid (before trainable head)",
tasks, categories, NOISE_LEVELS,
model_id=model_id, device=device, head=None,
)
run_accuracy(
"Fitted head (after, Platt-calibrated)",
tasks, categories, NOISE_LEVELS,
model_id=model_id, device=device, head=head,
)
return 0
if __name__ == "__main__":
raise SystemExit(main())