#!/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 prediction_agreement( records: Sequence[dict[str, Any]], actual: Sequence[int], reference: Sequence[int], candidate: Sequence[int], ) -> dict[str, Any]: """Report label drift from a reference runtime to a candidate artifact.""" if not ( len(records) == len(actual) == len(reference) == len(candidate) and records ): raise ValueError("prediction agreement needs equally sized, non-empty vectors") agreements = [ expected == got for expected, got in zip(reference, candidate) ] scorable_indexes = [ index for index, record in enumerate(records) if record.get("slice") != "vague-eval" ] transitions = Counter() disagreements = [] for index, agrees in enumerate(agreements): if agrees: continue reference_correct = reference[index] == actual[index] candidate_correct = candidate[index] == actual[index] if reference_correct and not candidate_correct: transition = "correctToIncorrect" elif not reference_correct and candidate_correct: transition = "incorrectToCorrect" else: transition = "differentIncorrectLabel" transitions[transition] += 1 disagreements.append( { "promptHash": prompt_hash(records[index]["prompt"]), "slice": records[index].get("slice", "unknown"), "expected": LABELS[actual[index]], "reference": LABELS[reference[index]], "candidate": LABELS[candidate[index]], "transition": transition, } ) by_slice = {} 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 ] by_slice[slice_name] = { "records": len(indexes), "labelAgreement": ( sum(agreements[index] for index in indexes) / len(indexes) ), } return { "records": len(records), "labelAgreement": sum(agreements) / len(agreements), "scoredLabelAgreement": ( sum(agreements[index] for index in scorable_indexes) / len(scorable_indexes) ), "transitionCounts": dict(sorted(transitions.items())), "bySlice": by_slice, "disagreements": disagreements, } 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 reference_model = 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" if args.compare_pytorch: reference_model = AutoModelForSequenceClassification.from_pretrained( args.model_dir, local_files_only=True, ).to(device) reference_model.eval() else: if args.compare_pytorch: raise DataError("--compare-pytorch requires --onnx-model") 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] = [] reference_predictions: list[int] = [] 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 if reference_model is not None: reference_logits = reference_model(**encoded).logits reference_predictions.extend( reference_logits.argmax(dim=-1).tolist() ) 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, } if reference_model is not None: report["pytorchParity"] = prediction_agreement( records, actual, reference_predictions, predicted, ) 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( "--compare-pytorch", action="store_true", help="include label-level drift from --model-dir when scoring ONNX", ) 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())