Merge nucleic/sleek-ember-seal-uady into dev

This commit is contained in:
2026-07-30 17:28:44 -07:00
parent 93e9c838bb
commit e4403b570d
7 changed files with 408 additions and 19 deletions
+32 -3
View File
@@ -72,7 +72,23 @@ platform selector, then install `requirements-base.txt`.
Training writes a local checkpoint, `calibration.json`, and `metrics.json` under
`outputs/purpose-lite-v1/`. It selects checkpoints and fits temperature on label-scorable
validation records. When deriving nested HIGH/MEDIUM/LOW cutoffs, every `vague-eval`
record counts as an abstention miss even if its synthetic label happens to match.
record counts as an abstention miss even if its synthetic label happens to match. The
incoming checkpoint is scored and retained as epoch zero, so a continuation run cannot
silently replace it with a regression. Validation early stopping defaults to two epochs
without an improvement greater than 0.05 points.
Continuation training accepts a local checkpoint. `--boundary-weight` is an opt-in,
validation-selected loss weight for the measured weakest slice; it does not add held-out
fixtures to training:
```bash
ml/purpose-classifier/.venv/bin/python ml/purpose-classifier/train.py \
--model ml/purpose-classifier/outputs/purpose-lite-v1/model \
--epochs 3 --learning-rate 3e-6 --warmup-ratio 0 \
--boundary-weight 2 \
--output-dir ml/purpose-classifier/outputs/purpose-lite-v1-boundary-tune \
--overwrite-output
```
For a wiring smoke test, use a small deterministic prefix:
@@ -104,8 +120,21 @@ ml/purpose-classifier/.venv/bin/python ml/purpose-classifier/eval.py \
`export.py` emits fixed-shape opset-17 fp16 and int8-QDQ graphs, a tokenizer/
normalization contract, golden tokenizations, shared calibration config, graph checks,
artifact hashes, and a size report. The int8 graph is the ≤25 MiB shipping candidate;
the fp16 graph remains the accelerator-oriented conversion input.
artifact hashes, and a size report. Its default 256-record quantization calibration sample
is deterministic and stratified by purpose, slice, and primary language; the export report
records the seed, distribution, and prompt hashes. The int8 graph is the ≤25 MiB shipping
candidate; the fp16 graph remains the accelerator-oriented conversion input.
When scoring ONNX, add `--compare-pytorch` to measure artifact drift against
`--model-dir`. The report then includes overall, label-scorable, and per-slice label
agreement plus every correct→incorrect, incorrect→correct, and changed-wrong-label
transition:
```bash
ml/purpose-classifier/.venv/bin/python ml/purpose-classifier/eval.py \
--onnx-model ml/purpose-classifier/outputs/purpose-lite-v1/export/purpose-lite-v1-int8-qdq.onnx \
--compare-pytorch --no-gate
```
## Audit curation
+101
View File
@@ -140,6 +140,80 @@ def routing_tier_drift(
}
def prediction_agreement(
records: Sequence[dict[str, Any]],
actual: Sequence[int],
reference: Sequence[int],
candidate: Sequence[int],
) -> dict[str, Any]:
"""Report label drift from a reference runtime to a candidate artifact."""
if not (
len(records) == len(actual) == len(reference) == len(candidate)
and records
):
raise ValueError("prediction agreement needs equally sized, non-empty vectors")
agreements = [
expected == got for expected, got in zip(reference, candidate)
]
scorable_indexes = [
index
for index, record in enumerate(records)
if record.get("slice") != "vague-eval"
]
transitions = Counter()
disagreements = []
for index, agrees in enumerate(agreements):
if agrees:
continue
reference_correct = reference[index] == actual[index]
candidate_correct = candidate[index] == actual[index]
if reference_correct and not candidate_correct:
transition = "correctToIncorrect"
elif not reference_correct and candidate_correct:
transition = "incorrectToCorrect"
else:
transition = "differentIncorrectLabel"
transitions[transition] += 1
disagreements.append(
{
"promptHash": prompt_hash(records[index]["prompt"]),
"slice": records[index].get("slice", "unknown"),
"expected": LABELS[actual[index]],
"reference": LABELS[reference[index]],
"candidate": LABELS[candidate[index]],
"transition": transition,
}
)
by_slice = {}
for slice_name in sorted(
{record.get("slice", "unknown") for record in records}
):
indexes = [
index
for index, record in enumerate(records)
if record.get("slice", "unknown") == slice_name
]
by_slice[slice_name] = {
"records": len(indexes),
"labelAgreement": (
sum(agreements[index] for index in indexes) / len(indexes)
),
}
return {
"records": len(records),
"labelAgreement": sum(agreements) / len(agreements),
"scoredLabelAgreement": (
sum(agreements[index] for index in scorable_indexes)
/ len(scorable_indexes)
),
"transitionCounts": dict(sorted(transitions.items())),
"bySlice": by_slice,
"disagreements": disagreements,
}
def evaluate(args: argparse.Namespace) -> dict[str, Any]:
try:
import torch
@@ -171,6 +245,7 @@ def evaluate(args: argparse.Namespace) -> dict[str, Any]:
label_to_id = {label: index for index, label in enumerate(LABELS)}
tokenizer = AutoTokenizer.from_pretrained(args.model_dir, local_files_only=True)
onnx_session = None
reference_model = None
if args.onnx_model is not None:
try:
import onnxruntime as ort
@@ -191,7 +266,15 @@ def evaluate(args: argparse.Namespace) -> dict[str, Any]:
device = torch.device("cpu")
model = None
runtime_name = "onnxruntime-cpu"
if args.compare_pytorch:
reference_model = AutoModelForSequenceClassification.from_pretrained(
args.model_dir,
local_files_only=True,
).to(device)
reference_model.eval()
else:
if args.compare_pytorch:
raise DataError("--compare-pytorch requires --onnx-model")
device = _device(torch, args.device)
model = AutoModelForSequenceClassification.from_pretrained(
args.model_dir, local_files_only=True
@@ -216,6 +299,7 @@ def evaluate(args: argparse.Namespace) -> dict[str, Any]:
probabilities: list[float] = []
confidences: list[str] = []
margins: list[float] = []
reference_predictions: list[int] = []
inference_batch_size = 1 if onnx_session is not None else args.batch_size
with torch.inference_mode():
for start in range(0, len(records), inference_batch_size):
@@ -226,6 +310,11 @@ def evaluate(args: argparse.Namespace) -> dict[str, Any]:
torch,
)
logits = predict_logits(encoded) / temperature
if reference_model is not None:
reference_logits = reference_model(**encoded).logits
reference_predictions.extend(
reference_logits.argmax(dim=-1).tolist()
)
distribution = torch.softmax(logits, dim=-1)
top = torch.topk(distribution, k=2, dim=-1)
batch_probabilities = top.values[:, 0].tolist()
@@ -419,6 +508,13 @@ def evaluate(args: argparse.Namespace) -> dict[str, Any]:
},
"misclassifications": misclassifications,
}
if reference_model is not None:
report["pytorchParity"] = prediction_agreement(
records,
actual,
reference_predictions,
predicted,
)
write_json(args.report, report)
return report
@@ -431,6 +527,11 @@ def build_parser() -> argparse.ArgumentParser:
type=Path,
help="score a fixed-shape ONNX artifact instead of the PyTorch checkpoint",
)
parser.add_argument(
"--compare-pytorch",
action="store_true",
help="include label-level drift from --model-dir when scoring ONNX",
)
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)
+83 -3
View File
@@ -5,15 +5,23 @@ from __future__ import annotations
import argparse
import hashlib
import math
import shutil
import sys
import tempfile
from collections import Counter, defaultdict
from pathlib import Path
from typing import Any, Sequence
import numpy as np
from purpose_data import DataError, load_jsonl, normalize_prompt, write_json
from purpose_data import (
DataError,
load_jsonl,
normalize_prompt,
prompt_hash,
write_json,
)
from train import (
HEAD_TOKENS,
MAX_LENGTH,
@@ -41,6 +49,59 @@ GOLDEN_PROMPTS = (
)
def stratified_calibration_sample(
records: Sequence[dict[str, Any]],
count: int,
*,
seed: int,
) -> list[dict[str, Any]]:
"""Select an exact, deterministic purpose/slice/language calibration sample."""
if not records:
raise DataError("cannot calibrate quantization from an empty validation split")
if count <= 0:
raise DataError("calibration record count must be positive")
target = min(count, len(records))
groups: dict[tuple[str, str, str], list[dict[str, Any]]] = defaultdict(list)
for record in records:
language = str(record.get("lang", "unknown")).split("-", 1)[0].casefold()
key = (
str(record.get("purpose", "unknown")),
str(record.get("slice", "unknown")),
language,
)
groups[key].append(record)
allocations = {}
remainders = []
allocated = 0
for key in sorted(groups):
quota = len(groups[key]) * target / len(records)
base = math.floor(quota)
allocations[key] = base
allocated += base
tie_break = hashlib.sha256(f"{seed}\0{key}".encode("utf-8")).hexdigest()
remainders.append((quota - base, tie_break, key))
for _, _, key in sorted(remainders, reverse=True)[: target - allocated]:
allocations[key] += 1
selected = []
for key in sorted(groups):
ranked = sorted(
groups[key],
key=lambda record: hashlib.sha256(
f"{seed}\0{prompt_hash(record['prompt'])}".encode("utf-8")
).hexdigest(),
)
selected.extend(ranked[: allocations[key]])
return sorted(
selected,
key=lambda record: hashlib.sha256(
f"{seed + 1}\0{prompt_hash(record['prompt'])}".encode("utf-8")
).hexdigest(),
)
def _sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
@@ -216,11 +277,16 @@ def export(args: argparse.Namespace) -> dict[str, Any]:
onnx.save(fp16_model, fp16_path)
validation = load_jsonl(args.validation)
calibration_samples = stratified_calibration_sample(
validation,
args.calibration_records,
seed=args.calibration_seed,
)
class Reader(CalibrationDataReader):
def __init__(self) -> None:
self.index = 0
self.samples = validation[: args.calibration_records]
self.samples = calibration_samples
def get_next(self) -> dict[str, np.ndarray] | None:
if self.index >= len(self.samples):
@@ -281,7 +347,20 @@ def export(args: argparse.Namespace) -> dict[str, Any]:
"opset": args.opset,
"fixedInputShape": [1, MAX_LENGTH],
"inputNames": input_names,
"calibrationRecords": min(args.calibration_records, len(validation)),
"calibrationRecords": len(calibration_samples),
"calibrationSeed": args.calibration_seed,
"calibrationSample": {
"strategy": "stratified by purpose, slice, and primary language",
"purposeCounts": dict(
sorted(Counter(item["purpose"] for item in calibration_samples).items())
),
"sliceCounts": dict(
sorted(Counter(item["slice"] for item in calibration_samples).items())
),
"promptHashes": sorted(
prompt_hash(item["prompt"]) for item in calibration_samples
),
},
"shippingArtifact": "int8QDQ",
"shippingBudgetBytes": args.shipping_budget_bytes,
"shippingBudgetPassed": int8_size <= args.shipping_budget_bytes,
@@ -303,6 +382,7 @@ def build_parser() -> argparse.ArgumentParser:
parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR)
parser.add_argument("--opset", type=int, default=17)
parser.add_argument("--calibration-records", type=int, default=256)
parser.add_argument("--calibration-seed", type=int, default=20260730)
parser.add_argument(
"--shipping-budget-bytes", type=int, default=SHIPPING_BUDGET_BYTES
)
+28
View File
@@ -23,5 +23,33 @@ class TierDriftTests(unittest.TestCase):
self.assertLessEqual(report["maximumTierDrift"], 1)
class PredictionAgreementTests(unittest.TestCase):
def test_reports_accuracy_transitions_and_scorable_agreement(self):
records = [
{"prompt": "one", "slice": "core"},
{"prompt": "two", "slice": "boundary"},
{"prompt": "three", "slice": "vague-eval"},
{"prompt": "four", "slice": "core"},
]
report = purpose_eval.prediction_agreement(
records,
actual=[0, 1, 2, 3],
reference=[0, 0, 3, 4],
candidate=[1, 1, 4, 5],
)
self.assertEqual(0.0, report["labelAgreement"])
self.assertEqual(0.0, report["scoredLabelAgreement"])
self.assertEqual(
{
"correctToIncorrect": 1,
"differentIncorrectLabel": 2,
"incorrectToCorrect": 1,
},
report["transitionCounts"],
)
self.assertEqual(2, report["bySlice"]["core"]["records"])
self.assertEqual(4, len(report["disagreements"]))
if __name__ == "__main__":
unittest.main()
+51
View File
@@ -0,0 +1,51 @@
import sys
import unittest
from collections import Counter
from pathlib import Path
MODULE_DIR = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(MODULE_DIR))
import export
class CalibrationSampleTests(unittest.TestCase):
def test_sample_is_exact_deterministic_and_stratified(self):
records = []
for index in range(100):
records.append(
{
"prompt": f"prompt {index}",
"purpose": "planning" if index < 80 else "writing",
"slice": "core" if index % 2 else "boundary",
"lang": "en" if index % 5 else "fr",
}
)
first = export.stratified_calibration_sample(records, 25, seed=42)
second = export.stratified_calibration_sample(records, 25, seed=42)
self.assertEqual(25, len(first))
self.assertEqual(
[item["prompt"] for item in first],
[item["prompt"] for item in second],
)
purposes = Counter(item["purpose"] for item in first)
self.assertEqual({"planning": 20, "writing": 5}, dict(purposes))
def test_sample_caps_at_population(self):
records = [
{
"prompt": "one",
"purpose": "planning",
"slice": "core",
"lang": "en",
}
]
self.assertEqual(
records,
export.stratified_calibration_sample(records, 10, seed=1),
)
if __name__ == "__main__":
unittest.main()
+10
View File
@@ -12,6 +12,16 @@ import train
class MetricsTests(unittest.TestCase):
def test_boundary_training_weight_is_opt_in(self):
self.assertEqual(
2.0,
train.training_weight({"slice": "boundary"}, boundary_weight=2.0),
)
self.assertEqual(
1.0,
train.training_weight({"slice": "core"}, boundary_weight=2.0),
)
def test_classification_metrics_include_every_label(self):
actual = list(range(8))
predicted = [0, 1, 2, 3, 4, 5, 6, 0]
+103 -13
View File
@@ -31,6 +31,12 @@ def prepare_text(prompt: str) -> str:
return normalize_prompt(prompt)
def training_weight(record: dict[str, Any], boundary_weight: float) -> float:
"""Return the loss weight for one training record."""
return boundary_weight if record.get("slice") == "boundary" else 1.0
def encode_fixed_shape(
tokenizer: Any,
texts: Sequence[str],
@@ -284,6 +290,7 @@ def _evaluate(torch: Any, model: Any, loader: Any, device: Any) -> tuple[Any, An
with torch.inference_mode():
for batch in loader:
labels = batch.pop("labels")
batch.pop("sample_weights", None)
inputs = {key: value.to(device) for key, value in batch.items()}
logits = model(**inputs).logits.cpu()
all_logits.append(logits)
@@ -317,6 +324,17 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
validation_records = validation_records[: args.max_validation_records]
output_dir: Path = args.output_dir
local_model = Path(args.model).expanduser()
if local_model.exists():
try:
local_model.resolve().relative_to(output_dir.resolve())
except ValueError:
pass
else:
raise DataError(
"local --model must not be inside --output-dir; overwrite could "
"destroy the continuation checkpoint"
)
if output_dir.exists() and any(output_dir.iterdir()):
if not args.overwrite_output:
raise DataError(
@@ -329,16 +347,22 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
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()}
model_revision = None if local_model.exists() else args.model_revision
pretrained_options = (
{"local_files_only": True}
if local_model.exists()
else {"revision": model_revision}
)
tokenizer = AutoTokenizer.from_pretrained(
args.model, revision=args.model_revision, use_fast=True
args.model, use_fast=True, **pretrained_options
)
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,
**pretrained_options,
)
config = model.config
if getattr(config, "hidden_size", None) != 384 or getattr(
@@ -364,14 +388,19 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
def __len__(self) -> int:
return len(self.records)
def __getitem__(self, index: int) -> tuple[str, int]:
def __getitem__(self, index: int) -> tuple[str, int, float]:
record = self.records[index]
return prepare_text(record["prompt"]), label_to_id[record["purpose"]]
return (
prepare_text(record["prompt"]),
label_to_id[record["purpose"]],
training_weight(record, args.boundary_weight),
)
def collate(items: Sequence[tuple[str, int]]) -> dict[str, Any]:
texts, labels = zip(*items)
def collate(items: Sequence[tuple[str, int, float]]) -> dict[str, Any]:
texts, labels, weights = zip(*items)
encoded = encode_fixed_shape(tokenizer, list(texts), torch)
encoded["labels"] = torch.tensor(labels, dtype=torch.long)
encoded["sample_weights"] = torch.tensor(weights, dtype=torch.float)
return encoded
generator = torch.Generator()
@@ -397,6 +426,36 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
[record.get("slice") != "vague-eval" for record in validation_records],
dtype=torch.bool,
)
initial_logits, initial_labels = _evaluate(
torch,
model,
validation_loader,
device,
)
initial_predictions = initial_logits.argmax(dim=-1).tolist()
initial_metrics = classification_metrics(
initial_labels[validation_scorable].tolist(),
[
prediction
for prediction, scorable in zip(
initial_predictions,
validation_scorable.tolist(),
)
if scorable
],
)
best_accuracy = initial_metrics["accuracy"]
epochs_without_improvement = 0
stopped_early = False
history = []
best_dir = output_dir / "model"
model.save_pretrained(best_dir, safe_serialization=True)
tokenizer.save_pretrained(best_dir)
print(
f"epoch 0: validation_accuracy={initial_metrics['accuracy']:.4%} "
f"macro_recall={initial_metrics['macroRecall']:.4%}",
flush=True,
)
optimizer = torch.optim.AdamW(
model.parameters(), lr=args.learning_rate, weight_decay=args.weight_decay
@@ -411,17 +470,23 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
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
labels = batch.pop("labels").to(device)
sample_weights = batch.pop("sample_weights").to(device)
inputs = {key: value.to(device) for key, value in batch.items()}
per_record_loss = torch.nn.functional.cross_entropy(
model(**inputs).logits,
labels,
reduction="none",
)
loss = (
(per_record_loss * sample_weights).sum() / sample_weights.sum()
) / args.gradient_accumulation_steps
loss.backward()
running_loss += float(loss.item()) * args.gradient_accumulation_steps
should_update = (
@@ -454,10 +519,23 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
f"macro_recall={metrics['macroRecall']:.4%}",
flush=True,
)
if metrics["accuracy"] > best_accuracy:
improvement = metrics["accuracy"] - best_accuracy
if improvement > args.minimum_improvement:
best_accuracy = metrics["accuracy"]
epochs_without_improvement = 0
model.save_pretrained(best_dir, safe_serialization=True)
tokenizer.save_pretrained(best_dir)
else:
epochs_without_improvement += 1
if epochs_without_improvement >= args.early_stopping_patience:
stopped_early = True
print(
f"early stopping after epoch {epoch}: no validation improvement "
f"greater than {args.minimum_improvement:.4%} for "
f"{args.early_stopping_patience} epoch(s)",
flush=True,
)
break
model = AutoModelForSequenceClassification.from_pretrained(best_dir).to(device)
logits, labels = _evaluate(torch, model, validation_loader, device)
@@ -497,7 +575,7 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
metrics = {
"modelVersion": "purpose-lite-v1",
"baseModel": args.model,
"baseModelRevision": args.model_revision,
"baseModelRevision": model_revision or "local-checkpoint",
"fixedInputShape": [1, MAX_LENGTH],
"truncation": {
"strategy": "head-tail-pair",
@@ -507,12 +585,16 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
"device": str(device),
"trainingSeconds": time.perf_counter() - started,
"trainRecords": len(train_records),
"boundaryTrainingWeight": args.boundary_weight,
"validationRecords": len(validation_records),
"scoredValidationRecords": int(validation_scorable.sum().item()),
"vagueAbstentionValidationRecords": int(
(~validation_scorable).sum().item()
),
"bestValidationAccuracy": best_accuracy,
"initialValidation": initial_metrics,
"epochsCompleted": len(history),
"stoppedEarly": stopped_early,
"bestValidation": classification_metrics(
labels[validation_scorable].tolist(),
[
@@ -555,6 +637,9 @@ def build_parser() -> argparse.ArgumentParser:
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("--early-stopping-patience", type=int, default=2)
parser.add_argument("--minimum-improvement", type=float, default=0.0005)
parser.add_argument("--boundary-weight", type=float, default=1.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)
@@ -576,10 +661,15 @@ def main(argv: Sequence[str] | None = None) -> int:
"batch_size",
"eval_batch_size",
"gradient_accumulation_steps",
"early_stopping_patience",
):
_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 args.minimum_improvement < 0.0:
parser.error("--minimum-improvement must be non-negative")
if args.boundary_weight <= 0.0:
parser.error("--boundary-weight must be positive")
if not 0.0 < args.accepted_precision <= args.high_precision <= 1.0:
parser.error(
"precision targets must satisfy 0 < accepted <= high <= 1"