Merge nucleic/lucid-north-quail-rnvt into dev
This commit is contained in:
@@ -0,0 +1,292 @@
|
||||
#!/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())
|
||||
Reference in New Issue
Block a user