Files
nucleic-purpose-classifier/eval.py
T

471 lines
18 KiB
Python
Raw Normal View History

#!/usr/bin/env python3
"""Evaluate a trained purpose-lite checkpoint on the frozen v1 test set."""
from __future__ import annotations
import argparse
import json
import math
import statistics
import sys
import time
from collections import Counter
from pathlib import Path
from typing import Any, Sequence
from purpose_data import (
HARD_SLICES,
LABELS,
DataError,
load_classifiable_fixtures,
load_jsonl,
prompt_hash,
write_json,
)
from train import (
MAX_LENGTH,
classification_metrics,
confidence_score,
encode_fixed_shape,
expected_calibration_error,
)
SCRIPT_DIR = Path(__file__).resolve().parent
REPOSITORY_ROOT = SCRIPT_DIR.parent.parent
DEFAULT_MODEL_DIR = SCRIPT_DIR / "outputs" / "purpose-lite-v1" / "model"
DEFAULT_CALIBRATION = SCRIPT_DIR / "outputs" / "purpose-lite-v1" / "calibration.json"
DEFAULT_TEST = SCRIPT_DIR / "data" / "frozen-test-v1.jsonl"
DEFAULT_FIXTURES = (
REPOSITORY_ROOT
/ "Tests"
/ "NucleicCoreTests"
/ "Fixtures"
/ "purpose-prompts.json"
)
DEFAULT_REPORT = SCRIPT_DIR / "outputs" / "purpose-lite-v1" / "frozen-eval.json"
ROUTING_LEVELS = ("quick", "light", "balanced", "deep", "max")
# Cost tiers mirrored from IntelligenceRouter.matrix for the two provider lanes. Keep this
# table in sync with Sources/NucleicCore/IntelligenceRouting.swift; the eval report names
# every violating prompt/level/lane so a matrix change cannot fail opaquely.
ROUTING_COST_TIERS = {
"planning": {"claude": (1, 2, 2, 3, 3), "codex": (1, 1, 2, 2, 2)},
"backendImpl": {"claude": (1, 1, 2, 3, 3), "codex": (0, 1, 2, 2, 2)},
"frontendImpl": {"claude": (1, 1, 2, 2, 3), "codex": (0, 1, 1, 2, 2)},
"quickFix": {"claude": (1, 1, 1, 2, 2), "codex": (0, 0, 1, 1, 2)},
"refactor": {"claude": (1, 1, 1, 2, 3), "codex": (0, 1, 1, 2, 2)},
"debugging": {"claude": (1, 1, 2, 3, 3), "codex": (1, 1, 2, 2, 2)},
"review": {"claude": (0, 1, 1, 2, 2), "codex": (0, 1, 1, 2, 2)},
"writing": {"claude": (0, 1, 1, 2, 2), "codex": (0, 0, 1, 1, 1)},
}
def _device(torch: Any, requested: str) -> Any:
if requested != "auto":
return torch.device(requested)
if torch.cuda.is_available():
return torch.device("cuda")
if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
return torch.device("mps")
return torch.device("cpu")
def _load_calibration(path: Path) -> dict[str, Any]:
try:
value = json.loads(path.read_text(encoding="utf-8"))
temperature = float(value["temperature"])
high = float(value["confidence"]["high"]["minimumScore"])
medium = float(value["confidence"]["medium"]["minimumScore"])
except (OSError, UnicodeError, json.JSONDecodeError, KeyError, TypeError, ValueError) as exc:
raise DataError(f"{path}: invalid calibration config: {exc}") from exc
if not math.isfinite(temperature) or temperature <= 0:
raise DataError(f"{path}: temperature must be finite and positive")
if not 0 <= medium <= high:
raise DataError(f"{path}: expected 0 <= medium <= high confidence thresholds")
return value
def _percentile(values: Sequence[float], percentile: float) -> float:
if not values:
return 0.0
ordered = sorted(values)
index = min(len(ordered) - 1, math.ceil(percentile * len(ordered)) - 1)
return ordered[index]
def _synchronize(torch: Any, device: Any) -> None:
if device.type == "cuda":
torch.cuda.synchronize()
elif device.type == "mps":
torch.mps.synchronize()
def routing_tier_drift(
records: Sequence[dict[str, Any]],
actual: Sequence[int],
predicted: Sequence[int],
) -> dict[str, Any]:
violations = []
maximum = 0
checked_misroutes = 0
for record, expected_index, predicted_index in zip(records, actual, predicted):
if expected_index == predicted_index:
continue
checked_misroutes += 1
expected = LABELS[expected_index]
got = LABELS[predicted_index]
for lane in ("claude", "codex"):
for level_index, level in enumerate(ROUTING_LEVELS):
drift = abs(
ROUTING_COST_TIERS[expected][lane][level_index]
- ROUTING_COST_TIERS[got][lane][level_index]
)
maximum = max(maximum, drift)
if drift > 1:
violations.append(
{
"promptHash": prompt_hash(record["prompt"]),
"expected": expected,
"predicted": got,
"lane": lane,
"level": level,
"tierDrift": drift,
}
)
return {
"checkedMisroutes": checked_misroutes,
"maximumTierDrift": maximum,
"violations": violations,
"passed": not violations,
}
def evaluate(args: argparse.Namespace) -> dict[str, Any]:
try:
import torch
from transformers import AutoModelForSequenceClassification, AutoTokenizer
except ImportError as exc:
raise DataError(
"evaluation dependencies are missing; install requirements.txt in a virtualenv"
) from exc
synthetic = load_jsonl(args.test)
fixtures = load_classifiable_fixtures(args.fixtures)
records: list[dict[str, Any]] = synthetic + [
{
"prompt": fixture["prompt"],
"purpose": fixture["purpose"],
"slice": "shipped-fixture",
"origin": "shipped-fixture",
}
for fixture in fixtures
]
for index, record in enumerate(records, 1):
if record.get("purpose") not in LABELS:
raise DataError(f"eval record {index}: invalid purpose")
calibration = _load_calibration(args.calibration)
temperature = float(calibration["temperature"])
high_threshold = float(calibration["confidence"]["high"]["minimumScore"])
medium_threshold = float(calibration["confidence"]["medium"]["minimumScore"])
label_to_id = {label: index for index, label in enumerate(LABELS)}
tokenizer = AutoTokenizer.from_pretrained(args.model_dir, local_files_only=True)
onnx_session = None
if args.onnx_model is not None:
try:
import onnxruntime as ort
except ImportError as exc:
raise DataError(
"ONNX evaluation requires onnxruntime from requirements.txt"
) from exc
if args.device not in ("auto", "cpu"):
raise DataError("ONNX evaluation currently measures the CPU provider")
session_options = ort.SessionOptions()
session_options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
onnx_session = ort.InferenceSession(
str(args.onnx_model),
sess_options=session_options,
providers=["CPUExecutionProvider"],
)
onnx_input_names = {item.name for item in onnx_session.get_inputs()}
device = torch.device("cpu")
model = None
runtime_name = "onnxruntime-cpu"
else:
device = _device(torch, args.device)
model = AutoModelForSequenceClassification.from_pretrained(
args.model_dir, local_files_only=True
).to(device)
model.eval()
onnx_input_names = set()
runtime_name = "pytorch"
def predict_logits(encoded: dict[str, Any]) -> Any:
if onnx_session is not None:
inputs = {
key: value.numpy()
for key, value in encoded.items()
if key in onnx_input_names
}
return torch.from_numpy(onnx_session.run(["logits"], inputs)[0])
moved = {key: value.to(device) for key, value in encoded.items()}
return model(**moved).logits.cpu()
actual: list[int] = []
predicted: list[int] = []
probabilities: list[float] = []
confidences: list[str] = []
margins: list[float] = []
inference_batch_size = 1 if onnx_session is not None else args.batch_size
with torch.inference_mode():
for start in range(0, len(records), inference_batch_size):
batch = records[start : start + inference_batch_size]
encoded = encode_fixed_shape(
tokenizer,
[record["prompt"] for record in batch],
torch,
)
logits = predict_logits(encoded) / temperature
distribution = torch.softmax(logits, dim=-1)
top = torch.topk(distribution, k=2, dim=-1)
batch_probabilities = top.values[:, 0].tolist()
batch_margins = (top.values[:, 0] - top.values[:, 1]).tolist()
batch_predictions = top.indices[:, 0].tolist()
for record, probability, margin, prediction in zip(
batch, batch_probabilities, batch_margins, batch_predictions
):
score = confidence_score(probability, margin)
confidence = (
"high"
if score >= high_threshold
else "medium"
if score >= medium_threshold
else "low"
)
actual.append(label_to_id[record["purpose"]])
predicted.append(prediction)
probabilities.append(probability)
margins.append(margin)
confidences.append(confidence)
metrics = classification_metrics(actual, predicted)
correctness = [want == got for want, got in zip(actual, predicted)]
hard_indexes = [
index
for index, record in enumerate(records)
if record.get("slice") in HARD_SLICES
]
hard_metrics = classification_metrics(
[actual[index] for index in hard_indexes],
[predicted[index] for index in hard_indexes],
)
scored_indexes = [
index
for index, record in enumerate(records)
if record.get("slice") != "vague-eval"
]
scored_metrics = classification_metrics(
[actual[index] for index in scored_indexes],
[predicted[index] for index in scored_indexes],
)
scored_hard_indexes = [
index
for index in hard_indexes
if records[index].get("slice") != "vague-eval"
]
scored_hard_metrics = classification_metrics(
[actual[index] for index in scored_hard_indexes],
[predicted[index] for index in scored_hard_indexes],
)
fixture_indexes = [
index
for index, record in enumerate(records)
if record.get("origin") == "shipped-fixture"
]
fixture_metrics = classification_metrics(
[actual[index] for index in fixture_indexes],
[predicted[index] for index in fixture_indexes],
)
accepted_indexes = [
index for index, confidence in enumerate(confidences) if confidence != "low"
]
accepted_precision = (
sum(correctness[index] for index in accepted_indexes) / len(accepted_indexes)
if accepted_indexes
else 1.0
)
def subset_report(indexes: Sequence[int]) -> dict[str, Any]:
subset_correct = [correctness[index] for index in indexes]
subset_accepted = [
index for index in indexes if confidences[index] != "low"
]
return {
**classification_metrics(
[actual[index] for index in indexes],
[predicted[index] for index in indexes],
),
"confidenceCounts": dict(
sorted(Counter(confidences[index] for index in indexes).items())
),
"acceptedPrecision": (
sum(correctness[index] for index in subset_accepted)
/ len(subset_accepted)
if subset_accepted
else 1.0
),
"acceptedCoverage": len(subset_accepted) / len(indexes),
"meanTopProbability": statistics.mean(
probabilities[index] for index in indexes
),
"meanTopTwoMargin": statistics.mean(margins[index] for index in indexes),
"expectedCalibrationError": expected_calibration_error(
[probabilities[index] for index in indexes],
subset_correct,
),
}
slice_reports = {}
for slice_name in sorted({record.get("slice", "unknown") for record in records}):
indexes = [
index
for index, record in enumerate(records)
if record.get("slice", "unknown") == slice_name
]
slice_reports[slice_name] = subset_report(indexes)
misclassifications = [
{
"promptHash": prompt_hash(records[index]["prompt"]),
"slice": records[index].get("slice", "unknown"),
"expected": LABELS[actual[index]],
"predicted": LABELS[predicted[index]],
"confidence": confidences[index],
"topProbability": round(probabilities[index], 6),
"topTwoMargin": round(margins[index], 6),
}
for index in range(len(records))
if not correctness[index]
]
latency_samples: list[float] = []
latency_records = records[: args.latency_samples]
if latency_records:
with torch.inference_mode():
for record in latency_records[: min(5, len(latency_records))]:
encoded = encode_fixed_shape(
tokenizer,
[record["prompt"]],
torch,
)
predict_logits(encoded)
_synchronize(torch, device)
for record in latency_records:
started = time.perf_counter()
encoded = encode_fixed_shape(
tokenizer,
[record["prompt"]],
torch,
)
predict_logits(encoded)
_synchronize(torch, device)
latency_samples.append((time.perf_counter() - started) * 1000)
tier_drift = routing_tier_drift(records, actual, predicted)
report = {
"modelVersion": calibration.get("modelVersion", args.model_dir.name),
"device": str(device),
"runtime": runtime_name,
"artifact": str(args.onnx_model or args.model_dir),
"fixedInputShape": [1, MAX_LENGTH],
"overall": metrics,
"scoredClassification": scored_metrics,
"hardSlice": hard_metrics,
"scoredHardSlice": scored_hard_metrics,
"shippedFixtures": fixture_metrics,
"bySlice": slice_reports,
"calibration": {
"temperature": temperature,
"expectedCalibrationError": expected_calibration_error(
probabilities, correctness
),
"confidenceCounts": dict(sorted(Counter(confidences).items())),
"acceptedPrecision": accepted_precision,
"acceptedCoverage": len(accepted_indexes) / len(records),
},
"latencyMilliseconds": {
"samples": len(latency_samples),
"median": statistics.median(latency_samples) if latency_samples else 0.0,
"p95": _percentile(latency_samples, 0.95),
},
"routingTierDrift": tier_drift,
"gates": {
"scoredAccuracyAtLeast95Percent": scored_metrics["accuracy"] >= 0.95,
"everyPurposeRecallAtLeast85Percent": min(
scored_metrics["perPurposeRecall"].values()
)
>= 0.85,
"vagueEvalLowConfidenceAtLeast90Percent": (
slice_reports.get("vague-eval", {}).get("confidenceCounts", {}).get(
"low", 0
)
/ max(1, slice_reports.get("vague-eval", {}).get("records", 0))
>= 0.90
),
"latencyP95AtMost20Milliseconds": (
not latency_samples or _percentile(latency_samples, 0.95) <= 20.0
),
"misroutesStayWithinOneCostTier": tier_drift["passed"],
},
"misclassifications": misclassifications,
}
write_json(args.report, report)
return report
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--model-dir", type=Path, default=DEFAULT_MODEL_DIR)
parser.add_argument(
"--onnx-model",
type=Path,
help="score a fixed-shape ONNX artifact instead of the PyTorch checkpoint",
)
parser.add_argument("--calibration", type=Path, default=DEFAULT_CALIBRATION)
parser.add_argument("--test", type=Path, default=DEFAULT_TEST)
parser.add_argument("--fixtures", type=Path, default=DEFAULT_FIXTURES)
parser.add_argument("--report", type=Path, default=DEFAULT_REPORT)
parser.add_argument("--device", default="auto")
parser.add_argument("--batch-size", type=int, default=64)
parser.add_argument("--latency-samples", type=int, default=100)
parser.add_argument(
"--no-gate",
action="store_true",
help="write metrics without returning failure when rollout gates miss",
)
return parser
def main(argv: Sequence[str] | None = None) -> int:
parser = build_parser()
args = parser.parse_args(argv)
if args.batch_size <= 0 or args.latency_samples < 0:
parser.error("batch size must be positive and latency samples non-negative")
try:
report = evaluate(args)
except (DataError, OSError, ValueError) as exc:
print(f"error: {exc}", file=sys.stderr)
return 1
print(
f"Frozen eval: scored={report['scoredClassification']['accuracy']:.4%}, "
f"scored-hard={report['scoredHardSlice']['accuracy']:.4%}, "
f"p95={report['latencyMilliseconds']['p95']:.2f} ms."
)
if not args.no_gate and not all(report["gates"].values()):
return 1
return 0
if __name__ == "__main__":
raise SystemExit(main())