Files
nucleic-purpose-classifier/eval.py
T

293 lines
11 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,
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())