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
|
||||
`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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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):
|
||||
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]
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user