#!/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, normalize_prompt, write_json, ) from train import ( MAX_LENGTH, classification_metrics, confidence_score, 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" 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 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)} device = _device(torch, args.device) tokenizer = AutoTokenizer.from_pretrained(args.model_dir, local_files_only=True) model = AutoModelForSequenceClassification.from_pretrained( args.model_dir, local_files_only=True ).to(device) model.eval() actual: list[int] = [] predicted: list[int] = [] probabilities: list[float] = [] confidences: list[str] = [] with torch.inference_mode(): for start in range(0, len(records), args.batch_size): batch = records[start : start + args.batch_size] encoded = tokenizer( [normalize_prompt(record["prompt"]) for record in batch], padding="max_length", truncation=True, max_length=MAX_LENGTH, return_tensors="pt", ) encoded = {key: value.to(device) for key, value in encoded.items()} logits = model(**encoded).logits.cpu() / 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) 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], ) 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 ) 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 = tokenizer( normalize_prompt(record["prompt"]), padding="max_length", truncation=True, max_length=MAX_LENGTH, return_tensors="pt", ) model(**{key: value.to(device) for key, value in encoded.items()}) _synchronize(torch, device) for record in latency_records: started = time.perf_counter() encoded = tokenizer( normalize_prompt(record["prompt"]), padding="max_length", truncation=True, max_length=MAX_LENGTH, return_tensors="pt", ) model(**{key: value.to(device) for key, value in encoded.items()}) _synchronize(torch, device) latency_samples.append((time.perf_counter() - started) * 1000) report = { "modelVersion": calibration.get("modelVersion", args.model_dir.name), "device": str(device), "fixedInputShape": [1, MAX_LENGTH], "overall": metrics, "hardSlice": hard_metrics, "shippedFixtures": fixture_metrics, "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), }, "gates": { "accuracyAtLeast95Percent": metrics["accuracy"] >= 0.95, "everyPurposeRecallAtLeast85Percent": min( metrics["perPurposeRecall"].values() ) >= 0.85, "latencyP95AtMost20Milliseconds": ( not latency_samples or _percentile(latency_samples, 0.95) <= 20.0 ), }, } 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("--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: accuracy={report['overall']['accuracy']:.4%}, " f"hard={report['hardSlice']['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())