Merge nucleic/lucid-north-quail-rnvt into dev
This commit is contained in:
@@ -0,0 +1,5 @@
|
||||
.artifacts/
|
||||
.venv/
|
||||
outputs/
|
||||
__pycache__/
|
||||
tests/__pycache__/
|
||||
@@ -0,0 +1,94 @@
|
||||
# Purpose classifier
|
||||
|
||||
This directory is the reproducible data and training pipeline for
|
||||
`docs/PURPOSE_CLASSIFIER.md`. The current slice covers work item 2 and the first part of
|
||||
work item 3: deterministic curation/splitting, a frozen v1 eval set, MiniLM fine-tuning,
|
||||
temperature calibration, shared confidence thresholds, and the frozen-set accuracy,
|
||||
recall, hard-slice, calibration, and latency report.
|
||||
|
||||
## Data contract
|
||||
|
||||
The canonical generated sources are listed in `data/generation-manifest.json`.
|
||||
`round2-NN.jsonl` files are retained generation batches and intentionally duplicate
|
||||
`purpose-prompts-round2.jsonl`; they are provenance, not additional training input.
|
||||
|
||||
`prepare_data.py`:
|
||||
|
||||
- validates the strict generated-record schema;
|
||||
- removes exact and high-overlap word-trigram duplicates;
|
||||
- fails for review if a high-overlap pair has conflicting labels;
|
||||
- keeps shipped fixtures completely outside source data;
|
||||
- holds every `vague-eval` record out of training;
|
||||
- stratifies by primary purpose, slice, and primary language; and
|
||||
- verifies that the deterministic test partition still matches the versioned
|
||||
`data/frozen-test-v1.jsonl`.
|
||||
|
||||
The frozen test set is the synthetic JSONL plus the 87 classifiable records in
|
||||
`Tests/NucleicCoreTests/Fixtures/purpose-prompts.json`. The fixture file's five `general`
|
||||
records are excluded because `general` is deliberately not a model label. The exact
|
||||
membership and hashes are locked in `data/dataset-v1-manifest.json`.
|
||||
|
||||
## Prepare
|
||||
|
||||
From the repository root:
|
||||
|
||||
```bash
|
||||
python3 ml/purpose-classifier/validate-data.py
|
||||
python3 ml/purpose-classifier/prepare_data.py
|
||||
python3 -m unittest discover -s ml/purpose-classifier/tests -p 'test_*.py'
|
||||
```
|
||||
|
||||
For a newly generated raw 200-record batch, enable batch-shape checks explicitly with
|
||||
`validate-data.py path/to/batch.jsonl --batch-size 200 --expected-total 200`.
|
||||
|
||||
The generated train/validation copies land under `.artifacts/dataset-v1/` and are
|
||||
gitignored. A source, curation, seed, or split-policy change that moves the frozen test
|
||||
set fails closed. After reviewing such a change, intentionally version it with:
|
||||
|
||||
```bash
|
||||
python3 ml/purpose-classifier/prepare_data.py --refresh-frozen-test
|
||||
```
|
||||
|
||||
## Train purpose-lite
|
||||
|
||||
Use a dedicated virtual environment. The base model is pinned to a specific
|
||||
`sentence-transformers/all-MiniLM-L6-v2` commit: a 6-layer, 384-dimensional encoder. The
|
||||
training collator always pads/truncates to 128 tokens so the later ONNX/Core ML export
|
||||
can expose a fixed `1 x 128` runtime shape.
|
||||
|
||||
```bash
|
||||
python3 -m venv ml/purpose-classifier/.venv
|
||||
ml/purpose-classifier/.venv/bin/pip install -r ml/purpose-classifier/requirements.txt
|
||||
ml/purpose-classifier/.venv/bin/python ml/purpose-classifier/train.py
|
||||
```
|
||||
|
||||
The default requirements use PyTorch's CPU-only wheel on Linux, avoiding an accidental
|
||||
multi-gigabyte CUDA install in CI and development containers. For NVIDIA, AMD, or Intel
|
||||
accelerator training, install the platform's `torch==2.13.0` build using PyTorch's
|
||||
platform selector, then install `requirements-base.txt`.
|
||||
|
||||
Training writes a local checkpoint, `calibration.json`, and `metrics.json` under
|
||||
`outputs/purpose-lite-v1/`. It fits one validation-only temperature and derives nested
|
||||
HIGH/MEDIUM/LOW cutoffs from calibrated top-one probability plus top-two margin.
|
||||
|
||||
For a wiring smoke test, use a small deterministic prefix:
|
||||
|
||||
```bash
|
||||
ml/purpose-classifier/.venv/bin/python ml/purpose-classifier/train.py \
|
||||
--epochs 1 --max-train-records 64 --max-validation-records 64 \
|
||||
--output-dir ml/purpose-classifier/outputs/smoke --overwrite-output
|
||||
```
|
||||
|
||||
## Evaluate
|
||||
|
||||
```bash
|
||||
ml/purpose-classifier/.venv/bin/python ml/purpose-classifier/eval.py
|
||||
```
|
||||
|
||||
The command returns failure unless frozen accuracy is at least 95%, every purpose recall
|
||||
is at least 85%, and measured batch-one p95 is at most 20 ms. Use `--no-gate` only for
|
||||
diagnostic runs. Accelerator residency, ONNX export/quantization, tokenizer golden tests,
|
||||
tier-drift evaluation, and Core ML parity remain follow-on work. Before calling dataset
|
||||
work item 2 complete, also run a semantic embedding duplicate audit and record the planned
|
||||
10% human label spot-check; the current dependency-free word-trigram pass is deliberately
|
||||
conservative.
|
||||
@@ -0,0 +1,147 @@
|
||||
{
|
||||
"curation": {
|
||||
"excludedDuplicates": {
|
||||
"near": 18
|
||||
},
|
||||
"inputRecords": 12214,
|
||||
"nearDuplicateMethod": "word-trigram Jaccard after SimHash LSH candidate search",
|
||||
"nearDuplicateThreshold": 0.92,
|
||||
"retainedRecords": 12196,
|
||||
"vagueEvalPolicy": "validation/test only"
|
||||
},
|
||||
"datasetVersion": "purpose-dataset-v1",
|
||||
"frozenEval": {
|
||||
"classifiableShippedFixtureRecords": 87,
|
||||
"excludedGeneralFixtureRecords": 5,
|
||||
"hardSliceDefinition": [
|
||||
"boundary",
|
||||
"mixed",
|
||||
"pasted-context",
|
||||
"vague-eval"
|
||||
],
|
||||
"shippedFixtureRecords": 92,
|
||||
"shippedFixturesPath": "Tests/NucleicCoreTests/Fixtures/purpose-prompts.json",
|
||||
"shippedFixturesSha256": "6c8f0eedce35f9da7a407934009e87737b5e85c7371f6a118982d9fae2ea7aae",
|
||||
"syntheticPath": "ml/purpose-classifier/data/frozen-test-v1.jsonl",
|
||||
"syntheticSha256": "e5a6b501a43ce256092932ab05b0af01ff9137aee42d85b7d58bb4ce65f2943f"
|
||||
},
|
||||
"ratios": {
|
||||
"test": 0.1,
|
||||
"train": 0.8,
|
||||
"validation": 0.1
|
||||
},
|
||||
"schemaVersion": 1,
|
||||
"seed": 3248837105,
|
||||
"sources": [
|
||||
{
|
||||
"path": "ml/purpose-classifier/data/purpose-prompts.jsonl",
|
||||
"records": 9214,
|
||||
"sha256": "3c7f8b496ef2014794a50433a78e3125c0dd4a03dc2a7773b52ee25587ac48e7"
|
||||
},
|
||||
{
|
||||
"path": "ml/purpose-classifier/data/purpose-prompts-round2.jsonl",
|
||||
"records": 3000,
|
||||
"sha256": "0a35eac2549e95b31518cd0eee99b83b03761c26a0fb258f427fe04a632699ed"
|
||||
}
|
||||
],
|
||||
"splits": {
|
||||
"test": {
|
||||
"classifiableFixtureRecords": 87,
|
||||
"distribution": {
|
||||
"language": {
|
||||
"de": 7,
|
||||
"en": 1111,
|
||||
"es": 9,
|
||||
"fr": 5,
|
||||
"ja": 3,
|
||||
"pt": 4,
|
||||
"zh": 3
|
||||
},
|
||||
"purpose": {
|
||||
"backendImpl": 171,
|
||||
"debugging": 149,
|
||||
"frontendImpl": 157,
|
||||
"planning": 132,
|
||||
"quickFix": 140,
|
||||
"refactor": 137,
|
||||
"review": 127,
|
||||
"writing": 129
|
||||
},
|
||||
"slice": {
|
||||
"boundary": 179,
|
||||
"core": 507,
|
||||
"mixed": 80,
|
||||
"pasted-context": 84,
|
||||
"vague-eval": 292
|
||||
}
|
||||
},
|
||||
"hardSyntheticRecords": 635,
|
||||
"logicalRecords": 1229,
|
||||
"sha256": "e5a6b501a43ce256092932ab05b0af01ff9137aee42d85b7d58bb4ce65f2943f",
|
||||
"syntheticRecords": 1142
|
||||
},
|
||||
"train": {
|
||||
"distribution": {
|
||||
"language": {
|
||||
"de": 100,
|
||||
"en": 9326,
|
||||
"es": 100,
|
||||
"fr": 81,
|
||||
"ja": 77,
|
||||
"pt": 70,
|
||||
"zh": 72
|
||||
},
|
||||
"purpose": {
|
||||
"backendImpl": 1225,
|
||||
"debugging": 1245,
|
||||
"frontendImpl": 1235,
|
||||
"planning": 1223,
|
||||
"quickFix": 1230,
|
||||
"refactor": 1225,
|
||||
"review": 1230,
|
||||
"writing": 1213
|
||||
},
|
||||
"slice": {
|
||||
"boundary": 2065,
|
||||
"core": 5820,
|
||||
"mixed": 906,
|
||||
"pasted-context": 1035
|
||||
}
|
||||
},
|
||||
"records": 9826,
|
||||
"sha256": "ca0ec9fe6fef6bc4013bf09bdfed5877327f69c98f8e5552c082cbb4e76d9b7d"
|
||||
},
|
||||
"validation": {
|
||||
"distribution": {
|
||||
"language": {
|
||||
"de": 8,
|
||||
"en": 1178,
|
||||
"es": 12,
|
||||
"fr": 7,
|
||||
"ja": 8,
|
||||
"pt": 8,
|
||||
"zh": 7
|
||||
},
|
||||
"purpose": {
|
||||
"backendImpl": 184,
|
||||
"debugging": 157,
|
||||
"frontendImpl": 168,
|
||||
"planning": 141,
|
||||
"quickFix": 151,
|
||||
"refactor": 148,
|
||||
"review": 140,
|
||||
"writing": 139
|
||||
},
|
||||
"slice": {
|
||||
"boundary": 196,
|
||||
"core": 541,
|
||||
"mixed": 83,
|
||||
"pasted-context": 95,
|
||||
"vague-eval": 313
|
||||
}
|
||||
},
|
||||
"records": 1228,
|
||||
"sha256": "7eca32f98105eec1c1a34c324851d52004d8dc720ca5853a1e91fa0a0dd620f4"
|
||||
}
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,57 @@
|
||||
{
|
||||
"schemaVersion": 1,
|
||||
"dataset": "purpose-classifier-source-v1",
|
||||
"canonicalFiles": [
|
||||
"purpose-prompts.jsonl",
|
||||
"purpose-prompts-round2.jsonl"
|
||||
],
|
||||
"derivedBatchFiles": [
|
||||
"round2-01.jsonl",
|
||||
"round2-02.jsonl",
|
||||
"round2-03.jsonl",
|
||||
"round2-04.jsonl",
|
||||
"round2-05.jsonl",
|
||||
"round2-06.jsonl",
|
||||
"round2-07.jsonl",
|
||||
"round2-08.jsonl",
|
||||
"round2-09.jsonl",
|
||||
"round2-10.jsonl",
|
||||
"round2-11.jsonl",
|
||||
"round2-12.jsonl",
|
||||
"round2-13.jsonl",
|
||||
"round2-14.jsonl",
|
||||
"round2-15.jsonl"
|
||||
],
|
||||
"generations": [
|
||||
{
|
||||
"file": "purpose-prompts.jsonl",
|
||||
"model": "mixed frontier-model runs (legacy sol/opus aliases; exact model IDs were not retained)",
|
||||
"date": "2026-07-29",
|
||||
"prompt": "../datagen-prompt.md",
|
||||
"topics": [
|
||||
"web and mobile",
|
||||
"backend and data",
|
||||
"infrastructure and systems",
|
||||
"developer tooling"
|
||||
],
|
||||
"notes": "The canonical aggregate was curated from the original per-model batches. Three exact overlaps with the shipped eval fixture were removed on 2026-07-30."
|
||||
},
|
||||
{
|
||||
"file": "purpose-prompts-round2.jsonl",
|
||||
"model": "Nucleic frontier-model generator (exact underlying model ID was not retained)",
|
||||
"date": "2026-07-30",
|
||||
"prompt": "../datagen-prompt-2.md",
|
||||
"topics": [
|
||||
"boundary confusion pairs",
|
||||
"pasted context",
|
||||
"mixed intent",
|
||||
"non-English developer prompts"
|
||||
],
|
||||
"notes": "Corrective generation that counterbalances round-one label, slice, length, opener, and language drift."
|
||||
}
|
||||
],
|
||||
"limitations": [
|
||||
"The exact generator model IDs and sampling parameters were not recorded when the source corpora were created.",
|
||||
"The round2-NN files are retained generation batches and duplicate the round-two canonical aggregate; dataset tooling must read canonicalFiles only."
|
||||
]
|
||||
}
|
||||
@@ -22,7 +22,6 @@
|
||||
{"prompt":"feature flag `new_billing_summary` is still gated to staff only, open it to everyone","purpose":"quickFix","secondary":null,"mixed":false,"difficulty":0.15,"slice":"core","lang":"en"}
|
||||
{"prompt":"dedupe the three copies of retryWithBackoff","purpose":"refactor","secondary":null,"mixed":false,"difficulty":0.35,"slice":"core","lang":"en"}
|
||||
{"prompt":"the /search endpoint got 4x slower after we shipped last thursday and nothing in that diff touches search. p99 went 180ms -> 750ms. where do i even start","purpose":"debugging","secondary":null,"mixed":false,"difficulty":0.75,"slice":"boundary","lang":"en"}
|
||||
{"prompt":"make it pop","purpose":"frontendImpl","secondary":null,"mixed":false,"difficulty":0.3,"slice":"vague-eval","lang":"en"}
|
||||
{"prompt":"is our jwt refresh flow safe against replay if someone grabs a refresh token off a stolen device? read through auth/refresh.go and tell me what you think, don't change anything yet","purpose":"review","secondary":null,"mixed":false,"difficulty":0.65,"slice":"core","lang":"en"}
|
||||
{"prompt":"compare our current celery setup against just using postgres SKIP LOCKED for the job queue. we do maybe 200 jobs/min, mostly short","purpose":"review","secondary":null,"mixed":false,"difficulty":0.6,"slice":"core","lang":"en"}
|
||||
{"prompt":"explain what pkg/ledger/reconcile.go actually does, line by line if you have to. inherited it and nobody left knows","purpose":"review","secondary":null,"mixed":false,"difficulty":0.55,"slice":"boundary","lang":"en"}
|
||||
@@ -72,7 +71,6 @@
|
||||
{"prompt":"review the diff on my branch before i open the PR, specifically the transaction boundaries in the transfer path","purpose":"review","secondary":null,"mixed":false,"difficulty":0.6,"slice":"core","lang":"en"}
|
||||
{"prompt":"what's the difference between our useSyncedQuery hook and just using react-query's useQuery with our fetcher? feels like we reimplemented it","purpose":"review","secondary":null,"mixed":false,"difficulty":0.45,"slice":"core","lang":"en"}
|
||||
{"prompt":"document the event payload schemas for the pubsub topics in docs/events.md AND add the missing JSON schema files under schemas/ so we can validate in CI","purpose":"writing","secondary":"backendImpl","mixed":true,"difficulty":0.55,"slice":"mixed","lang":"en"}
|
||||
{"prompt":"continue","purpose":"backendImpl","secondary":null,"mixed":false,"difficulty":0.4,"slice":"vague-eval","lang":"en"}
|
||||
{"prompt":"figure out why the checkout total is off by a cent for some carts and then write up the postmortem, we told the customer we'd have both by friday","purpose":"debugging","secondary":"writing","mixed":true,"difficulty":0.7,"slice":"mixed","lang":"en"}
|
||||
{"prompt":"picking the payment-flow bug back up from yesterday — the 3DS redirect lands on a blank page maybe a third of the time in safari. no console error, network tab shows the POST to /confirm returning 302 and then nothing","purpose":"debugging","secondary":null,"mixed":false,"difficulty":0.75,"slice":"core","lang":"en"}
|
||||
{"prompt":"come up with the migration plan for moving our 4TB mysql instance to aurora with under 15 min of write downtime, and list the go/no-go checks","purpose":"planning","secondary":null,"mixed":false,"difficulty":0.9,"slice":"core","lang":"en"}
|
||||
@@ -99,7 +97,6 @@
|
||||
{"prompt":"internal blog post about the latency work we did last quarter. audience is other engineers here, ~800 words, i can give you the numbers","purpose":"writing","secondary":null,"mixed":false,"difficulty":0.45,"slice":"core","lang":"en"}
|
||||
{"prompt":"runbook for the on-call rotation covering the four alerts that actually page us","purpose":"writing","secondary":null,"mixed":false,"difficulty":0.5,"slice":"core","lang":"en"}
|
||||
{"prompt":"the CONTRIBUTING.md is 3 lines. write a real one — branch naming, how to run the test suite, what we expect in a PR, and the codegen step people always forget","purpose":"writing","secondary":null,"mixed":false,"difficulty":0.4,"slice":"core","lang":"en"}
|
||||
{"prompt":"do the thing we discussed","purpose":"backendImpl","secondary":null,"mixed":false,"difficulty":0.4,"slice":"vague-eval","lang":"en"}
|
||||
{"prompt":"spec out the multiplayer lobby then get the netcode skeleton in — matchmaking rules first as a doc, then the actual go server with the room state machine","purpose":"planning","secondary":"backendImpl","mixed":true,"difficulty":0.85,"slice":"mixed","lang":"en"}
|
||||
{"prompt":"evaluate whether we should adopt bazel. monorepo, ~60 services, mixed go/ts/python, current builds are makefiles and 14 min of CI. i want a recommendation with a phased rollout, not a yes/no","purpose":"planning","secondary":null,"mixed":false,"difficulty":0.85,"slice":"core","lang":"en"}
|
||||
{"prompt":"what's a reasonable retention + partitioning strategy for the raw telemetry table? we ingest ~90GB/day and only ever query the last 14 days interactively","purpose":"planning","secondary":null,"mixed":false,"difficulty":0.7,"slice":"core","lang":"en"}
|
||||
|
||||
+1
-1
@@ -2,7 +2,7 @@
|
||||
|
||||
The prompt below is fed verbatim to a frontier-model agent to produce training/eval data
|
||||
per docs/PURPOSE_CLASSIFIER.md §4.1. Record the generating model, date, and batch topics
|
||||
in the generation manifest alongside the output. The 82 shipped fixtures
|
||||
in the generation manifest alongside the output. The 92 shipped fixtures
|
||||
(Tests/NucleicCoreTests/Fixtures/purpose-prompts.json) are eval-only and must NOT be
|
||||
pasted into the generator's context (contamination).
|
||||
|
||||
|
||||
@@ -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())
|
||||
+325
@@ -0,0 +1,325 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Curate source prompts and build deterministic purpose-classifier splits."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import hashlib
|
||||
import json
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Any, Sequence
|
||||
|
||||
from purpose_data import (
|
||||
HARD_SLICES,
|
||||
LABELS,
|
||||
DataError,
|
||||
SourceRecord,
|
||||
curate_records,
|
||||
distribution,
|
||||
file_sha256,
|
||||
jsonl_bytes,
|
||||
load_classifiable_fixtures,
|
||||
load_sources,
|
||||
prompt_hash,
|
||||
split_records,
|
||||
write_json,
|
||||
write_jsonl,
|
||||
)
|
||||
|
||||
|
||||
SCRIPT_DIR = Path(__file__).resolve().parent
|
||||
REPOSITORY_ROOT = SCRIPT_DIR.parent.parent
|
||||
DATA_DIR = SCRIPT_DIR / "data"
|
||||
GENERATION_MANIFEST = DATA_DIR / "generation-manifest.json"
|
||||
DEFAULT_FIXTURES = (
|
||||
REPOSITORY_ROOT
|
||||
/ "Tests"
|
||||
/ "NucleicCoreTests"
|
||||
/ "Fixtures"
|
||||
/ "purpose-prompts.json"
|
||||
)
|
||||
DEFAULT_OUTPUT_DIR = SCRIPT_DIR / ".artifacts" / "dataset-v1"
|
||||
DEFAULT_FROZEN_TEST = DATA_DIR / "frozen-test-v1.jsonl"
|
||||
DEFAULT_SPLIT_MANIFEST = DATA_DIR / "dataset-v1-manifest.json"
|
||||
DATASET_VERSION = "purpose-dataset-v1"
|
||||
DEFAULT_SEED = 0xC1A551F1
|
||||
|
||||
|
||||
def _relative(path: Path) -> str:
|
||||
try:
|
||||
return str(path.resolve().relative_to(REPOSITORY_ROOT))
|
||||
except ValueError:
|
||||
return str(path.resolve())
|
||||
|
||||
|
||||
def default_source_paths() -> list[Path]:
|
||||
try:
|
||||
value = json.loads(GENERATION_MANIFEST.read_text(encoding="utf-8"))
|
||||
names = value["canonicalFiles"]
|
||||
except (OSError, UnicodeError, json.JSONDecodeError, KeyError, TypeError) as exc:
|
||||
raise DataError(f"{GENERATION_MANIFEST}: cannot read canonicalFiles: {exc}") from exc
|
||||
if not isinstance(names, list) or not names or not all(
|
||||
isinstance(name, str) and name for name in names
|
||||
):
|
||||
raise DataError(
|
||||
f"{GENERATION_MANIFEST}: canonicalFiles must be a non-empty string array"
|
||||
)
|
||||
return [(DATA_DIR / name).resolve() for name in names]
|
||||
|
||||
|
||||
def _fixture_counts(path: Path) -> tuple[int, int]:
|
||||
try:
|
||||
value = json.loads(path.read_text(encoding="utf-8"))
|
||||
except (OSError, UnicodeError, json.JSONDecodeError) as exc:
|
||||
raise DataError(f"{path}: cannot read fixtures: {exc}") from exc
|
||||
if not isinstance(value, list):
|
||||
raise DataError(f"{path}: fixture root must be an array")
|
||||
classifiable = sum(
|
||||
isinstance(item, dict) and item.get("purpose") in LABELS for item in value
|
||||
)
|
||||
return len(value), classifiable
|
||||
|
||||
|
||||
def _sha256_bytes(value: bytes) -> str:
|
||||
return hashlib.sha256(value).hexdigest()
|
||||
|
||||
|
||||
def _records(values: Sequence[SourceRecord]) -> list[dict[str, Any]]:
|
||||
return [record.value for record in values]
|
||||
|
||||
|
||||
def _build_manifest(
|
||||
*,
|
||||
sources: Sequence[Path],
|
||||
source_record_count: int,
|
||||
curated_record_count: int,
|
||||
duplicate_counts: dict[str, int],
|
||||
splits: Any,
|
||||
fixtures_path: Path,
|
||||
total_fixture_count: int,
|
||||
classifiable_fixture_count: int,
|
||||
output_hashes: dict[str, str],
|
||||
seed: int,
|
||||
near_duplicate_threshold: float,
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"schemaVersion": 1,
|
||||
"datasetVersion": DATASET_VERSION,
|
||||
"seed": seed,
|
||||
"ratios": {"train": 0.8, "validation": 0.1, "test": 0.1},
|
||||
"sources": [
|
||||
{
|
||||
"path": _relative(path),
|
||||
"sha256": file_sha256(path),
|
||||
"records": len(load_sources([path])),
|
||||
}
|
||||
for path in sources
|
||||
],
|
||||
"curation": {
|
||||
"inputRecords": source_record_count,
|
||||
"retainedRecords": curated_record_count,
|
||||
"excludedDuplicates": duplicate_counts,
|
||||
"nearDuplicateMethod": "word-trigram Jaccard after SimHash LSH candidate search",
|
||||
"nearDuplicateThreshold": near_duplicate_threshold,
|
||||
"vagueEvalPolicy": "validation/test only",
|
||||
},
|
||||
"frozenEval": {
|
||||
"syntheticPath": _relative(DEFAULT_FROZEN_TEST),
|
||||
"syntheticSha256": output_hashes["test"],
|
||||
"shippedFixturesPath": _relative(fixtures_path),
|
||||
"shippedFixturesSha256": file_sha256(fixtures_path),
|
||||
"shippedFixtureRecords": total_fixture_count,
|
||||
"classifiableShippedFixtureRecords": classifiable_fixture_count,
|
||||
"excludedGeneralFixtureRecords": (
|
||||
total_fixture_count - classifiable_fixture_count
|
||||
),
|
||||
"hardSliceDefinition": sorted(HARD_SLICES),
|
||||
},
|
||||
"splits": {
|
||||
"train": {
|
||||
"records": len(splits.train),
|
||||
"sha256": output_hashes["train"],
|
||||
"distribution": distribution(splits.train),
|
||||
},
|
||||
"validation": {
|
||||
"records": len(splits.validation),
|
||||
"sha256": output_hashes["validation"],
|
||||
"distribution": distribution(splits.validation),
|
||||
},
|
||||
"test": {
|
||||
"syntheticRecords": len(splits.test),
|
||||
"classifiableFixtureRecords": classifiable_fixture_count,
|
||||
"logicalRecords": splits.logical_test_count,
|
||||
"hardSyntheticRecords": sum(
|
||||
record.value["slice"] in HARD_SLICES for record in splits.test
|
||||
),
|
||||
"sha256": output_hashes["test"],
|
||||
"distribution": distribution(splits.test),
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def prepare(
|
||||
*,
|
||||
sources: Sequence[Path],
|
||||
fixtures_path: Path,
|
||||
output_dir: Path,
|
||||
frozen_test_path: Path,
|
||||
manifest_path: Path,
|
||||
refresh_frozen_test: bool,
|
||||
seed: int,
|
||||
near_duplicate_threshold: float,
|
||||
) -> dict[str, Any]:
|
||||
records = load_sources(sources)
|
||||
fixtures = load_classifiable_fixtures(fixtures_path)
|
||||
total_fixture_count, classifiable_fixture_count = _fixture_counts(fixtures_path)
|
||||
if classifiable_fixture_count != len(fixtures):
|
||||
raise DataError("fixture accounting mismatch")
|
||||
|
||||
curated = curate_records(
|
||||
records,
|
||||
fixtures,
|
||||
near_duplicate_threshold=near_duplicate_threshold,
|
||||
)
|
||||
splits = split_records(
|
||||
curated.records,
|
||||
fixture_count=classifiable_fixture_count,
|
||||
seed=seed,
|
||||
)
|
||||
|
||||
train_values = _records(splits.train)
|
||||
validation_values = _records(splits.validation)
|
||||
test_values = _records(splits.test)
|
||||
candidate_frozen_test = jsonl_bytes(test_values)
|
||||
if refresh_frozen_test:
|
||||
frozen_test_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
frozen_test_path.write_bytes(candidate_frozen_test)
|
||||
elif not frozen_test_path.exists():
|
||||
raise DataError(
|
||||
f"{frozen_test_path}: frozen test is missing; review the candidate then run "
|
||||
"--refresh-frozen-test"
|
||||
)
|
||||
elif frozen_test_path.read_bytes() != candidate_frozen_test:
|
||||
raise DataError(
|
||||
f"{frozen_test_path}: deterministic test split changed; inspect source/seed "
|
||||
"changes and use --refresh-frozen-test only when intentionally versioning it"
|
||||
)
|
||||
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
write_jsonl(output_dir / "train.jsonl", train_values)
|
||||
write_jsonl(output_dir / "validation.jsonl", validation_values)
|
||||
write_jsonl(output_dir / "test.jsonl", test_values)
|
||||
write_jsonl(
|
||||
output_dir / "exclusions.jsonl",
|
||||
(
|
||||
{
|
||||
"promptHash": prompt_hash(duplicate.dropped.value["prompt"]),
|
||||
"matchedPromptHash": duplicate.matched_prompt_hash,
|
||||
"source": _relative(duplicate.dropped.source),
|
||||
"line": duplicate.dropped.line,
|
||||
"kind": duplicate.kind,
|
||||
"similarity": round(duplicate.similarity, 6),
|
||||
}
|
||||
for duplicate in curated.duplicates
|
||||
),
|
||||
)
|
||||
|
||||
duplicate_counts: dict[str, int] = {}
|
||||
for duplicate in curated.duplicates:
|
||||
duplicate_counts[duplicate.kind] = duplicate_counts.get(duplicate.kind, 0) + 1
|
||||
output_hashes = {
|
||||
"train": _sha256_bytes(jsonl_bytes(train_values)),
|
||||
"validation": _sha256_bytes(jsonl_bytes(validation_values)),
|
||||
"test": _sha256_bytes(candidate_frozen_test),
|
||||
}
|
||||
manifest = _build_manifest(
|
||||
sources=sources,
|
||||
source_record_count=len(records),
|
||||
curated_record_count=len(curated.records),
|
||||
duplicate_counts=dict(sorted(duplicate_counts.items())),
|
||||
splits=splits,
|
||||
fixtures_path=fixtures_path,
|
||||
total_fixture_count=total_fixture_count,
|
||||
classifiable_fixture_count=classifiable_fixture_count,
|
||||
output_hashes=output_hashes,
|
||||
seed=seed,
|
||||
near_duplicate_threshold=near_duplicate_threshold,
|
||||
)
|
||||
# The frozen path can be overridden in tests or experiments.
|
||||
manifest["frozenEval"]["syntheticPath"] = _relative(frozen_test_path)
|
||||
write_json(output_dir / "manifest.json", manifest)
|
||||
if refresh_frozen_test:
|
||||
write_json(manifest_path, manifest)
|
||||
elif manifest_path.exists():
|
||||
existing = json.loads(manifest_path.read_text(encoding="utf-8"))
|
||||
if existing != manifest:
|
||||
raise DataError(
|
||||
f"{manifest_path}: split manifest changed; inspect and refresh the frozen "
|
||||
"test intentionally"
|
||||
)
|
||||
else:
|
||||
raise DataError(
|
||||
f"{manifest_path}: frozen split manifest is missing; use --refresh-frozen-test"
|
||||
)
|
||||
return manifest
|
||||
|
||||
|
||||
def build_parser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument(
|
||||
"--source",
|
||||
action="append",
|
||||
type=Path,
|
||||
help="canonical source JSONL; repeat for multiple files (default: generation manifest)",
|
||||
)
|
||||
parser.add_argument("--fixtures", type=Path, default=DEFAULT_FIXTURES)
|
||||
parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR)
|
||||
parser.add_argument("--frozen-test", type=Path, default=DEFAULT_FROZEN_TEST)
|
||||
parser.add_argument("--manifest", type=Path, default=DEFAULT_SPLIT_MANIFEST)
|
||||
parser.add_argument(
|
||||
"--refresh-frozen-test",
|
||||
action="store_true",
|
||||
help="replace the versioned test split and manifest after intentional review",
|
||||
)
|
||||
parser.add_argument("--seed", type=int, default=DEFAULT_SEED)
|
||||
parser.add_argument("--near-duplicate-threshold", type=float, default=0.92)
|
||||
return parser
|
||||
|
||||
|
||||
def main(argv: Sequence[str] | None = None) -> int:
|
||||
args = build_parser().parse_args(argv)
|
||||
try:
|
||||
sources = (
|
||||
[path.resolve() for path in args.source]
|
||||
if args.source
|
||||
else default_source_paths()
|
||||
)
|
||||
manifest = prepare(
|
||||
sources=sources,
|
||||
fixtures_path=args.fixtures.resolve(),
|
||||
output_dir=args.output_dir.resolve(),
|
||||
frozen_test_path=args.frozen_test.resolve(),
|
||||
manifest_path=args.manifest.resolve(),
|
||||
refresh_frozen_test=args.refresh_frozen_test,
|
||||
seed=args.seed,
|
||||
near_duplicate_threshold=args.near_duplicate_threshold,
|
||||
)
|
||||
except (DataError, OSError, UnicodeError, json.JSONDecodeError) as exc:
|
||||
print(f"error: {exc}", file=sys.stderr)
|
||||
return 1
|
||||
|
||||
splits = manifest["splits"]
|
||||
print(
|
||||
f"Prepared {manifest['curation']['retainedRecords']} curated records: "
|
||||
f"{splits['train']['records']} train, "
|
||||
f"{splits['validation']['records']} validation, "
|
||||
f"{splits['test']['logicalRecords']} frozen test "
|
||||
f"({splits['test']['classifiableFixtureRecords']} shipped fixtures)."
|
||||
)
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
+516
@@ -0,0 +1,516 @@
|
||||
"""Shared data contracts for the purpose-classifier pipeline."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import math
|
||||
import re
|
||||
import unicodedata
|
||||
from collections import Counter, defaultdict
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any, Callable, Iterable, Sequence
|
||||
|
||||
|
||||
LABELS = (
|
||||
"planning",
|
||||
"backendImpl",
|
||||
"frontendImpl",
|
||||
"quickFix",
|
||||
"refactor",
|
||||
"debugging",
|
||||
"review",
|
||||
"writing",
|
||||
)
|
||||
LABEL_SET = frozenset(LABELS)
|
||||
SOURCE_FIELDS = frozenset(
|
||||
{"prompt", "purpose", "secondary", "mixed", "difficulty", "slice", "lang"}
|
||||
)
|
||||
SLICES = frozenset(
|
||||
{"core", "boundary", "mixed", "pasted-context", "vague-eval"}
|
||||
)
|
||||
HARD_SLICES = frozenset({"boundary", "mixed", "pasted-context", "vague-eval"})
|
||||
WORD_RE = re.compile(r"\w+", re.UNICODE)
|
||||
|
||||
|
||||
class DataError(ValueError):
|
||||
"""A deterministic data-contract failure."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SourceRecord:
|
||||
value: dict[str, Any]
|
||||
source: Path
|
||||
line: int
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Duplicate:
|
||||
dropped: SourceRecord
|
||||
matched_prompt_hash: str
|
||||
kind: str
|
||||
similarity: float
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CurationResult:
|
||||
records: list[SourceRecord]
|
||||
duplicates: list[Duplicate]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SplitResult:
|
||||
train: list[SourceRecord]
|
||||
validation: list[SourceRecord]
|
||||
test: list[SourceRecord]
|
||||
fixture_count: int
|
||||
|
||||
@property
|
||||
def logical_test_count(self) -> int:
|
||||
return len(self.test) + self.fixture_count
|
||||
|
||||
|
||||
def normalize_prompt(prompt: str) -> str:
|
||||
"""Match the runtime's whitespace collapse and add stable Unicode normalization."""
|
||||
|
||||
normalized = unicodedata.normalize("NFKC", prompt)
|
||||
return " ".join(normalized.split())
|
||||
|
||||
|
||||
def normalized_key(prompt: str) -> str:
|
||||
return normalize_prompt(prompt).casefold()
|
||||
|
||||
|
||||
def prompt_hash(prompt: str) -> str:
|
||||
return hashlib.sha256(normalized_key(prompt).encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
def file_sha256(path: Path) -> str:
|
||||
digest = hashlib.sha256()
|
||||
with path.open("rb") as handle:
|
||||
for chunk in iter(lambda: handle.read(1024 * 1024), b""):
|
||||
digest.update(chunk)
|
||||
return digest.hexdigest()
|
||||
|
||||
|
||||
def canonical_json(record: dict[str, Any]) -> str:
|
||||
return json.dumps(record, ensure_ascii=False, separators=(",", ":"))
|
||||
|
||||
|
||||
def jsonl_bytes(records: Iterable[dict[str, Any]]) -> bytes:
|
||||
return ("".join(f"{canonical_json(record)}\n" for record in records)).encode("utf-8")
|
||||
|
||||
|
||||
def write_jsonl(path: Path, records: Iterable[dict[str, Any]]) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_bytes(jsonl_bytes(records))
|
||||
|
||||
|
||||
def write_json(path: Path, value: Any) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_text(
|
||||
json.dumps(value, ensure_ascii=False, indent=2, sort_keys=True) + "\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
|
||||
def load_jsonl(path: Path) -> list[dict[str, Any]]:
|
||||
records: list[dict[str, Any]] = []
|
||||
try:
|
||||
lines = path.read_text(encoding="utf-8").splitlines()
|
||||
except (OSError, UnicodeError) as exc:
|
||||
raise DataError(f"{path}: cannot read UTF-8 JSONL: {exc}") from exc
|
||||
|
||||
for line_number, line in enumerate(lines, 1):
|
||||
if not line.strip():
|
||||
raise DataError(f"{path}:{line_number}: blank JSONL line")
|
||||
try:
|
||||
value = json.loads(line)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise DataError(f"{path}:{line_number}: invalid JSON: {exc}") from exc
|
||||
if not isinstance(value, dict):
|
||||
raise DataError(f"{path}:{line_number}: expected a JSON object")
|
||||
records.append(value)
|
||||
return records
|
||||
|
||||
|
||||
def validate_source_record(record: dict[str, Any], location: str) -> None:
|
||||
keys = set(record)
|
||||
if keys != SOURCE_FIELDS:
|
||||
missing = sorted(SOURCE_FIELDS - keys)
|
||||
extra = sorted(keys - SOURCE_FIELDS)
|
||||
details = []
|
||||
if missing:
|
||||
details.append(f"missing {', '.join(missing)}")
|
||||
if extra:
|
||||
details.append(f"unexpected {', '.join(extra)}")
|
||||
raise DataError(f"{location}: invalid fields ({'; '.join(details)})")
|
||||
|
||||
prompt = record["prompt"]
|
||||
if not isinstance(prompt, str) or not prompt.strip():
|
||||
raise DataError(f"{location}: prompt must be a non-empty string")
|
||||
if record["purpose"] not in LABEL_SET:
|
||||
raise DataError(f"{location}: invalid purpose {record['purpose']!r}")
|
||||
|
||||
secondary = record["secondary"]
|
||||
mixed = record["mixed"]
|
||||
if type(mixed) is not bool:
|
||||
raise DataError(f"{location}: mixed must be a boolean")
|
||||
if secondary is not None and secondary not in LABEL_SET:
|
||||
raise DataError(f"{location}: invalid secondary purpose {secondary!r}")
|
||||
if mixed != (secondary is not None):
|
||||
raise DataError(f"{location}: mixed and secondary disagree")
|
||||
if secondary == record["purpose"]:
|
||||
raise DataError(f"{location}: secondary must differ from purpose")
|
||||
|
||||
difficulty = record["difficulty"]
|
||||
if (
|
||||
isinstance(difficulty, bool)
|
||||
or not isinstance(difficulty, (int, float))
|
||||
or not math.isfinite(difficulty)
|
||||
or not 0.0 <= difficulty <= 1.0
|
||||
):
|
||||
raise DataError(f"{location}: difficulty must be a finite value from 0 to 1")
|
||||
if record["slice"] not in SLICES:
|
||||
raise DataError(f"{location}: invalid slice {record['slice']!r}")
|
||||
if (record["slice"] == "mixed") != mixed:
|
||||
raise DataError(f"{location}: the mixed slice and mixed field disagree")
|
||||
if not isinstance(record["lang"], str) or not record["lang"]:
|
||||
raise DataError(f"{location}: lang must be a non-empty string")
|
||||
|
||||
|
||||
def load_sources(paths: Sequence[Path]) -> list[SourceRecord]:
|
||||
records: list[SourceRecord] = []
|
||||
for path in paths:
|
||||
for line, value in enumerate(load_jsonl(path), 1):
|
||||
validate_source_record(value, f"{path}:{line}")
|
||||
records.append(SourceRecord(value=value, source=path, line=line))
|
||||
return records
|
||||
|
||||
|
||||
def load_classifiable_fixtures(path: Path) -> list[dict[str, str]]:
|
||||
try:
|
||||
value = json.loads(path.read_text(encoding="utf-8"))
|
||||
except (OSError, UnicodeError, json.JSONDecodeError) as exc:
|
||||
raise DataError(f"{path}: cannot load fixtures: {exc}") from exc
|
||||
if not isinstance(value, list):
|
||||
raise DataError(f"{path}: fixture root must be an array")
|
||||
|
||||
fixtures: list[dict[str, str]] = []
|
||||
for index, item in enumerate(value):
|
||||
if not isinstance(item, dict):
|
||||
raise DataError(f"{path}: fixture {index} must be an object")
|
||||
prompt = item.get("prompt")
|
||||
purpose = item.get("purpose")
|
||||
if not isinstance(prompt, str) or not prompt.strip():
|
||||
raise DataError(f"{path}: fixture {index} has an invalid prompt")
|
||||
if purpose == "general":
|
||||
continue
|
||||
if purpose not in LABEL_SET:
|
||||
raise DataError(f"{path}: fixture {index} has invalid purpose {purpose!r}")
|
||||
fixtures.append({"prompt": prompt, "purpose": purpose})
|
||||
return fixtures
|
||||
|
||||
|
||||
def _word_shingles(prompt: str) -> frozenset[str]:
|
||||
words = WORD_RE.findall(normalized_key(prompt))
|
||||
if len(words) < 8:
|
||||
return frozenset()
|
||||
return frozenset(
|
||||
"\x1f".join(words[index : index + 3])
|
||||
for index in range(len(words) - 2)
|
||||
)
|
||||
|
||||
|
||||
def _simhash(features: frozenset[str]) -> int:
|
||||
weights = [0] * 64
|
||||
for feature in features:
|
||||
value = int.from_bytes(
|
||||
hashlib.blake2b(feature.encode("utf-8"), digest_size=8).digest(), "big"
|
||||
)
|
||||
for bit in range(64):
|
||||
weights[bit] += 1 if value & (1 << bit) else -1
|
||||
|
||||
signature = 0
|
||||
for bit, weight in enumerate(weights):
|
||||
if weight >= 0:
|
||||
signature |= 1 << bit
|
||||
return signature
|
||||
|
||||
|
||||
class _NearDuplicateIndex:
|
||||
"""Small dependency-free LSH pass for generated template duplicates.
|
||||
|
||||
This is intentionally conservative. Short prompts are handled only by exact-match
|
||||
checks because a one-word change can completely change their purpose. Longer prompts
|
||||
become word-trigram sets. SimHash bands produce candidates; exact Jaccard similarity
|
||||
decides whether a row is dropped.
|
||||
"""
|
||||
|
||||
def __init__(self, threshold: float) -> None:
|
||||
self.threshold = threshold
|
||||
self.items: list[tuple[frozenset[str], str, str]] = []
|
||||
self.buckets: dict[tuple[int, int], list[int]] = defaultdict(list)
|
||||
|
||||
@staticmethod
|
||||
def _bands(signature: int) -> Iterable[tuple[int, int]]:
|
||||
# Requiring two matching 8-bit bands keeps random candidate sets small while
|
||||
# retaining every pair whose SimHashes differ in at most six bands.
|
||||
for band in range(8):
|
||||
yield band, (signature >> (band * 8)) & 0xFF
|
||||
|
||||
def find(self, prompt: str) -> tuple[str, str, float] | None:
|
||||
features = _word_shingles(prompt)
|
||||
if not features:
|
||||
return None
|
||||
signature = _simhash(features)
|
||||
hits: Counter[int] = Counter()
|
||||
for band in self._bands(signature):
|
||||
hits.update(self.buckets.get(band, ()))
|
||||
|
||||
best: tuple[str, str, float] | None = None
|
||||
for index, matching_bands in hits.items():
|
||||
if matching_bands < 2:
|
||||
continue
|
||||
other_features, other_label, other_hash = self.items[index]
|
||||
union = len(features | other_features)
|
||||
similarity = len(features & other_features) / union if union else 1.0
|
||||
if similarity >= self.threshold and (
|
||||
best is None or similarity > best[2]
|
||||
):
|
||||
best = (other_label, other_hash, similarity)
|
||||
return best
|
||||
|
||||
def add(self, prompt: str, label: str) -> None:
|
||||
features = _word_shingles(prompt)
|
||||
if not features:
|
||||
return
|
||||
signature = _simhash(features)
|
||||
index = len(self.items)
|
||||
self.items.append((features, label, prompt_hash(prompt)))
|
||||
for band in self._bands(signature):
|
||||
self.buckets[band].append(index)
|
||||
|
||||
|
||||
def curate_records(
|
||||
records: Sequence[SourceRecord],
|
||||
fixtures: Sequence[dict[str, str]],
|
||||
*,
|
||||
near_duplicate_threshold: float = 0.92,
|
||||
) -> CurationResult:
|
||||
if not 0.0 < near_duplicate_threshold <= 1.0:
|
||||
raise DataError("near-duplicate threshold must be in (0, 1]")
|
||||
|
||||
exact: dict[str, tuple[str, str]] = {}
|
||||
near = _NearDuplicateIndex(near_duplicate_threshold)
|
||||
for fixture in fixtures:
|
||||
key = normalized_key(fixture["prompt"])
|
||||
previous = exact.get(key)
|
||||
if previous is not None and previous[0] != fixture["purpose"]:
|
||||
raise DataError("shipped fixtures contain an exact prompt with two labels")
|
||||
exact[key] = (fixture["purpose"], prompt_hash(fixture["prompt"]))
|
||||
near.add(fixture["prompt"], fixture["purpose"])
|
||||
|
||||
kept: list[SourceRecord] = []
|
||||
duplicates: list[Duplicate] = []
|
||||
conflicts: list[str] = []
|
||||
for record in records:
|
||||
prompt = record.value["prompt"]
|
||||
purpose = record.value["purpose"]
|
||||
key = normalized_key(prompt)
|
||||
previous = exact.get(key)
|
||||
if previous is not None:
|
||||
previous_label, previous_hash = previous
|
||||
if previous_label != purpose:
|
||||
conflicts.append(
|
||||
f"{record.source}:{record.line}: exact duplicate has labels "
|
||||
f"{previous_label!r} and {purpose!r}"
|
||||
)
|
||||
else:
|
||||
duplicates.append(
|
||||
Duplicate(record, previous_hash, "exact", 1.0)
|
||||
)
|
||||
continue
|
||||
|
||||
match = near.find(prompt)
|
||||
if match is not None:
|
||||
previous_label, previous_hash, similarity = match
|
||||
if previous_label != purpose:
|
||||
conflicts.append(
|
||||
f"{record.source}:{record.line}: {similarity:.1%}-similar prompt "
|
||||
f"has labels {previous_label!r} and {purpose!r}"
|
||||
)
|
||||
# Keep the row for now so all conflicts are reported without causing a
|
||||
# cascade of duplicates against a record that may later be relabeled.
|
||||
exact[key] = (purpose, prompt_hash(prompt))
|
||||
near.add(prompt, purpose)
|
||||
kept.append(record)
|
||||
else:
|
||||
duplicates.append(
|
||||
Duplicate(record, previous_hash, "near", similarity)
|
||||
)
|
||||
continue
|
||||
|
||||
exact[key] = (purpose, prompt_hash(prompt))
|
||||
near.add(prompt, purpose)
|
||||
kept.append(record)
|
||||
|
||||
if conflicts:
|
||||
preview = "\n".join(conflicts[:20])
|
||||
remainder = len(conflicts) - min(20, len(conflicts))
|
||||
suffix = f"\n... {remainder} more conflict(s)" if remainder else ""
|
||||
raise DataError(f"near-duplicate label conflicts require review:\n{preview}{suffix}")
|
||||
return CurationResult(records=kept, duplicates=duplicates)
|
||||
|
||||
|
||||
def _stable_digest(seed: int, prompt: str) -> str:
|
||||
material = f"{seed}\0{normalized_key(prompt)}".encode("utf-8")
|
||||
return hashlib.sha256(material).hexdigest()
|
||||
|
||||
|
||||
def _stratified_order(
|
||||
records: Sequence[SourceRecord],
|
||||
*,
|
||||
seed: int,
|
||||
strata: Callable[[SourceRecord], tuple[str, ...]],
|
||||
) -> list[SourceRecord]:
|
||||
groups: dict[tuple[str, ...], list[SourceRecord]] = defaultdict(list)
|
||||
for record in records:
|
||||
groups[strata(record)].append(record)
|
||||
|
||||
ranked: list[tuple[float, str, SourceRecord]] = []
|
||||
for key in sorted(groups):
|
||||
group = sorted(
|
||||
groups[key],
|
||||
key=lambda record: _stable_digest(seed, record.value["prompt"]),
|
||||
)
|
||||
size = len(group)
|
||||
for index, record in enumerate(group):
|
||||
quantile = (index + 0.5) / size
|
||||
ranked.append(
|
||||
(quantile, _stable_digest(seed + 1, record.value["prompt"]), record)
|
||||
)
|
||||
return [item[2] for item in sorted(ranked, key=lambda item: (item[0], item[1]))]
|
||||
|
||||
|
||||
def split_records(
|
||||
records: Sequence[SourceRecord],
|
||||
*,
|
||||
fixture_count: int,
|
||||
seed: int = 0xC1A551F1,
|
||||
train_ratio: float = 0.8,
|
||||
validation_ratio: float = 0.1,
|
||||
) -> SplitResult:
|
||||
if not records:
|
||||
raise DataError("cannot split an empty dataset")
|
||||
if fixture_count < 0:
|
||||
raise DataError("fixture_count cannot be negative")
|
||||
if not 0.0 < train_ratio < 1.0 or not 0.0 < validation_ratio < 1.0:
|
||||
raise DataError("split ratios must be in (0, 1)")
|
||||
if train_ratio + validation_ratio >= 1.0:
|
||||
raise DataError("train + validation ratios must leave room for test")
|
||||
|
||||
logical_total = len(records) + fixture_count
|
||||
target_train = round(logical_total * train_ratio)
|
||||
target_validation = round(logical_total * validation_ratio)
|
||||
target_test = logical_total - target_train - target_validation
|
||||
if fixture_count > target_test:
|
||||
raise DataError("fixture count exceeds the target test split")
|
||||
|
||||
vague = [record for record in records if record.value["slice"] == "vague-eval"]
|
||||
regular = [record for record in records if record.value["slice"] != "vague-eval"]
|
||||
eval_capacity = target_validation + target_test - fixture_count
|
||||
if len(vague) > eval_capacity:
|
||||
raise DataError(
|
||||
"vague-eval records exceed validation/test capacity; lower train_ratio"
|
||||
)
|
||||
|
||||
# Allocate vague records between validation and test in proportion to each split's
|
||||
# remaining capacity. None may enter training.
|
||||
test_source_capacity = target_test - fixture_count
|
||||
vague_validation_count = round(
|
||||
len(vague) * target_validation / (target_validation + test_source_capacity)
|
||||
)
|
||||
vague_validation_count = min(vague_validation_count, target_validation)
|
||||
vague_test_count = len(vague) - vague_validation_count
|
||||
if vague_test_count > test_source_capacity:
|
||||
overflow = vague_test_count - test_source_capacity
|
||||
vague_validation_count += overflow
|
||||
vague_test_count -= overflow
|
||||
|
||||
vague_order = _stratified_order(
|
||||
vague,
|
||||
seed=seed + 7,
|
||||
strata=lambda record: (record.value["purpose"],),
|
||||
)
|
||||
vague_validation = vague_order[:vague_validation_count]
|
||||
vague_test = vague_order[vague_validation_count:]
|
||||
|
||||
regular_train_count = target_train
|
||||
regular_validation_count = target_validation - len(vague_validation)
|
||||
regular_test_count = test_source_capacity - len(vague_test)
|
||||
if (
|
||||
regular_train_count + regular_validation_count + regular_test_count
|
||||
!= len(regular)
|
||||
):
|
||||
raise DataError("internal split accounting mismatch")
|
||||
|
||||
regular_order = _stratified_order(
|
||||
regular,
|
||||
seed=seed,
|
||||
strata=lambda record: (
|
||||
record.value["purpose"],
|
||||
record.value["slice"],
|
||||
record.value["lang"].split("-", 1)[0].casefold()
|
||||
if record.value["lang"].split("-", 1)[0].casefold() != "en"
|
||||
else "en",
|
||||
),
|
||||
)
|
||||
train = regular_order[:regular_train_count]
|
||||
validation_end = regular_train_count + regular_validation_count
|
||||
validation = regular_order[regular_train_count:validation_end] + vague_validation
|
||||
test = regular_order[validation_end:] + vague_test
|
||||
|
||||
# A second stable ordering makes file content independent of stratum dictionary order
|
||||
# and gives the trainer a deterministic shuffle before its epoch sampler takes over.
|
||||
def stable_sort(items: Sequence[SourceRecord], offset: int) -> list[SourceRecord]:
|
||||
return sorted(
|
||||
items,
|
||||
key=lambda record: _stable_digest(seed + offset, record.value["prompt"]),
|
||||
)
|
||||
|
||||
result = SplitResult(
|
||||
train=stable_sort(train, 11),
|
||||
validation=stable_sort(validation, 13),
|
||||
test=stable_sort(test, 17),
|
||||
fixture_count=fixture_count,
|
||||
)
|
||||
if any(record.value["slice"] == "vague-eval" for record in result.train):
|
||||
raise DataError("vague-eval leakage into training")
|
||||
if len(result.train) != target_train:
|
||||
raise DataError("train split missed its target size")
|
||||
if len(result.validation) != target_validation:
|
||||
raise DataError("validation split missed its target size")
|
||||
if result.logical_test_count != target_test:
|
||||
raise DataError("test split missed its target size")
|
||||
return result
|
||||
|
||||
|
||||
def distribution(records: Sequence[SourceRecord]) -> dict[str, dict[str, int]]:
|
||||
return {
|
||||
"purpose": dict(
|
||||
sorted(Counter(record.value["purpose"] for record in records).items())
|
||||
),
|
||||
"slice": dict(
|
||||
sorted(Counter(record.value["slice"] for record in records).items())
|
||||
),
|
||||
"language": dict(
|
||||
sorted(
|
||||
Counter(
|
||||
record.value["lang"].split("-", 1)[0].casefold()
|
||||
for record in records
|
||||
).items()
|
||||
)
|
||||
),
|
||||
}
|
||||
@@ -0,0 +1,2 @@
|
||||
numpy==2.5.1
|
||||
transformers==5.14.1
|
||||
@@ -0,0 +1,8 @@
|
||||
-r requirements-base.txt
|
||||
|
||||
# PyPI's Linux torch wheel pulls the full CUDA stack. The reproducible default is
|
||||
# deliberately CPU-only; accelerator training environments should install the matching
|
||||
# torch==2.13.0 build from pytorch.org, then install requirements-base.txt.
|
||||
--extra-index-url https://download.pytorch.org/whl/cpu
|
||||
torch==2.13.0+cpu ; sys_platform == "linux"
|
||||
torch==2.13.0 ; sys_platform != "linux"
|
||||
@@ -0,0 +1,77 @@
|
||||
import json
|
||||
import sys
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
MODULE_DIR = Path(__file__).resolve().parents[1]
|
||||
sys.path.insert(0, str(MODULE_DIR))
|
||||
|
||||
import prepare_data
|
||||
import purpose_data
|
||||
|
||||
|
||||
def example(index: int):
|
||||
return {
|
||||
"prompt": f"Implement sample endpoint number {index} with stable pagination",
|
||||
"purpose": purpose_data.LABELS[index % len(purpose_data.LABELS)],
|
||||
"secondary": None,
|
||||
"mixed": False,
|
||||
"difficulty": 0.4,
|
||||
"slice": "vague-eval" if index < 5 else "core",
|
||||
"lang": "en",
|
||||
}
|
||||
|
||||
|
||||
class PrepareIntegrationTests(unittest.TestCase):
|
||||
def test_refresh_then_verify_frozen_split(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory)
|
||||
source = root / "source.jsonl"
|
||||
purpose_data.write_jsonl(source, (example(index) for index in range(80)))
|
||||
fixtures = root / "fixtures.json"
|
||||
fixtures.write_text(
|
||||
json.dumps(
|
||||
[
|
||||
{"prompt": "Plan the cache migration", "purpose": "planning"},
|
||||
{"prompt": "Anything else?", "purpose": "general"},
|
||||
]
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
output = root / "output"
|
||||
frozen = root / "frozen.jsonl"
|
||||
manifest = root / "manifest.json"
|
||||
|
||||
first = prepare_data.prepare(
|
||||
sources=[source],
|
||||
fixtures_path=fixtures,
|
||||
output_dir=output,
|
||||
frozen_test_path=frozen,
|
||||
manifest_path=manifest,
|
||||
refresh_frozen_test=True,
|
||||
seed=23,
|
||||
near_duplicate_threshold=0.92,
|
||||
)
|
||||
second = prepare_data.prepare(
|
||||
sources=[source],
|
||||
fixtures_path=fixtures,
|
||||
output_dir=output,
|
||||
frozen_test_path=frozen,
|
||||
manifest_path=manifest,
|
||||
refresh_frozen_test=False,
|
||||
seed=23,
|
||||
near_duplicate_threshold=0.92,
|
||||
)
|
||||
|
||||
self.assertEqual(first, second)
|
||||
self.assertEqual(65, first["splits"]["train"]["records"])
|
||||
self.assertEqual(8, first["splits"]["validation"]["records"])
|
||||
self.assertEqual(8, first["splits"]["test"]["logicalRecords"])
|
||||
train = purpose_data.load_jsonl(output / "train.jsonl")
|
||||
self.assertFalse(any(row["slice"] == "vague-eval" for row in train))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,120 @@
|
||||
import sys
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
MODULE_DIR = Path(__file__).resolve().parents[1]
|
||||
sys.path.insert(0, str(MODULE_DIR))
|
||||
|
||||
import purpose_data
|
||||
|
||||
|
||||
def example(index: int, **overrides):
|
||||
value = {
|
||||
"prompt": f"Implement sample endpoint number {index} with stable pagination",
|
||||
"purpose": purpose_data.LABELS[index % len(purpose_data.LABELS)],
|
||||
"secondary": None,
|
||||
"mixed": False,
|
||||
"difficulty": 0.4,
|
||||
"slice": "core",
|
||||
"lang": "en",
|
||||
}
|
||||
value.update(overrides)
|
||||
return value
|
||||
|
||||
|
||||
def source(value, line=1):
|
||||
return purpose_data.SourceRecord(value, Path("source.jsonl"), line)
|
||||
|
||||
|
||||
class NormalizationTests(unittest.TestCase):
|
||||
def test_normalization_matches_runtime_whitespace_contract(self):
|
||||
self.assertEqual(
|
||||
"Café deploy now",
|
||||
purpose_data.normalize_prompt(" Cafe\u0301\tdeploy\nnow "),
|
||||
)
|
||||
self.assertEqual(
|
||||
purpose_data.normalized_key("FIX spacing"),
|
||||
purpose_data.normalized_key(" fix spacing "),
|
||||
)
|
||||
|
||||
|
||||
class CurationTests(unittest.TestCase):
|
||||
def test_excludes_exact_fixture_overlap(self):
|
||||
record = source(example(0, prompt="Make the toolbar nicer", purpose="frontendImpl"))
|
||||
result = purpose_data.curate_records(
|
||||
[record],
|
||||
[{"prompt": " make the toolbar nicer ", "purpose": "frontendImpl"}],
|
||||
)
|
||||
self.assertEqual([], result.records)
|
||||
self.assertEqual("exact", result.duplicates[0].kind)
|
||||
|
||||
def test_excludes_high_overlap_generated_template(self):
|
||||
words = [f"token{index}" for index in range(100)]
|
||||
first = " ".join(words)
|
||||
words[50] = "replacement"
|
||||
second = " ".join(words)
|
||||
result = purpose_data.curate_records(
|
||||
[
|
||||
source(example(0, prompt=first, purpose="backendImpl"), 1),
|
||||
source(example(1, prompt=second, purpose="backendImpl"), 2),
|
||||
],
|
||||
[],
|
||||
near_duplicate_threshold=0.92,
|
||||
)
|
||||
self.assertEqual(1, len(result.records))
|
||||
self.assertEqual(1, len(result.duplicates))
|
||||
self.assertEqual("near", result.duplicates[0].kind)
|
||||
self.assertGreaterEqual(result.duplicates[0].similarity, 0.92)
|
||||
|
||||
def test_near_duplicate_label_conflict_requires_review(self):
|
||||
words = [f"token{index}" for index in range(100)]
|
||||
first = " ".join(words)
|
||||
words[50] = "replacement"
|
||||
second = " ".join(words)
|
||||
with self.assertRaisesRegex(
|
||||
purpose_data.DataError, "label conflicts require review"
|
||||
):
|
||||
purpose_data.curate_records(
|
||||
[
|
||||
source(example(0, prompt=first, purpose="backendImpl"), 1),
|
||||
source(example(1, prompt=second, purpose="writing"), 2),
|
||||
],
|
||||
[],
|
||||
near_duplicate_threshold=0.92,
|
||||
)
|
||||
|
||||
|
||||
class SplitTests(unittest.TestCase):
|
||||
def test_split_is_deterministic_stratified_and_keeps_vague_out_of_train(self):
|
||||
records = []
|
||||
for index in range(1_000):
|
||||
slice_name = "vague-eval" if index < 50 else (
|
||||
"boundary" if index % 5 == 0 else "core"
|
||||
)
|
||||
records.append(source(example(index, slice=slice_name), index + 1))
|
||||
|
||||
first = purpose_data.split_records(records, fixture_count=10, seed=17)
|
||||
second = purpose_data.split_records(records, fixture_count=10, seed=17)
|
||||
|
||||
self.assertEqual(
|
||||
[row.value["prompt"] for row in first.train],
|
||||
[row.value["prompt"] for row in second.train],
|
||||
)
|
||||
self.assertEqual(808, len(first.train))
|
||||
self.assertEqual(101, len(first.validation))
|
||||
self.assertEqual(101, first.logical_test_count)
|
||||
self.assertFalse(
|
||||
any(row.value["slice"] == "vague-eval" for row in first.train)
|
||||
)
|
||||
split_prompts = [
|
||||
{row.value["prompt"] for row in split}
|
||||
for split in (first.train, first.validation, first.test)
|
||||
]
|
||||
self.assertFalse(split_prompts[0] & split_prompts[1])
|
||||
self.assertFalse(split_prompts[0] & split_prompts[2])
|
||||
self.assertFalse(split_prompts[1] & split_prompts[2])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,46 @@
|
||||
import sys
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
MODULE_DIR = Path(__file__).resolve().parents[1]
|
||||
sys.path.insert(0, str(MODULE_DIR))
|
||||
|
||||
import train
|
||||
|
||||
|
||||
class MetricsTests(unittest.TestCase):
|
||||
def test_classification_metrics_include_every_label(self):
|
||||
actual = list(range(8))
|
||||
predicted = [0, 1, 2, 3, 4, 5, 6, 0]
|
||||
metrics = train.classification_metrics(actual, predicted)
|
||||
self.assertEqual(7 / 8, metrics["accuracy"])
|
||||
self.assertEqual(0.0, metrics["perPurposeRecall"]["writing"])
|
||||
self.assertEqual(1.0, metrics["perPurposeRecall"]["planning"])
|
||||
|
||||
def test_thresholds_preserve_confidence_nesting(self):
|
||||
probabilities = [0.99, 0.95, 0.85, 0.75, 0.65]
|
||||
margins = [0.95, 0.85, 0.60, 0.40, 0.20]
|
||||
correct = [True, True, True, False, False]
|
||||
thresholds = train.choose_confidence_thresholds(
|
||||
probabilities,
|
||||
margins,
|
||||
correct,
|
||||
high_precision=1.0,
|
||||
accepted_precision=0.75,
|
||||
)
|
||||
self.assertGreaterEqual(
|
||||
thresholds["high"]["minimumScore"],
|
||||
thresholds["medium"]["minimumScore"],
|
||||
)
|
||||
self.assertGreater(thresholds["medium"]["validationAcceptedCoverage"], 0)
|
||||
|
||||
def test_expected_calibration_error_is_zero_for_perfect_extremes(self):
|
||||
self.assertEqual(
|
||||
0.0,
|
||||
train.expected_calibration_error([1.0, 0.0], [True, False]),
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,480 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Fine-tune the fixed-shape purpose-lite MiniLM classifier."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import math
|
||||
import random
|
||||
import shutil
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any, Sequence
|
||||
|
||||
from purpose_data import LABELS, DataError, load_jsonl, normalize_prompt, write_json
|
||||
|
||||
|
||||
SCRIPT_DIR = Path(__file__).resolve().parent
|
||||
DEFAULT_DATASET_DIR = SCRIPT_DIR / ".artifacts" / "dataset-v1"
|
||||
DEFAULT_OUTPUT_DIR = SCRIPT_DIR / "outputs" / "purpose-lite-v1"
|
||||
DEFAULT_MODEL = "sentence-transformers/all-MiniLM-L6-v2"
|
||||
# Reproducibility requires a model commit, not a mutable `main` branch.
|
||||
DEFAULT_MODEL_REVISION = "1110a243fdf4706b3f48f1d95db1a4f5529b4d41"
|
||||
MAX_LENGTH = 128
|
||||
|
||||
|
||||
def prepare_text(prompt: str) -> str:
|
||||
return normalize_prompt(prompt)
|
||||
|
||||
|
||||
def classification_metrics(
|
||||
actual: Sequence[int], predicted: Sequence[int]
|
||||
) -> dict[str, Any]:
|
||||
if len(actual) != len(predicted) or not actual:
|
||||
raise ValueError("metrics need equally sized, non-empty vectors")
|
||||
correct = sum(want == got for want, got in zip(actual, predicted))
|
||||
recalls: dict[str, float] = {}
|
||||
confusion = [[0 for _ in LABELS] for _ in LABELS]
|
||||
for want, got in zip(actual, predicted):
|
||||
confusion[want][got] += 1
|
||||
for index, label in enumerate(LABELS):
|
||||
total = sum(confusion[index])
|
||||
recalls[label] = confusion[index][index] / total if total else 0.0
|
||||
return {
|
||||
"records": len(actual),
|
||||
"accuracy": correct / len(actual),
|
||||
"macroRecall": sum(recalls.values()) / len(recalls),
|
||||
"perPurposeRecall": recalls,
|
||||
"confusionMatrix": {
|
||||
"labels": list(LABELS),
|
||||
"rows": confusion,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def confidence_score(top_probability: float, top_two_margin: float) -> float:
|
||||
"""One monotonic score that keeps both calibration signals in the contract."""
|
||||
|
||||
return top_probability * (0.5 + 0.5 * top_two_margin)
|
||||
|
||||
|
||||
def _threshold_for_precision(
|
||||
scores: Sequence[float],
|
||||
correct: Sequence[bool],
|
||||
target_precision: float,
|
||||
) -> tuple[float, float, float]:
|
||||
ranked = sorted(zip(scores, correct), key=lambda item: item[0], reverse=True)
|
||||
accepted = 0
|
||||
accepted_correct = 0
|
||||
best: tuple[float, float, float] | None = None
|
||||
index = 0
|
||||
while index < len(ranked):
|
||||
score = ranked[index][0]
|
||||
while index < len(ranked) and ranked[index][0] == score:
|
||||
accepted += 1
|
||||
accepted_correct += int(ranked[index][1])
|
||||
index += 1
|
||||
precision = accepted_correct / accepted
|
||||
if precision >= target_precision:
|
||||
best = (score, precision, accepted / len(ranked))
|
||||
if best is None:
|
||||
return 1.000001, 1.0, 0.0
|
||||
return best
|
||||
|
||||
|
||||
def choose_confidence_thresholds(
|
||||
top_probabilities: Sequence[float],
|
||||
top_two_margins: Sequence[float],
|
||||
correct: Sequence[bool],
|
||||
*,
|
||||
high_precision: float = 0.98,
|
||||
accepted_precision: float = 0.95,
|
||||
) -> dict[str, Any]:
|
||||
if not (
|
||||
len(top_probabilities) == len(top_two_margins) == len(correct)
|
||||
and top_probabilities
|
||||
):
|
||||
raise ValueError("threshold calibration needs equally sized, non-empty vectors")
|
||||
scores = [
|
||||
confidence_score(probability, margin)
|
||||
for probability, margin in zip(top_probabilities, top_two_margins)
|
||||
]
|
||||
high = _threshold_for_precision(scores, correct, high_precision)
|
||||
medium = _threshold_for_precision(scores, correct, accepted_precision)
|
||||
# HIGH must always be a subset of the accepted MEDIUM-or-better population.
|
||||
high_threshold = max(high[0], medium[0])
|
||||
return {
|
||||
"score": {
|
||||
"formula": "topProbability * (0.5 + 0.5 * topTwoMargin)",
|
||||
"probabilityWeight": 0.5,
|
||||
"marginInteractionWeight": 0.5,
|
||||
},
|
||||
"high": {
|
||||
"minimumScore": high_threshold,
|
||||
"targetPrecision": high_precision,
|
||||
"validationPrecision": high[1],
|
||||
"validationCoverage": high[2] if high_threshold == high[0] else 0.0,
|
||||
},
|
||||
"medium": {
|
||||
"minimumScore": medium[0],
|
||||
"targetAcceptedPrecision": accepted_precision,
|
||||
"validationAcceptedPrecision": medium[1],
|
||||
"validationAcceptedCoverage": medium[2],
|
||||
},
|
||||
"low": {"minimumScore": 0.0},
|
||||
}
|
||||
|
||||
|
||||
def expected_calibration_error(
|
||||
probabilities: Sequence[float],
|
||||
correct: Sequence[bool],
|
||||
bins: int = 15,
|
||||
) -> float:
|
||||
if len(probabilities) != len(correct) or not probabilities:
|
||||
raise ValueError("ECE needs equally sized, non-empty vectors")
|
||||
total_error = 0.0
|
||||
for lower_index in range(bins):
|
||||
lower = lower_index / bins
|
||||
upper = (lower_index + 1) / bins
|
||||
members = [
|
||||
index
|
||||
for index, value in enumerate(probabilities)
|
||||
if lower <= value < upper or (upper == 1.0 and value == 1.0)
|
||||
]
|
||||
if not members:
|
||||
continue
|
||||
confidence = sum(probabilities[index] for index in members) / len(members)
|
||||
accuracy = sum(correct[index] for index in members) / len(members)
|
||||
total_error += len(members) / len(probabilities) * abs(confidence - accuracy)
|
||||
return total_error
|
||||
|
||||
|
||||
def _validate_split(records: Sequence[dict[str, Any]], path: Path) -> None:
|
||||
if not records:
|
||||
raise DataError(f"{path}: split is empty")
|
||||
for index, record in enumerate(records, 1):
|
||||
if record.get("purpose") not in LABELS:
|
||||
raise DataError(f"{path}:{index}: invalid purpose")
|
||||
if not isinstance(record.get("prompt"), str) or not record["prompt"].strip():
|
||||
raise DataError(f"{path}:{index}: invalid prompt")
|
||||
|
||||
|
||||
def _select_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 _set_seeds(torch: Any, seed: int) -> None:
|
||||
random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
|
||||
|
||||
def _fit_temperature(torch: Any, logits: Any, labels: Any) -> float:
|
||||
log_temperature = torch.zeros(1, requires_grad=True)
|
||||
optimizer = torch.optim.LBFGS(
|
||||
[log_temperature], lr=0.05, max_iter=100, line_search_fn="strong_wolfe"
|
||||
)
|
||||
|
||||
def closure() -> Any:
|
||||
optimizer.zero_grad()
|
||||
temperature = log_temperature.exp().clamp(0.05, 20.0)
|
||||
loss = torch.nn.functional.cross_entropy(logits / temperature, labels)
|
||||
loss.backward()
|
||||
return loss
|
||||
|
||||
optimizer.step(closure)
|
||||
return float(log_temperature.detach().exp().clamp(0.05, 20.0).item())
|
||||
|
||||
|
||||
def _evaluate(torch: Any, model: Any, loader: Any, device: Any) -> tuple[Any, Any]:
|
||||
model.eval()
|
||||
all_logits = []
|
||||
all_labels = []
|
||||
with torch.inference_mode():
|
||||
for batch in loader:
|
||||
labels = batch.pop("labels")
|
||||
inputs = {key: value.to(device) for key, value in batch.items()}
|
||||
logits = model(**inputs).logits.cpu()
|
||||
all_logits.append(logits)
|
||||
all_labels.append(labels)
|
||||
return torch.cat(all_logits), torch.cat(all_labels)
|
||||
|
||||
|
||||
def train(args: argparse.Namespace) -> dict[str, Any]:
|
||||
try:
|
||||
import torch
|
||||
from torch.utils.data import DataLoader, Dataset
|
||||
from transformers import (
|
||||
AutoModelForSequenceClassification,
|
||||
AutoTokenizer,
|
||||
get_linear_schedule_with_warmup,
|
||||
)
|
||||
except ImportError as exc:
|
||||
raise DataError(
|
||||
"training dependencies are missing; install requirements.txt in a virtualenv"
|
||||
) from exc
|
||||
|
||||
train_path = args.dataset_dir / "train.jsonl"
|
||||
validation_path = args.dataset_dir / "validation.jsonl"
|
||||
train_records = load_jsonl(train_path)
|
||||
validation_records = load_jsonl(validation_path)
|
||||
_validate_split(train_records, train_path)
|
||||
_validate_split(validation_records, validation_path)
|
||||
if args.max_train_records:
|
||||
train_records = train_records[: args.max_train_records]
|
||||
if args.max_validation_records:
|
||||
validation_records = validation_records[: args.max_validation_records]
|
||||
|
||||
output_dir: Path = args.output_dir
|
||||
if output_dir.exists() and any(output_dir.iterdir()):
|
||||
if not args.overwrite_output:
|
||||
raise DataError(
|
||||
f"{output_dir}: output is not empty; pass --overwrite-output intentionally"
|
||||
)
|
||||
shutil.rmtree(output_dir)
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
_set_seeds(torch, args.seed)
|
||||
device = _select_device(torch, args.device)
|
||||
label_to_id = {label: index for index, label in enumerate(LABELS)}
|
||||
id_to_label = {index: label for label, index in label_to_id.items()}
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
args.model, revision=args.model_revision, use_fast=True
|
||||
)
|
||||
model = AutoModelForSequenceClassification.from_pretrained(
|
||||
args.model,
|
||||
revision=args.model_revision,
|
||||
num_labels=len(LABELS),
|
||||
label2id=label_to_id,
|
||||
id2label=id_to_label,
|
||||
ignore_mismatched_sizes=True,
|
||||
)
|
||||
config = model.config
|
||||
if getattr(config, "hidden_size", None) != 384 or getattr(
|
||||
config, "num_hidden_layers", None
|
||||
) != 6:
|
||||
raise DataError(
|
||||
"purpose-lite must remain a 6-layer, 384-dimensional MiniLM encoder"
|
||||
)
|
||||
config.purpose_classifier_version = "purpose-lite-v1"
|
||||
config.purpose_classifier_max_length = MAX_LENGTH
|
||||
config.purpose_classifier_fixed_shape = [1, MAX_LENGTH]
|
||||
model.to(device)
|
||||
|
||||
class PromptDataset(Dataset):
|
||||
def __init__(self, records: Sequence[dict[str, Any]]) -> None:
|
||||
self.records = records
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self.records)
|
||||
|
||||
def __getitem__(self, index: int) -> tuple[str, int]:
|
||||
record = self.records[index]
|
||||
return prepare_text(record["prompt"]), label_to_id[record["purpose"]]
|
||||
|
||||
def collate(items: Sequence[tuple[str, int]]) -> dict[str, Any]:
|
||||
texts, labels = zip(*items)
|
||||
encoded = tokenizer(
|
||||
list(texts),
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
max_length=MAX_LENGTH,
|
||||
return_tensors="pt",
|
||||
)
|
||||
encoded["labels"] = torch.tensor(labels, dtype=torch.long)
|
||||
return encoded
|
||||
|
||||
generator = torch.Generator()
|
||||
generator.manual_seed(args.seed)
|
||||
train_loader = DataLoader(
|
||||
PromptDataset(train_records),
|
||||
batch_size=args.batch_size,
|
||||
shuffle=True,
|
||||
generator=generator,
|
||||
collate_fn=collate,
|
||||
num_workers=args.workers,
|
||||
pin_memory=device.type == "cuda",
|
||||
)
|
||||
validation_loader = DataLoader(
|
||||
PromptDataset(validation_records),
|
||||
batch_size=args.eval_batch_size,
|
||||
shuffle=False,
|
||||
collate_fn=collate,
|
||||
num_workers=args.workers,
|
||||
pin_memory=device.type == "cuda",
|
||||
)
|
||||
|
||||
optimizer = torch.optim.AdamW(
|
||||
model.parameters(), lr=args.learning_rate, weight_decay=args.weight_decay
|
||||
)
|
||||
update_steps_per_epoch = math.ceil(
|
||||
len(train_loader) / args.gradient_accumulation_steps
|
||||
)
|
||||
total_steps = update_steps_per_epoch * args.epochs
|
||||
scheduler = get_linear_schedule_with_warmup(
|
||||
optimizer,
|
||||
num_warmup_steps=round(total_steps * args.warmup_ratio),
|
||||
num_training_steps=total_steps,
|
||||
)
|
||||
|
||||
best_accuracy = -1.0
|
||||
history = []
|
||||
best_dir = output_dir / "model"
|
||||
started = time.perf_counter()
|
||||
for epoch in range(1, args.epochs + 1):
|
||||
model.train()
|
||||
optimizer.zero_grad(set_to_none=True)
|
||||
running_loss = 0.0
|
||||
for step, batch in enumerate(train_loader, 1):
|
||||
batch = {key: value.to(device) for key, value in batch.items()}
|
||||
loss = model(**batch).loss / args.gradient_accumulation_steps
|
||||
loss.backward()
|
||||
running_loss += float(loss.item()) * args.gradient_accumulation_steps
|
||||
should_update = (
|
||||
step % args.gradient_accumulation_steps == 0
|
||||
or step == len(train_loader)
|
||||
)
|
||||
if should_update:
|
||||
torch.nn.utils.clip_grad_norm_(model.parameters(), args.max_grad_norm)
|
||||
optimizer.step()
|
||||
scheduler.step()
|
||||
optimizer.zero_grad(set_to_none=True)
|
||||
|
||||
logits, labels = _evaluate(torch, model, validation_loader, device)
|
||||
predictions = logits.argmax(dim=-1).tolist()
|
||||
metrics = classification_metrics(labels.tolist(), predictions)
|
||||
metrics["epoch"] = epoch
|
||||
metrics["meanTrainingLoss"] = running_loss / len(train_loader)
|
||||
history.append(metrics)
|
||||
print(
|
||||
f"epoch {epoch}: loss={metrics['meanTrainingLoss']:.4f} "
|
||||
f"validation_accuracy={metrics['accuracy']:.4%} "
|
||||
f"macro_recall={metrics['macroRecall']:.4%}",
|
||||
flush=True,
|
||||
)
|
||||
if metrics["accuracy"] > best_accuracy:
|
||||
best_accuracy = metrics["accuracy"]
|
||||
model.save_pretrained(best_dir, safe_serialization=True)
|
||||
tokenizer.save_pretrained(best_dir)
|
||||
|
||||
model = AutoModelForSequenceClassification.from_pretrained(best_dir).to(device)
|
||||
logits, labels = _evaluate(torch, model, validation_loader, device)
|
||||
temperature = _fit_temperature(torch, logits, labels)
|
||||
calibrated = torch.softmax(logits / temperature, dim=-1)
|
||||
top = torch.topk(calibrated, k=2, dim=-1)
|
||||
top_probabilities = top.values[:, 0].tolist()
|
||||
margins = (top.values[:, 0] - top.values[:, 1]).tolist()
|
||||
predictions = top.indices[:, 0].tolist()
|
||||
correct = [
|
||||
prediction == actual
|
||||
for prediction, actual in zip(predictions, labels.tolist())
|
||||
]
|
||||
thresholds = choose_confidence_thresholds(
|
||||
top_probabilities,
|
||||
margins,
|
||||
correct,
|
||||
high_precision=args.high_precision,
|
||||
accepted_precision=args.accepted_precision,
|
||||
)
|
||||
calibration = {
|
||||
"schemaVersion": 1,
|
||||
"modelVersion": "purpose-lite-v1",
|
||||
"labels": list(LABELS),
|
||||
"temperature": temperature,
|
||||
"confidence": thresholds,
|
||||
"validationECE": expected_calibration_error(top_probabilities, correct),
|
||||
}
|
||||
metrics = {
|
||||
"modelVersion": "purpose-lite-v1",
|
||||
"baseModel": args.model,
|
||||
"baseModelRevision": args.model_revision,
|
||||
"fixedInputShape": [1, MAX_LENGTH],
|
||||
"device": str(device),
|
||||
"trainingSeconds": time.perf_counter() - started,
|
||||
"trainRecords": len(train_records),
|
||||
"validationRecords": len(validation_records),
|
||||
"bestValidationAccuracy": best_accuracy,
|
||||
"bestValidation": classification_metrics(labels.tolist(), predictions),
|
||||
"history": history,
|
||||
"calibration": calibration,
|
||||
}
|
||||
write_json(output_dir / "calibration.json", calibration)
|
||||
write_json(output_dir / "metrics.json", metrics)
|
||||
write_json(
|
||||
output_dir / "training-config.json",
|
||||
{
|
||||
key: str(value) if isinstance(value, Path) else value
|
||||
for key, value in vars(args).items()
|
||||
},
|
||||
)
|
||||
return metrics
|
||||
|
||||
|
||||
def build_parser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--dataset-dir", type=Path, default=DEFAULT_DATASET_DIR)
|
||||
parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR)
|
||||
parser.add_argument("--model", default=DEFAULT_MODEL)
|
||||
parser.add_argument("--model-revision", default=DEFAULT_MODEL_REVISION)
|
||||
parser.add_argument("--device", default="auto")
|
||||
parser.add_argument("--seed", type=int, default=20260730)
|
||||
parser.add_argument("--epochs", type=int, default=3)
|
||||
parser.add_argument("--batch-size", type=int, default=32)
|
||||
parser.add_argument("--eval-batch-size", type=int, default=64)
|
||||
parser.add_argument("--gradient-accumulation-steps", type=int, default=1)
|
||||
parser.add_argument("--learning-rate", type=float, default=2e-5)
|
||||
parser.add_argument("--weight-decay", type=float, default=0.01)
|
||||
parser.add_argument("--warmup-ratio", type=float, default=0.1)
|
||||
parser.add_argument("--max-grad-norm", type=float, default=1.0)
|
||||
parser.add_argument("--workers", type=int, default=0)
|
||||
parser.add_argument("--high-precision", type=float, default=0.98)
|
||||
parser.add_argument("--accepted-precision", type=float, default=0.95)
|
||||
parser.add_argument("--max-train-records", type=int)
|
||||
parser.add_argument("--max-validation-records", type=int)
|
||||
parser.add_argument("--overwrite-output", action="store_true")
|
||||
return parser
|
||||
|
||||
|
||||
def _positive(parser: argparse.ArgumentParser, name: str, value: int) -> None:
|
||||
if value <= 0:
|
||||
parser.error(f"{name} must be positive")
|
||||
|
||||
|
||||
def main(argv: Sequence[str] | None = None) -> int:
|
||||
parser = build_parser()
|
||||
args = parser.parse_args(argv)
|
||||
for name in (
|
||||
"epochs",
|
||||
"batch_size",
|
||||
"eval_batch_size",
|
||||
"gradient_accumulation_steps",
|
||||
):
|
||||
_positive(parser, f"--{name.replace('_', '-')}", getattr(args, name))
|
||||
if not 0.0 <= args.warmup_ratio < 1.0:
|
||||
parser.error("--warmup-ratio must be in [0, 1)")
|
||||
if not 0.0 < args.accepted_precision <= args.high_precision <= 1.0:
|
||||
parser.error(
|
||||
"precision targets must satisfy 0 < accepted <= high <= 1"
|
||||
)
|
||||
try:
|
||||
metrics = train(args)
|
||||
except (DataError, OSError, ValueError) as exc:
|
||||
print(f"error: {exc}", file=sys.stderr)
|
||||
return 1
|
||||
print(
|
||||
f"Saved purpose-lite-v1; best validation accuracy "
|
||||
f"{metrics['bestValidationAccuracy']:.4%}."
|
||||
)
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
+13
-9
@@ -7,7 +7,7 @@ process violations are errors. Approximate targets (the requirements written as
|
||||
|
||||
Usage:
|
||||
python3 ml/purpose-classifier/validate-data.py
|
||||
python3 ml/purpose-classifier/validate-data.py path/to/batch.jsonl
|
||||
python3 ml/purpose-classifier/validate-data.py path/to/batch.jsonl --batch-size 200
|
||||
python3 ml/purpose-classifier/validate-data.py --strict
|
||||
"""
|
||||
|
||||
@@ -28,6 +28,10 @@ from typing import Any, Iterable
|
||||
|
||||
SCRIPT_DIR = Path(__file__).resolve().parent
|
||||
DEFAULT_DATA_DIR = SCRIPT_DIR / "data"
|
||||
DEFAULT_CANONICAL_FILES = (
|
||||
DEFAULT_DATA_DIR / "purpose-prompts.jsonl",
|
||||
DEFAULT_DATA_DIR / "purpose-prompts-round2.jsonl",
|
||||
)
|
||||
DEFAULT_FIXTURES = (
|
||||
SCRIPT_DIR.parent.parent
|
||||
/ "Tests"
|
||||
@@ -177,8 +181,8 @@ class Validator:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
batch_size: int = 200,
|
||||
expected_total: int = 8_000,
|
||||
batch_size: int = 0,
|
||||
expected_total: int = 12_214,
|
||||
fixture_path: Path | None = DEFAULT_FIXTURES,
|
||||
process_checks: bool = True,
|
||||
) -> None:
|
||||
@@ -520,7 +524,7 @@ class Validator:
|
||||
def _check_manifests(self, roots: list[Path]) -> None:
|
||||
directories = sorted({root if root.is_dir() else root.parent for root in roots})
|
||||
for directory in directories:
|
||||
candidates = sorted(directory.glob("*manifest*.json"))
|
||||
candidates = sorted(directory.glob("*generation-manifest*.json"))
|
||||
if not candidates:
|
||||
self.warning(
|
||||
directory,
|
||||
@@ -687,14 +691,14 @@ def build_parser() -> argparse.ArgumentParser:
|
||||
parser.add_argument(
|
||||
"--batch-size",
|
||||
type=int,
|
||||
default=200,
|
||||
help="required generation batch size; use 0 to disable (default: 200)",
|
||||
default=0,
|
||||
help="required generation batch size; use 200 for raw batches (default: disabled)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--expected-total",
|
||||
type=int,
|
||||
default=8_000,
|
||||
help="expected total record count; use 0 to disable (default: 8000)",
|
||||
default=12_214,
|
||||
help="expected canonical record count; use 0 to disable (default: 12214)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--fixtures",
|
||||
@@ -727,7 +731,7 @@ def main(argv: list[str] | None = None) -> int:
|
||||
print("error: numeric options must be non-negative", file=sys.stderr)
|
||||
return 2
|
||||
|
||||
targets = args.paths or [DEFAULT_DATA_DIR]
|
||||
targets = args.paths or list(DEFAULT_CANONICAL_FILES)
|
||||
paths, roots, discovery_errors = discover_paths(targets)
|
||||
if discovery_errors:
|
||||
for message in discovery_errors:
|
||||
|
||||
Reference in New Issue
Block a user