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":"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":"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":"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":"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":"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"}
|
{"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":"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":"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":"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":"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":"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"}
|
{"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":"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":"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":"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":"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":"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"}
|
{"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
|
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
|
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
|
(Tests/NucleicCoreTests/Fixtures/purpose-prompts.json) are eval-only and must NOT be
|
||||||
pasted into the generator's context (contamination).
|
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:
|
Usage:
|
||||||
python3 ml/purpose-classifier/validate-data.py
|
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
|
python3 ml/purpose-classifier/validate-data.py --strict
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@@ -28,6 +28,10 @@ from typing import Any, Iterable
|
|||||||
|
|
||||||
SCRIPT_DIR = Path(__file__).resolve().parent
|
SCRIPT_DIR = Path(__file__).resolve().parent
|
||||||
DEFAULT_DATA_DIR = SCRIPT_DIR / "data"
|
DEFAULT_DATA_DIR = SCRIPT_DIR / "data"
|
||||||
|
DEFAULT_CANONICAL_FILES = (
|
||||||
|
DEFAULT_DATA_DIR / "purpose-prompts.jsonl",
|
||||||
|
DEFAULT_DATA_DIR / "purpose-prompts-round2.jsonl",
|
||||||
|
)
|
||||||
DEFAULT_FIXTURES = (
|
DEFAULT_FIXTURES = (
|
||||||
SCRIPT_DIR.parent.parent
|
SCRIPT_DIR.parent.parent
|
||||||
/ "Tests"
|
/ "Tests"
|
||||||
@@ -177,8 +181,8 @@ class Validator:
|
|||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
batch_size: int = 200,
|
batch_size: int = 0,
|
||||||
expected_total: int = 8_000,
|
expected_total: int = 12_214,
|
||||||
fixture_path: Path | None = DEFAULT_FIXTURES,
|
fixture_path: Path | None = DEFAULT_FIXTURES,
|
||||||
process_checks: bool = True,
|
process_checks: bool = True,
|
||||||
) -> None:
|
) -> None:
|
||||||
@@ -520,7 +524,7 @@ class Validator:
|
|||||||
def _check_manifests(self, roots: list[Path]) -> None:
|
def _check_manifests(self, roots: list[Path]) -> None:
|
||||||
directories = sorted({root if root.is_dir() else root.parent for root in roots})
|
directories = sorted({root if root.is_dir() else root.parent for root in roots})
|
||||||
for directory in directories:
|
for directory in directories:
|
||||||
candidates = sorted(directory.glob("*manifest*.json"))
|
candidates = sorted(directory.glob("*generation-manifest*.json"))
|
||||||
if not candidates:
|
if not candidates:
|
||||||
self.warning(
|
self.warning(
|
||||||
directory,
|
directory,
|
||||||
@@ -687,14 +691,14 @@ def build_parser() -> argparse.ArgumentParser:
|
|||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--batch-size",
|
"--batch-size",
|
||||||
type=int,
|
type=int,
|
||||||
default=200,
|
default=0,
|
||||||
help="required generation batch size; use 0 to disable (default: 200)",
|
help="required generation batch size; use 200 for raw batches (default: disabled)",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--expected-total",
|
"--expected-total",
|
||||||
type=int,
|
type=int,
|
||||||
default=8_000,
|
default=12_214,
|
||||||
help="expected total record count; use 0 to disable (default: 8000)",
|
help="expected canonical record count; use 0 to disable (default: 12214)",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--fixtures",
|
"--fixtures",
|
||||||
@@ -727,7 +731,7 @@ def main(argv: list[str] | None = None) -> int:
|
|||||||
print("error: numeric options must be non-negative", file=sys.stderr)
|
print("error: numeric options must be non-negative", file=sys.stderr)
|
||||||
return 2
|
return 2
|
||||||
|
|
||||||
targets = args.paths or [DEFAULT_DATA_DIR]
|
targets = args.paths or list(DEFAULT_CANONICAL_FILES)
|
||||||
paths, roots, discovery_errors = discover_paths(targets)
|
paths, roots, discovery_errors = discover_paths(targets)
|
||||||
if discovery_errors:
|
if discovery_errors:
|
||||||
for message in discovery_errors:
|
for message in discovery_errors:
|
||||||
|
|||||||
Reference in New Issue
Block a user