Files
nucleic-purpose-classifier/train.py
T

830 lines
31 KiB
Python
Raw Normal View History

#!/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
HEAD_TAIL_SPECIAL_TOKENS = 3
HEAD_TOKENS = (MAX_LENGTH - HEAD_TAIL_SPECIAL_TOKENS + 1) // 2
TAIL_TOKENS = MAX_LENGTH - HEAD_TAIL_SPECIAL_TOKENS - HEAD_TOKENS
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 enable_quantization_aware_training(torch: Any, model: Any) -> dict[str, int]:
"""Mirror the export graph's int8 policy with straight-through fake quantization.
ONNX Runtime emits per-tensor uint8 embedding weights, per-channel symmetric int8
linear weights, and per-tensor uint8 activations. The replacement modules keep the
original parameter names, so the selected checkpoint loads as an ordinary Transformers
model for export after fake-quantization-aware fine-tuning.
"""
functional = torch.nn.functional
def affine_parameters(value: Any) -> tuple[float, int]:
detached = value.detach().float()
# ONNX Runtime extends affine calibration ranges to include exact zero.
minimum = min(0.0, float(detached.amin().item()))
maximum = max(0.0, float(detached.amax().item()))
scale = max((maximum - minimum) / 255.0, torch.finfo(torch.float32).eps)
zero_point = max(0, min(255, round(-minimum / scale)))
return scale, zero_point
def fake_quantize_activation(value: Any) -> Any:
scale, zero_point = affine_parameters(value)
return torch.fake_quantize_per_tensor_affine(
value,
scale,
zero_point,
0,
255,
)
def fake_quantize_linear_weight(weight: Any) -> Any:
detached = weight.detach().float()
scales = detached.abs().amax(dim=1).div(127.0).clamp_min(
torch.finfo(torch.float32).eps
)
zero_points = torch.zeros_like(scales, dtype=torch.int32)
return torch.fake_quantize_per_channel_affine(
weight,
scales,
zero_points,
0,
-127,
127,
)
class QATLinear(torch.nn.Linear):
def forward(self, value: Any) -> Any:
result = functional.linear(
fake_quantize_activation(value),
fake_quantize_linear_weight(self.weight),
self.bias,
)
return fake_quantize_activation(result)
class QATEmbedding(torch.nn.Embedding):
def forward(self, indexes: Any) -> Any:
embedded = functional.embedding(
indexes,
self.weight,
self.padding_idx,
self.max_norm,
self.norm_type,
self.scale_grad_by_freq,
self.sparse,
)
# Quantize only the selected rows using the full table's scale. This is
# numerically equivalent to dequantizing the whole table before Gather but
# avoids materializing a 30k x 384 fake-quantized embedding every batch.
weight_scale, weight_zero_point = affine_parameters(self.weight)
embedded = torch.fake_quantize_per_tensor_affine(
embedded,
weight_scale,
weight_zero_point,
0,
255,
)
return fake_quantize_activation(embedded)
counts = {"linear": 0, "embedding": 0}
def replace(parent: Any) -> None:
for name, child in list(parent.named_children()):
replacement = None
if isinstance(child, torch.nn.Linear):
replacement = QATLinear(
child.in_features,
child.out_features,
bias=child.bias is not None,
device=child.weight.device,
dtype=child.weight.dtype,
)
counts["linear"] += 1
elif isinstance(child, torch.nn.Embedding):
replacement = QATEmbedding(
child.num_embeddings,
child.embedding_dim,
padding_idx=child.padding_idx,
max_norm=child.max_norm,
norm_type=child.norm_type,
scale_grad_by_freq=child.scale_grad_by_freq,
sparse=child.sparse,
device=child.weight.device,
dtype=child.weight.dtype,
)
counts["embedding"] += 1
if replacement is not None:
replacement.weight = child.weight
if isinstance(child, torch.nn.Linear):
replacement.bias = child.bias
setattr(parent, name, replacement)
else:
replace(child)
replace(model)
return counts
def encode_fixed_shape(
tokenizer: Any,
texts: Sequence[str],
torch: Any,
) -> dict[str, Any]:
"""Tokenize to 1x128 while retaining both context and a tail-buried request.
Pasted logs and stack traces frequently put the actual ask after the context. Plain
right truncation made generated boundary examples identical even when their final
request — and therefore their label — differed. Long inputs use BERT's sentence-pair
framing: [CLS] first 63 content tokens [SEP] last 62 content tokens [SEP].
"""
normalized = [prepare_text(text) for text in texts]
raw = tokenizer(
normalized,
add_special_tokens=False,
padding=False,
truncation=False,
return_attention_mask=False,
return_token_type_ids=False,
verbose=False,
)
if not isinstance(raw.get("input_ids"), list):
raise DataError("tokenizer did not return input_ids")
if tokenizer.pad_token_id is None:
raise DataError("purpose-lite tokenizer must define a padding token")
if tokenizer.cls_token_id is None or tokenizer.sep_token_id is None:
raise DataError("purpose-lite tokenizer must define BERT CLS and SEP tokens")
if tokenizer.padding_side != "right":
raise DataError("purpose-lite tokenizer must use right padding")
input_rows: list[list[int]] = []
mask_rows: list[list[int]] = []
type_rows: list[list[int]] = []
include_token_types = "token_type_ids" in tokenizer.model_input_names
single_budget = MAX_LENGTH - tokenizer.num_special_tokens_to_add(pair=False)
pair_budget = MAX_LENGTH - tokenizer.num_special_tokens_to_add(pair=True)
if pair_budget != HEAD_TOKENS + TAIL_TOKENS:
raise DataError(
"purpose-lite tokenizer special-token layout changed; expected three "
"tokens for head-tail inputs"
)
for content in raw["input_ids"]:
if len(content) <= single_budget:
first = content
second = None
else:
first = content[:HEAD_TOKENS]
second = content[-TAIL_TOKENS:]
if second is None:
input_ids = [tokenizer.cls_token_id] + first + [tokenizer.sep_token_id]
token_types = [0] * len(input_ids)
else:
input_ids = (
[tokenizer.cls_token_id]
+ first
+ [tokenizer.sep_token_id]
+ second
+ [tokenizer.sep_token_id]
)
token_types = [0] * (len(first) + 2) + [1] * (len(second) + 1)
if len(input_ids) > MAX_LENGTH:
raise DataError("fixed-shape tokenizer exceeded its 128-token contract")
padding = MAX_LENGTH - len(input_ids)
input_rows.append(input_ids + [tokenizer.pad_token_id] * padding)
mask_rows.append([1] * len(input_ids) + [0] * padding)
if include_token_types:
type_rows.append(token_types + [0] * padding)
encoded = {
"input_ids": torch.tensor(input_rows, dtype=torch.long),
"attention_mask": torch.tensor(mask_rows, dtype=torch.long),
}
if include_token_types:
encoded["token_type_ids"] = torch.tensor(type_rows, dtype=torch.long)
return encoded
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")
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)
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
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(
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()}
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, use_fast=True, **pretrained_options
)
model = AutoModelForSequenceClassification.from_pretrained(
args.model,
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(
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]
config.purpose_classifier_truncation = {
"strategy": "head-tail-pair",
"headTokens": HEAD_TOKENS,
"tailTokens": TAIL_TOKENS,
}
config.purpose_classifier_quantization_aware_training = bool(
args.quantization_aware
)
qat_modules = {"linear": 0, "embedding": 0}
if args.quantization_aware:
qat_modules = enable_quantization_aware_training(torch, model)
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, float]:
record = self.records[index]
return (
prepare_text(record["prompt"]),
label_to_id[record["purpose"]],
training_weight(record, args.boundary_weight),
)
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()
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",
)
validation_scorable = torch.tensor(
[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
)
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,
)
started = time.perf_counter()
for epoch in range(1, args.epochs + 1):
epoch_started = time.perf_counter()
model.train()
optimizer.zero_grad(set_to_none=True)
running_loss = 0.0
for step, batch in enumerate(train_loader, 1):
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 = (
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)
if args.progress_steps and (
step % args.progress_steps == 0 or step == len(train_loader)
):
print(
f"epoch {epoch} step {step}/{len(train_loader)} "
f"mean_loss={running_loss / step:.4f} "
f"elapsed={time.perf_counter() - epoch_started:.1f}s",
flush=True,
)
logits, labels = _evaluate(torch, model, validation_loader, device)
predictions = logits.argmax(dim=-1).tolist()
scored_labels = labels[validation_scorable].tolist()
scored_predictions = [
prediction
for prediction, scorable in zip(
predictions, validation_scorable.tolist()
)
if scorable
]
metrics = classification_metrics(scored_labels, scored_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,
)
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)
temperature = _fit_temperature(
torch,
logits[validation_scorable],
labels[validation_scorable],
)
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 and scorable
for prediction, actual, scorable in zip(
predictions,
labels.tolist(),
validation_scorable.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": model_revision or "local-checkpoint",
"fixedInputShape": [1, MAX_LENGTH],
"truncation": {
"strategy": "head-tail-pair",
"headTokens": HEAD_TOKENS,
"tailTokens": TAIL_TOKENS,
},
"device": str(device),
"trainingSeconds": time.perf_counter() - started,
"trainRecords": len(train_records),
"boundaryTrainingWeight": args.boundary_weight,
"quantizationAwareTraining": args.quantization_aware,
"quantizationAwareModules": qat_modules,
"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(),
[
prediction
for prediction, scorable in zip(
predictions, validation_scorable.tolist()
)
if scorable
],
),
"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("--progress-steps", type=int, default=50)
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("--quantization-aware", action="store_true")
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",
"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.progress_steps < 0:
parser.error("--progress-steps 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"
)
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())