Files
nucleic-purpose-classifier/train_mlx.py
T

867 lines
30 KiB
Python
Raw Normal View History

#!/usr/bin/env python3
"""Fine-tune purpose-lite with MLX, using Metal by default."""
from __future__ import annotations
import argparse
import json
import math
import random
import shutil
import sys
import time
from pathlib import Path
from typing import Any, Iterator, Sequence
import numpy as np
from purpose_data import LABELS, DataError, load_jsonl, write_json
from train import (
HEAD_TOKENS,
MAX_LENGTH,
TAIL_TOKENS,
_fit_temperature,
_validate_split,
choose_confidence_thresholds,
classification_metrics,
distillation_record_keys,
expected_calibration_error,
prepare_text,
training_weight,
)
SCRIPT_DIR = Path(__file__).resolve().parent
DEFAULT_DATASET_DIR = SCRIPT_DIR / ".artifacts" / "dataset-v1"
DEFAULT_OUTPUT_DIR = SCRIPT_DIR / "outputs" / "purpose-lite-v1-mlx"
def _configure_mlx_device(mx: Any, device: str) -> None:
if device == "metal":
if not mx.metal.is_available():
raise DataError("MLX Metal training requires Apple Silicon")
mx.set_default_device(mx.gpu)
return
if device == "cpu":
mx.set_default_device(mx.cpu)
return
raise DataError(f"unsupported MLX device {device!r}")
def _load_mlx(device: str) -> tuple[Any, Any, Any]:
try:
import mlx.core as mx
import mlx.nn as nn
import mlx.optimizers as optim
except ImportError as exc:
raise DataError(
"MLX training requires requirements-mlx.txt"
) from exc
_configure_mlx_device(mx, device)
return mx, nn, optim
def encode_fixed_shape_numpy(
tokenizer: Any,
texts: Sequence[str],
) -> dict[str, np.ndarray]:
"""Apply the same fixed 128-token head-tail contract as train.py."""
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": np.asarray(input_rows, dtype=np.int32),
"attention_mask": np.asarray(mask_rows, dtype=np.int32),
}
if include_token_types:
encoded["token_type_ids"] = np.asarray(type_rows, dtype=np.int32)
return encoded
def _encode_records(
tokenizer: Any,
records: Sequence[dict[str, Any]],
*,
chunk_size: int = 256,
) -> dict[str, np.ndarray]:
chunks: dict[str, list[np.ndarray]] = {}
for start in range(0, len(records), chunk_size):
encoded = encode_fixed_shape_numpy(
tokenizer,
[record["prompt"] for record in records[start : start + chunk_size]],
)
for key, value in encoded.items():
chunks.setdefault(key, []).append(value)
return {key: np.concatenate(values) for key, values in chunks.items()}
def _batch_indexes(
size: int,
batch_size: int,
*,
permutation: np.ndarray | None = None,
) -> Iterator[np.ndarray]:
indexes = permutation if permutation is not None else np.arange(size)
for start in range(0, size, batch_size):
yield indexes[start : start + batch_size]
def _mlx_batch(
mx: Any,
encoded: dict[str, np.ndarray],
indexes: np.ndarray,
) -> dict[str, Any]:
return {key: mx.array(value[indexes]) for key, value in encoded.items()}
def _evaluate(
mx: Any,
model: Any,
encoded: dict[str, np.ndarray],
labels: np.ndarray,
batch_size: int,
) -> tuple[np.ndarray, np.ndarray]:
model.eval()
logits: list[np.ndarray] = []
for indexes in _batch_indexes(len(labels), batch_size):
batch = _mlx_batch(mx, encoded, indexes)
output = model(**batch)
mx.eval(output)
logits.append(np.asarray(output))
return np.concatenate(logits), labels.copy()
def _teacher_cache(
path: Path,
train_records: Sequence[dict[str, Any]],
validation_records: Sequence[dict[str, Any]],
) -> tuple[np.ndarray, np.ndarray]:
try:
import torch
except ImportError as exc:
raise DataError(
"loading the existing teacher cache requires PyTorch"
) from exc
if not path.is_file():
raise DataError(f"{path}: distillation cache is missing")
try:
cache = torch.load(path, map_location="cpu", weights_only=True)
if cache["schemaVersion"] != 1 or cache["labels"] != list(LABELS):
raise DataError("distillation cache contract does not match purpose-lite")
if cache["trainRecordKeys"][: len(train_records)] != distillation_record_keys(
train_records
):
raise DataError("distillation cache does not match the training split")
if cache["validationRecordKeys"][
: len(validation_records)
] != distillation_record_keys(validation_records):
raise DataError("distillation cache does not match the validation split")
train_logits = (
cache["trainLogits"][: len(train_records)].float().numpy().copy()
)
validation_logits = (
cache["validationLogits"][: len(validation_records)]
.float()
.numpy()
.copy()
)
except DataError:
raise
except (KeyError, TypeError, ValueError, RuntimeError) as exc:
raise DataError(f"{path}: cannot load distillation cache: {exc}") from exc
if train_logits.shape != (len(train_records), len(LABELS)):
raise DataError("distillation training logits have the wrong shape")
if validation_logits.shape != (len(validation_records), len(LABELS)):
raise DataError("distillation validation logits have the wrong shape")
return train_logits, validation_logits
def _checkpoint_config(model_dir: Path) -> dict[str, Any]:
config_path = model_dir / "config.json"
try:
config = json.loads(config_path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError) as exc:
raise DataError(f"{config_path}: cannot load model config: {exc}") from exc
if (
config.get("model_type") != "bert"
or config.get("hidden_size") != 384
or config.get("num_hidden_layers") != 6
or len(config.get("id2label", {})) != len(LABELS)
):
raise DataError("MLX purpose-lite requires the 6-layer 384-wide BERT classifier")
configured_labels = [
config["id2label"].get(str(index), config["id2label"].get(index))
for index in range(len(LABELS))
]
if configured_labels != list(LABELS):
raise DataError("MLX checkpoint label order does not match purpose-lite")
return config
def _save_checkpoint(
mx: Any,
model: Any,
source_dir: Path,
destination: Path,
config: dict[str, Any],
*,
quantization_aware: bool,
) -> None:
from mlx_model import save_hugging_face_weights
if destination.exists():
shutil.rmtree(destination)
destination.mkdir(parents=True)
for source in source_dir.iterdir():
if source.name.startswith("model") and source.suffix == ".safetensors":
continue
target = destination / source.name
if source.is_dir():
shutil.copytree(source, target)
else:
shutil.copy2(source, target)
output_config = dict(config)
output_config["purpose_classifier_training_backend"] = "mlx"
output_config["purpose_classifier_quantization_aware_training"] = bool(
quantization_aware
)
(destination / "config.json").write_text(
json.dumps(output_config, indent=2, sort_keys=True) + "\n",
encoding="utf-8",
)
save_hugging_face_weights(model, destination / "model.safetensors")
mx.eval(model.parameters())
def _softmax(values: np.ndarray) -> np.ndarray:
shifted = values - values.max(axis=-1, keepdims=True)
exponentials = np.exp(shifted)
return exponentials / exponentials.sum(axis=-1, keepdims=True)
def _linear_schedule(
mx: Any,
learning_rate: float,
total_steps: int,
warmup_steps: int,
) -> Any:
def schedule(step: Any) -> Any:
step = step.astype(mx.float32)
if warmup_steps:
warmup = learning_rate * step / warmup_steps
else:
warmup = mx.array(learning_rate)
remaining = max(total_steps - warmup_steps, 1)
decay = learning_rate * mx.maximum(
0.0,
(total_steps - step) / remaining,
)
if warmup_steps:
return mx.where(step < warmup_steps, warmup, decay)
return decay
return schedule
def train(args: argparse.Namespace) -> dict[str, Any]:
mx, nn, optim = _load_mlx(args.device)
try:
from transformers import AutoTokenizer
from mlx_model import (
BertClassifierConfig,
BertForSequenceClassification,
QATEmbedding,
QATLinear,
load_hugging_face_weights,
)
except ImportError as exc:
raise DataError(
"MLX training dependencies are missing; install requirements-mlx.txt"
) from exc
model_dir = args.model.expanduser()
checkpoint = model_dir / "model.safetensors"
if not checkpoint.is_file():
raise DataError("--model must be a local Hugging Face safetensors checkpoint")
output_dir: Path = args.output_dir
try:
model_dir.resolve().relative_to(output_dir.resolve())
except ValueError:
pass
else:
raise DataError("--model must not be inside --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)
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]
config_json = _checkpoint_config(model_dir)
config = BertClassifierConfig.from_hugging_face(config_json)
tokenizer = AutoTokenizer.from_pretrained(model_dir, local_files_only=True)
print("tokenizing train and validation splits", flush=True)
encoded_train = _encode_records(tokenizer, train_records)
encoded_validation = _encode_records(tokenizer, validation_records)
label_to_id = {label: index for index, label in enumerate(LABELS)}
train_labels = np.asarray(
[label_to_id[record["purpose"]] for record in train_records],
dtype=np.int32,
)
validation_labels = np.asarray(
[label_to_id[record["purpose"]] for record in validation_records],
dtype=np.int32,
)
sample_weights = np.asarray(
[
training_weight(record, args.boundary_weight)
for record in train_records
],
dtype=np.float32,
)
teacher_train_logits = None
teacher_validation_logits = None
if args.distillation_cache is not None:
teacher_train_logits, teacher_validation_logits = _teacher_cache(
args.distillation_cache.expanduser(),
train_records,
validation_records,
)
random.seed(args.seed)
np.random.seed(args.seed)
mx.random.seed(args.seed)
model = BertForSequenceClassification(
config,
quantization_aware=args.quantization_aware,
)
load_hugging_face_weights(model, checkpoint)
qat_modules = {
"linear": sum(isinstance(module, QATLinear) for module in model.modules()),
"embedding": sum(
isinstance(module, QATEmbedding) for module in model.modules()
),
}
validation_scorable = np.asarray(
[record.get("slice") != "vague-eval" for record in validation_records],
dtype=np.bool_,
)
initial_logits, _ = _evaluate(
mx,
model,
encoded_validation,
validation_labels,
args.eval_batch_size,
)
initial_predictions = initial_logits.argmax(axis=-1)
initial_metrics = classification_metrics(
validation_labels[validation_scorable].tolist(),
initial_predictions[validation_scorable].tolist(),
)
teacher_validation_predictions = (
teacher_validation_logits.argmax(axis=-1)
if teacher_validation_logits is not None
else None
)
def teacher_agreement(predictions: np.ndarray) -> float | None:
if teacher_validation_predictions is None:
return None
return float(
np.mean(
predictions[validation_scorable]
== teacher_validation_predictions[validation_scorable]
)
)
def selection_score(accuracy: float, agreement: float | None) -> float:
if agreement is None:
return accuracy
weight = args.distillation_selection_weight
return (accuracy + weight * agreement) / (1.0 + weight)
initial_agreement = teacher_agreement(initial_predictions)
initial_selection_score = selection_score(
initial_metrics["accuracy"],
initial_agreement,
)
if initial_agreement is not None:
initial_metrics["teacherAgreement"] = initial_agreement
initial_metrics["selectionScore"] = initial_selection_score
best_accuracy = initial_metrics["accuracy"]
best_selection_score = initial_selection_score
best_dir = output_dir / "model"
_save_checkpoint(
mx,
model,
model_dir,
best_dir,
config_json,
quantization_aware=args.quantization_aware,
)
print(
f"epoch 0: validation_accuracy={initial_metrics['accuracy']:.4%} "
f"macro_recall={initial_metrics['macroRecall']:.4%}"
+ (
f" teacher_agreement={initial_agreement:.4%} "
f"selection_score={initial_selection_score:.4%}"
if initial_agreement is not None
else ""
),
flush=True,
)
steps_per_epoch = math.ceil(len(train_records) / args.batch_size)
total_steps = steps_per_epoch * args.epochs
schedule = _linear_schedule(
mx,
args.learning_rate,
total_steps,
round(total_steps * args.warmup_ratio),
)
optimizer = optim.AdamW(
learning_rate=schedule,
weight_decay=args.weight_decay,
bias_correction=True,
)
def loss_function(
input_ids: Any,
attention_mask: Any,
token_type_ids: Any,
labels: Any,
weights: Any,
teacher_logits: Any | None,
) -> tuple[Any, Any, Any]:
logits = model(
input_ids=input_ids,
attention_mask=attention_mask,
token_type_ids=token_type_ids,
)
label_loss = nn.losses.cross_entropy(
logits,
labels,
reduction="none",
)
distillation_loss = mx.zeros_like(label_loss)
if teacher_logits is not None:
temperature = args.distillation_temperature
student_log_probabilities = (
logits / temperature
- mx.logsumexp(logits / temperature, axis=-1, keepdims=True)
)
teacher_probabilities = mx.softmax(
teacher_logits / temperature,
axis=-1,
)
teacher_log_probabilities = mx.log(
mx.maximum(teacher_probabilities, 1e-12)
)
distillation_loss = (
mx.sum(
teacher_probabilities
* (teacher_log_probabilities - student_log_probabilities),
axis=-1,
)
* temperature
* temperature
)
per_record_loss = (
(1.0 - args.distillation_weight) * label_loss
+ args.distillation_weight * distillation_loss
)
denominator = mx.sum(weights)
loss = mx.sum(per_record_loss * weights) / denominator
mean_label = mx.sum(label_loss * weights) / denominator
mean_distillation = mx.sum(distillation_loss * weights) / denominator
return loss, mean_label, mean_distillation
loss_and_grad = nn.value_and_grad(model, loss_function)
rng = np.random.default_rng(args.seed)
epochs_without_improvement = 0
stopped_early = False
history: list[dict[str, Any]] = []
started = time.perf_counter()
for epoch in range(1, args.epochs + 1):
epoch_started = time.perf_counter()
model.train()
running_loss = 0.0
running_label_loss = 0.0
running_distillation_loss = 0.0
permutation = rng.permutation(len(train_records))
for step, indexes in enumerate(
_batch_indexes(
len(train_records),
args.batch_size,
permutation=permutation,
),
1,
):
batch = _mlx_batch(mx, encoded_train, indexes)
labels = mx.array(train_labels[indexes])
weights = mx.array(sample_weights[indexes])
teacher_logits = (
mx.array(teacher_train_logits[indexes])
if teacher_train_logits is not None
else None
)
(loss, label_loss, distillation_loss), gradients = loss_and_grad(
batch["input_ids"],
batch["attention_mask"],
batch.get("token_type_ids"),
labels,
weights,
teacher_logits,
)
gradients, _ = optim.clip_grad_norm(gradients, args.max_grad_norm)
optimizer.update(model, gradients)
mx.eval(
model.parameters(),
optimizer.state,
loss,
label_loss,
distillation_loss,
)
running_loss += float(loss.item())
running_label_loss += float(label_loss.item())
running_distillation_loss += float(distillation_loss.item())
if args.progress_steps and (
step % args.progress_steps == 0 or step == steps_per_epoch
):
print(
f"epoch {epoch} step {step}/{steps_per_epoch} "
f"mean_loss={running_loss / step:.4f} "
f"label_loss={running_label_loss / step:.4f} "
f"distill_loss={running_distillation_loss / step:.4f} "
f"elapsed={time.perf_counter() - epoch_started:.1f}s",
flush=True,
)
logits, _ = _evaluate(
mx,
model,
encoded_validation,
validation_labels,
args.eval_batch_size,
)
predictions = logits.argmax(axis=-1)
metrics = classification_metrics(
validation_labels[validation_scorable].tolist(),
predictions[validation_scorable].tolist(),
)
agreement = teacher_agreement(predictions)
candidate_selection_score = selection_score(metrics["accuracy"], agreement)
if agreement is not None:
metrics["teacherAgreement"] = agreement
metrics["selectionScore"] = candidate_selection_score
metrics["epoch"] = epoch
metrics["meanTrainingLoss"] = running_loss / steps_per_epoch
metrics["meanLabelLoss"] = running_label_loss / steps_per_epoch
metrics["meanDistillationLoss"] = (
running_distillation_loss / steps_per_epoch
)
history.append(metrics)
print(
f"epoch {epoch}: loss={metrics['meanTrainingLoss']:.4f} "
f"validation_accuracy={metrics['accuracy']:.4%} "
f"macro_recall={metrics['macroRecall']:.4%}"
+ (
f" teacher_agreement={agreement:.4%} "
f"selection_score={candidate_selection_score:.4%}"
if agreement is not None
else ""
),
flush=True,
)
improvement = candidate_selection_score - best_selection_score
if improvement > args.minimum_improvement:
best_accuracy = metrics["accuracy"]
best_selection_score = candidate_selection_score
epochs_without_improvement = 0
_save_checkpoint(
mx,
model,
model_dir,
best_dir,
config_json,
quantization_aware=args.quantization_aware,
)
else:
epochs_without_improvement += 1
if epochs_without_improvement >= args.early_stopping_patience:
stopped_early = True
print(
f"early stopping after epoch {epoch}: no selection-score "
f"improvement greater than {args.minimum_improvement:.4%} "
f"for {args.early_stopping_patience} epoch(s)",
flush=True,
)
break
final_model = BertForSequenceClassification(
config,
quantization_aware=False,
)
load_hugging_face_weights(final_model, best_dir / "model.safetensors")
logits, labels = _evaluate(
mx,
final_model,
encoded_validation,
validation_labels,
args.eval_batch_size,
)
try:
import torch
except ImportError as exc:
raise DataError("final calibration requires PyTorch") from exc
temperature = _fit_temperature(
torch,
torch.from_numpy(logits),
torch.from_numpy(labels.astype(np.int64)),
)
calibrated = _softmax(logits / temperature)
sorted_indexes = np.argsort(calibrated, axis=-1)
top_indexes = sorted_indexes[:, -1]
second_indexes = sorted_indexes[:, -2]
row_indexes = np.arange(len(labels))
top_probabilities = calibrated[row_indexes, top_indexes]
margins = (
top_probabilities - calibrated[row_indexes, second_indexes]
)
correct = (
(top_indexes == labels) & validation_scorable
).tolist()
thresholds = choose_confidence_thresholds(
top_probabilities.tolist(),
margins.tolist(),
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.tolist(),
correct,
),
}
metrics = {
"modelVersion": "purpose-lite-v1",
"baseModel": str(model_dir),
"baseModelRevision": "local-checkpoint",
"trainingBackend": "mlx",
"device": args.device,
"fixedInputShape": [1, MAX_LENGTH],
"truncation": {
"strategy": "head-tail-pair",
"headTokens": HEAD_TOKENS,
"tailTokens": TAIL_TOKENS,
},
"trainingSeconds": time.perf_counter() - started,
"trainRecords": len(train_records),
"validationRecords": len(validation_records),
"scoredValidationRecords": int(validation_scorable.sum()),
"vagueAbstentionValidationRecords": int(
(~validation_scorable).sum()
),
"boundaryTrainingWeight": args.boundary_weight,
"quantizationAwareTraining": args.quantization_aware,
"quantizationAwareModules": qat_modules,
"distillation": {
"cache": (
str(args.distillation_cache)
if args.distillation_cache is not None
else None
),
"weight": args.distillation_weight,
"temperature": args.distillation_temperature,
"selectionAgreementWeight": args.distillation_selection_weight,
},
"bestValidationAccuracy": best_accuracy,
"bestValidationSelectionScore": best_selection_score,
"initialValidation": initial_metrics,
"epochsCompleted": len(history),
"stoppedEarly": stopped_early,
"bestValidation": classification_metrics(
labels[validation_scorable].tolist(),
top_indexes[validation_scorable].tolist(),
),
"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", type=Path, required=True)
parser.add_argument(
"--device",
choices=("metal", "cpu"),
default="metal",
help="MLX execution device (default: metal; cpu is a diagnostic fallback)",
)
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("--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("--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("--distillation-cache", type=Path)
parser.add_argument("--distillation-weight", type=float, default=0.0)
parser.add_argument("--distillation-temperature", type=float, default=2.0)
parser.add_argument("--distillation-selection-weight", type=float, default=0.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 main(argv: Sequence[str] | None = None) -> int:
parser = build_parser()
args = parser.parse_args(argv)
for name in (
"epochs",
"batch_size",
"eval_batch_size",
"early_stopping_patience",
):
if getattr(args, name) <= 0:
parser.error(f"--{name.replace('_', '-')} must be positive")
if args.progress_steps < 0:
parser.error("--progress-steps must be non-negative")
if args.learning_rate <= 0:
parser.error("--learning-rate must be positive")
if not 0 <= args.warmup_ratio < 1:
parser.error("--warmup-ratio must be in [0, 1)")
if args.boundary_weight <= 0:
parser.error("--boundary-weight must be positive")
if not 0 <= args.distillation_weight <= 1:
parser.error("--distillation-weight must be in [0, 1]")
if args.distillation_temperature <= 0:
parser.error("--distillation-temperature must be positive")
if not 0 <= args.distillation_selection_weight <= 1:
parser.error("--distillation-selection-weight must be in [0, 1]")
if (args.distillation_cache is None) != (args.distillation_weight == 0):
parser.error(
"--distillation-cache and a positive --distillation-weight "
"must be supplied together"
)
if args.distillation_selection_weight and args.distillation_cache is None:
parser.error(
"--distillation-selection-weight requires --distillation-cache"
)
try:
metrics = train(args)
except DataError as exc:
print(f"error: {exc}", file=sys.stderr)
return 2
print(
f"selected validation accuracy: {metrics['bestValidationAccuracy']:.4%}",
flush=True,
)
return 0
if __name__ == "__main__":
raise SystemExit(main())