Merge nucleic/lucid-north-quail-rnvt into dev

This commit is contained in:
2026-07-30 03:48:00 -07:00
parent e2a5314f1b
commit 9c43be6df1
17 changed files with 3325 additions and 13 deletions
+5
View File
@@ -0,0 +1,5 @@
.artifacts/
.venv/
outputs/
__pycache__/
tests/__pycache__/
+94
View File
@@ -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.
+147
View File
@@ -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
+57
View File
@@ -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."
]
}
-3
View File
@@ -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
View File
@@ -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).
+292
View File
@@ -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
View File
@@ -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
View File
@@ -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()
)
),
}
+2
View File
@@ -0,0 +1,2 @@
numpy==2.5.1
transformers==5.14.1
+8
View File
@@ -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"
+77
View File
@@ -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()
+120
View File
@@ -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()
+46
View File
@@ -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()
+480
View File
@@ -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
View File
@@ -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: