Merge nucleic/lucid-north-quail-rnvt into dev
This commit is contained in:
@@ -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())
|
||||
Reference in New Issue
Block a user