Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -72,7 +72,23 @@ platform selector, then install `requirements-base.txt`.
|
|||||||
Training writes a local checkpoint, `calibration.json`, and `metrics.json` under
|
Training writes a local checkpoint, `calibration.json`, and `metrics.json` under
|
||||||
`outputs/purpose-lite-v1/`. It selects checkpoints and fits temperature on label-scorable
|
`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`
|
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:
|
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/
|
`export.py` emits fixed-shape opset-17 fp16 and int8-QDQ graphs, a tokenizer/
|
||||||
normalization contract, golden tokenizations, shared calibration config, graph checks,
|
normalization contract, golden tokenizations, shared calibration config, graph checks,
|
||||||
artifact hashes, and a size report. The int8 graph is the ≤25 MiB shipping candidate;
|
artifact hashes, and a size report. Its default 256-record quantization calibration sample
|
||||||
the fp16 graph remains the accelerator-oriented conversion input.
|
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
|
## Audit curation
|
||||||
|
|
||||||
|
|||||||
@@ -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]:
|
def evaluate(args: argparse.Namespace) -> dict[str, Any]:
|
||||||
try:
|
try:
|
||||||
import torch
|
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)}
|
label_to_id = {label: index for index, label in enumerate(LABELS)}
|
||||||
tokenizer = AutoTokenizer.from_pretrained(args.model_dir, local_files_only=True)
|
tokenizer = AutoTokenizer.from_pretrained(args.model_dir, local_files_only=True)
|
||||||
onnx_session = None
|
onnx_session = None
|
||||||
|
reference_model = None
|
||||||
if args.onnx_model is not None:
|
if args.onnx_model is not None:
|
||||||
try:
|
try:
|
||||||
import onnxruntime as ort
|
import onnxruntime as ort
|
||||||
@@ -191,7 +266,15 @@ def evaluate(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
device = torch.device("cpu")
|
device = torch.device("cpu")
|
||||||
model = None
|
model = None
|
||||||
runtime_name = "onnxruntime-cpu"
|
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:
|
else:
|
||||||
|
if args.compare_pytorch:
|
||||||
|
raise DataError("--compare-pytorch requires --onnx-model")
|
||||||
device = _device(torch, args.device)
|
device = _device(torch, args.device)
|
||||||
model = AutoModelForSequenceClassification.from_pretrained(
|
model = AutoModelForSequenceClassification.from_pretrained(
|
||||||
args.model_dir, local_files_only=True
|
args.model_dir, local_files_only=True
|
||||||
@@ -216,6 +299,7 @@ def evaluate(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
probabilities: list[float] = []
|
probabilities: list[float] = []
|
||||||
confidences: list[str] = []
|
confidences: list[str] = []
|
||||||
margins: list[float] = []
|
margins: list[float] = []
|
||||||
|
reference_predictions: list[int] = []
|
||||||
inference_batch_size = 1 if onnx_session is not None else args.batch_size
|
inference_batch_size = 1 if onnx_session is not None else args.batch_size
|
||||||
with torch.inference_mode():
|
with torch.inference_mode():
|
||||||
for start in range(0, len(records), inference_batch_size):
|
for start in range(0, len(records), inference_batch_size):
|
||||||
@@ -226,6 +310,11 @@ def evaluate(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
torch,
|
torch,
|
||||||
)
|
)
|
||||||
logits = predict_logits(encoded) / temperature
|
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)
|
distribution = torch.softmax(logits, dim=-1)
|
||||||
top = torch.topk(distribution, k=2, dim=-1)
|
top = torch.topk(distribution, k=2, dim=-1)
|
||||||
batch_probabilities = top.values[:, 0].tolist()
|
batch_probabilities = top.values[:, 0].tolist()
|
||||||
@@ -419,6 +508,13 @@ def evaluate(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
},
|
},
|
||||||
"misclassifications": misclassifications,
|
"misclassifications": misclassifications,
|
||||||
}
|
}
|
||||||
|
if reference_model is not None:
|
||||||
|
report["pytorchParity"] = prediction_agreement(
|
||||||
|
records,
|
||||||
|
actual,
|
||||||
|
reference_predictions,
|
||||||
|
predicted,
|
||||||
|
)
|
||||||
write_json(args.report, report)
|
write_json(args.report, report)
|
||||||
return report
|
return report
|
||||||
|
|
||||||
@@ -431,6 +527,11 @@ def build_parser() -> argparse.ArgumentParser:
|
|||||||
type=Path,
|
type=Path,
|
||||||
help="score a fixed-shape ONNX artifact instead of the PyTorch checkpoint",
|
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("--calibration", type=Path, default=DEFAULT_CALIBRATION)
|
||||||
parser.add_argument("--test", type=Path, default=DEFAULT_TEST)
|
parser.add_argument("--test", type=Path, default=DEFAULT_TEST)
|
||||||
parser.add_argument("--fixtures", type=Path, default=DEFAULT_FIXTURES)
|
parser.add_argument("--fixtures", type=Path, default=DEFAULT_FIXTURES)
|
||||||
|
|||||||
@@ -5,15 +5,23 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import argparse
|
import argparse
|
||||||
import hashlib
|
import hashlib
|
||||||
|
import math
|
||||||
import shutil
|
import shutil
|
||||||
import sys
|
import sys
|
||||||
import tempfile
|
import tempfile
|
||||||
|
from collections import Counter, defaultdict
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Sequence
|
from typing import Any, Sequence
|
||||||
|
|
||||||
import numpy as np
|
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 (
|
from train import (
|
||||||
HEAD_TOKENS,
|
HEAD_TOKENS,
|
||||||
MAX_LENGTH,
|
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:
|
def _sha256(path: Path) -> str:
|
||||||
digest = hashlib.sha256()
|
digest = hashlib.sha256()
|
||||||
with path.open("rb") as handle:
|
with path.open("rb") as handle:
|
||||||
@@ -216,11 +277,16 @@ def export(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
onnx.save(fp16_model, fp16_path)
|
onnx.save(fp16_model, fp16_path)
|
||||||
|
|
||||||
validation = load_jsonl(args.validation)
|
validation = load_jsonl(args.validation)
|
||||||
|
calibration_samples = stratified_calibration_sample(
|
||||||
|
validation,
|
||||||
|
args.calibration_records,
|
||||||
|
seed=args.calibration_seed,
|
||||||
|
)
|
||||||
|
|
||||||
class Reader(CalibrationDataReader):
|
class Reader(CalibrationDataReader):
|
||||||
def __init__(self) -> None:
|
def __init__(self) -> None:
|
||||||
self.index = 0
|
self.index = 0
|
||||||
self.samples = validation[: args.calibration_records]
|
self.samples = calibration_samples
|
||||||
|
|
||||||
def get_next(self) -> dict[str, np.ndarray] | None:
|
def get_next(self) -> dict[str, np.ndarray] | None:
|
||||||
if self.index >= len(self.samples):
|
if self.index >= len(self.samples):
|
||||||
@@ -281,7 +347,20 @@ def export(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
"opset": args.opset,
|
"opset": args.opset,
|
||||||
"fixedInputShape": [1, MAX_LENGTH],
|
"fixedInputShape": [1, MAX_LENGTH],
|
||||||
"inputNames": input_names,
|
"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",
|
"shippingArtifact": "int8QDQ",
|
||||||
"shippingBudgetBytes": args.shipping_budget_bytes,
|
"shippingBudgetBytes": args.shipping_budget_bytes,
|
||||||
"shippingBudgetPassed": int8_size <= 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("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR)
|
||||||
parser.add_argument("--opset", type=int, default=17)
|
parser.add_argument("--opset", type=int, default=17)
|
||||||
parser.add_argument("--calibration-records", type=int, default=256)
|
parser.add_argument("--calibration-records", type=int, default=256)
|
||||||
|
parser.add_argument("--calibration-seed", type=int, default=20260730)
|
||||||
parser.add_argument(
|
parser.add_argument(
|
||||||
"--shipping-budget-bytes", type=int, default=SHIPPING_BUDGET_BYTES
|
"--shipping-budget-bytes", type=int, default=SHIPPING_BUDGET_BYTES
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -23,5 +23,33 @@ class TierDriftTests(unittest.TestCase):
|
|||||||
self.assertLessEqual(report["maximumTierDrift"], 1)
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -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()
|
||||||
@@ -12,6 +12,16 @@ import train
|
|||||||
|
|
||||||
|
|
||||||
class MetricsTests(unittest.TestCase):
|
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):
|
def test_classification_metrics_include_every_label(self):
|
||||||
actual = list(range(8))
|
actual = list(range(8))
|
||||||
predicted = [0, 1, 2, 3, 4, 5, 6, 0]
|
predicted = [0, 1, 2, 3, 4, 5, 6, 0]
|
||||||
|
|||||||
@@ -31,6 +31,12 @@ def prepare_text(prompt: str) -> str:
|
|||||||
return normalize_prompt(prompt)
|
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(
|
def encode_fixed_shape(
|
||||||
tokenizer: Any,
|
tokenizer: Any,
|
||||||
texts: Sequence[str],
|
texts: Sequence[str],
|
||||||
@@ -284,6 +290,7 @@ def _evaluate(torch: Any, model: Any, loader: Any, device: Any) -> tuple[Any, An
|
|||||||
with torch.inference_mode():
|
with torch.inference_mode():
|
||||||
for batch in loader:
|
for batch in loader:
|
||||||
labels = batch.pop("labels")
|
labels = batch.pop("labels")
|
||||||
|
batch.pop("sample_weights", None)
|
||||||
inputs = {key: value.to(device) for key, value in batch.items()}
|
inputs = {key: value.to(device) for key, value in batch.items()}
|
||||||
logits = model(**inputs).logits.cpu()
|
logits = model(**inputs).logits.cpu()
|
||||||
all_logits.append(logits)
|
all_logits.append(logits)
|
||||||
@@ -317,6 +324,17 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
validation_records = validation_records[: args.max_validation_records]
|
validation_records = validation_records[: args.max_validation_records]
|
||||||
|
|
||||||
output_dir: Path = args.output_dir
|
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 output_dir.exists() and any(output_dir.iterdir()):
|
||||||
if not args.overwrite_output:
|
if not args.overwrite_output:
|
||||||
raise DataError(
|
raise DataError(
|
||||||
@@ -329,16 +347,22 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
device = _select_device(torch, args.device)
|
device = _select_device(torch, args.device)
|
||||||
label_to_id = {label: index for index, label in enumerate(LABELS)}
|
label_to_id = {label: index for index, label in enumerate(LABELS)}
|
||||||
id_to_label = {index: label for label, index in label_to_id.items()}
|
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(
|
tokenizer = AutoTokenizer.from_pretrained(
|
||||||
args.model, revision=args.model_revision, use_fast=True
|
args.model, use_fast=True, **pretrained_options
|
||||||
)
|
)
|
||||||
model = AutoModelForSequenceClassification.from_pretrained(
|
model = AutoModelForSequenceClassification.from_pretrained(
|
||||||
args.model,
|
args.model,
|
||||||
revision=args.model_revision,
|
|
||||||
num_labels=len(LABELS),
|
num_labels=len(LABELS),
|
||||||
label2id=label_to_id,
|
label2id=label_to_id,
|
||||||
id2label=id_to_label,
|
id2label=id_to_label,
|
||||||
ignore_mismatched_sizes=True,
|
ignore_mismatched_sizes=True,
|
||||||
|
**pretrained_options,
|
||||||
)
|
)
|
||||||
config = model.config
|
config = model.config
|
||||||
if getattr(config, "hidden_size", None) != 384 or getattr(
|
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:
|
def __len__(self) -> int:
|
||||||
return len(self.records)
|
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]
|
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]:
|
def collate(items: Sequence[tuple[str, int, float]]) -> dict[str, Any]:
|
||||||
texts, labels = zip(*items)
|
texts, labels, weights = zip(*items)
|
||||||
encoded = encode_fixed_shape(tokenizer, list(texts), torch)
|
encoded = encode_fixed_shape(tokenizer, list(texts), torch)
|
||||||
encoded["labels"] = torch.tensor(labels, dtype=torch.long)
|
encoded["labels"] = torch.tensor(labels, dtype=torch.long)
|
||||||
|
encoded["sample_weights"] = torch.tensor(weights, dtype=torch.float)
|
||||||
return encoded
|
return encoded
|
||||||
|
|
||||||
generator = torch.Generator()
|
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],
|
[record.get("slice") != "vague-eval" for record in validation_records],
|
||||||
dtype=torch.bool,
|
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(
|
optimizer = torch.optim.AdamW(
|
||||||
model.parameters(), lr=args.learning_rate, weight_decay=args.weight_decay
|
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,
|
num_training_steps=total_steps,
|
||||||
)
|
)
|
||||||
|
|
||||||
best_accuracy = -1.0
|
|
||||||
history = []
|
|
||||||
best_dir = output_dir / "model"
|
|
||||||
started = time.perf_counter()
|
started = time.perf_counter()
|
||||||
for epoch in range(1, args.epochs + 1):
|
for epoch in range(1, args.epochs + 1):
|
||||||
model.train()
|
model.train()
|
||||||
optimizer.zero_grad(set_to_none=True)
|
optimizer.zero_grad(set_to_none=True)
|
||||||
running_loss = 0.0
|
running_loss = 0.0
|
||||||
for step, batch in enumerate(train_loader, 1):
|
for step, batch in enumerate(train_loader, 1):
|
||||||
batch = {key: value.to(device) for key, value in batch.items()}
|
labels = batch.pop("labels").to(device)
|
||||||
loss = model(**batch).loss / args.gradient_accumulation_steps
|
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()
|
loss.backward()
|
||||||
running_loss += float(loss.item()) * args.gradient_accumulation_steps
|
running_loss += float(loss.item()) * args.gradient_accumulation_steps
|
||||||
should_update = (
|
should_update = (
|
||||||
@@ -454,10 +519,23 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
f"macro_recall={metrics['macroRecall']:.4%}",
|
f"macro_recall={metrics['macroRecall']:.4%}",
|
||||||
flush=True,
|
flush=True,
|
||||||
)
|
)
|
||||||
if metrics["accuracy"] > best_accuracy:
|
improvement = metrics["accuracy"] - best_accuracy
|
||||||
|
if improvement > args.minimum_improvement:
|
||||||
best_accuracy = metrics["accuracy"]
|
best_accuracy = metrics["accuracy"]
|
||||||
|
epochs_without_improvement = 0
|
||||||
model.save_pretrained(best_dir, safe_serialization=True)
|
model.save_pretrained(best_dir, safe_serialization=True)
|
||||||
tokenizer.save_pretrained(best_dir)
|
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)
|
model = AutoModelForSequenceClassification.from_pretrained(best_dir).to(device)
|
||||||
logits, labels = _evaluate(torch, model, validation_loader, device)
|
logits, labels = _evaluate(torch, model, validation_loader, device)
|
||||||
@@ -497,7 +575,7 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
metrics = {
|
metrics = {
|
||||||
"modelVersion": "purpose-lite-v1",
|
"modelVersion": "purpose-lite-v1",
|
||||||
"baseModel": args.model,
|
"baseModel": args.model,
|
||||||
"baseModelRevision": args.model_revision,
|
"baseModelRevision": model_revision or "local-checkpoint",
|
||||||
"fixedInputShape": [1, MAX_LENGTH],
|
"fixedInputShape": [1, MAX_LENGTH],
|
||||||
"truncation": {
|
"truncation": {
|
||||||
"strategy": "head-tail-pair",
|
"strategy": "head-tail-pair",
|
||||||
@@ -507,12 +585,16 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
"device": str(device),
|
"device": str(device),
|
||||||
"trainingSeconds": time.perf_counter() - started,
|
"trainingSeconds": time.perf_counter() - started,
|
||||||
"trainRecords": len(train_records),
|
"trainRecords": len(train_records),
|
||||||
|
"boundaryTrainingWeight": args.boundary_weight,
|
||||||
"validationRecords": len(validation_records),
|
"validationRecords": len(validation_records),
|
||||||
"scoredValidationRecords": int(validation_scorable.sum().item()),
|
"scoredValidationRecords": int(validation_scorable.sum().item()),
|
||||||
"vagueAbstentionValidationRecords": int(
|
"vagueAbstentionValidationRecords": int(
|
||||||
(~validation_scorable).sum().item()
|
(~validation_scorable).sum().item()
|
||||||
),
|
),
|
||||||
"bestValidationAccuracy": best_accuracy,
|
"bestValidationAccuracy": best_accuracy,
|
||||||
|
"initialValidation": initial_metrics,
|
||||||
|
"epochsCompleted": len(history),
|
||||||
|
"stoppedEarly": stopped_early,
|
||||||
"bestValidation": classification_metrics(
|
"bestValidation": classification_metrics(
|
||||||
labels[validation_scorable].tolist(),
|
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("--warmup-ratio", type=float, default=0.1)
|
||||||
parser.add_argument("--max-grad-norm", type=float, default=1.0)
|
parser.add_argument("--max-grad-norm", type=float, default=1.0)
|
||||||
parser.add_argument("--workers", type=int, default=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("--high-precision", type=float, default=0.98)
|
||||||
parser.add_argument("--accepted-precision", type=float, default=0.95)
|
parser.add_argument("--accepted-precision", type=float, default=0.95)
|
||||||
parser.add_argument("--max-train-records", type=int)
|
parser.add_argument("--max-train-records", type=int)
|
||||||
@@ -576,10 +661,15 @@ def main(argv: Sequence[str] | None = None) -> int:
|
|||||||
"batch_size",
|
"batch_size",
|
||||||
"eval_batch_size",
|
"eval_batch_size",
|
||||||
"gradient_accumulation_steps",
|
"gradient_accumulation_steps",
|
||||||
|
"early_stopping_patience",
|
||||||
):
|
):
|
||||||
_positive(parser, f"--{name.replace('_', '-')}", getattr(args, name))
|
_positive(parser, f"--{name.replace('_', '-')}", getattr(args, name))
|
||||||
if not 0.0 <= args.warmup_ratio < 1.0:
|
if not 0.0 <= args.warmup_ratio < 1.0:
|
||||||
parser.error("--warmup-ratio must be in [0, 1)")
|
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:
|
if not 0.0 < args.accepted_precision <= args.high_precision <= 1.0:
|
||||||
parser.error(
|
parser.error(
|
||||||
"precision targets must satisfy 0 < accepted <= high <= 1"
|
"precision targets must satisfy 0 < accepted <= high <= 1"
|
||||||
|
|||||||
Reference in New Issue
Block a user