2026-07-30 03:48:00 -07:00
|
|
|
#!/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
|
2026-07-30 05:04:34 -07:00
|
|
|
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
|
2026-07-30 03:48:00 -07:00
|
|
|
|
|
|
|
|
|
|
|
|
|
def prepare_text(prompt: str) -> str:
|
|
|
|
|
return normalize_prompt(prompt)
|
|
|
|
|
|
|
|
|
|
|
2026-07-30 17:28:44 -07:00
|
|
|
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
|
|
|
|
|
|
|
|
|
|
|
2026-07-30 18:39:33 -07:00
|
|
|
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
|
|
|
|
|
|
|
|
|
|
|
2026-07-30 05:04:34 -07:00
|
|
|
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
|
|
|
|
|
|
|
|
|
|
|
2026-07-30 03:48:00 -07:00
|
|
|
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")
|
2026-07-30 17:28:44 -07:00
|
|
|
batch.pop("sample_weights", None)
|
2026-07-30 03:48:00 -07:00
|
|
|
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
|
2026-07-30 17:28:44 -07:00
|
|
|
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"
|
|
|
|
|
)
|
2026-07-30 03:48:00 -07:00
|
|
|
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()}
|
2026-07-30 17:28:44 -07:00
|
|
|
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}
|
|
|
|
|
)
|
2026-07-30 03:48:00 -07:00
|
|
|
tokenizer = AutoTokenizer.from_pretrained(
|
2026-07-30 17:28:44 -07:00
|
|
|
args.model, use_fast=True, **pretrained_options
|
2026-07-30 03:48:00 -07:00
|
|
|
)
|
|
|
|
|
model = AutoModelForSequenceClassification.from_pretrained(
|
|
|
|
|
args.model,
|
|
|
|
|
num_labels=len(LABELS),
|
|
|
|
|
label2id=label_to_id,
|
|
|
|
|
id2label=id_to_label,
|
|
|
|
|
ignore_mismatched_sizes=True,
|
2026-07-30 17:28:44 -07:00
|
|
|
**pretrained_options,
|
2026-07-30 03:48:00 -07:00
|
|
|
)
|
|
|
|
|
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]
|
2026-07-30 05:04:34 -07:00
|
|
|
config.purpose_classifier_truncation = {
|
|
|
|
|
"strategy": "head-tail-pair",
|
|
|
|
|
"headTokens": HEAD_TOKENS,
|
|
|
|
|
"tailTokens": TAIL_TOKENS,
|
|
|
|
|
}
|
2026-07-30 18:39:33 -07:00
|
|
|
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)
|
2026-07-30 03:48:00 -07:00
|
|
|
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)
|
|
|
|
|
|
2026-07-30 17:28:44 -07:00
|
|
|
def __getitem__(self, index: int) -> tuple[str, int, float]:
|
2026-07-30 03:48:00 -07:00
|
|
|
record = self.records[index]
|
2026-07-30 17:28:44 -07:00
|
|
|
return (
|
|
|
|
|
prepare_text(record["prompt"]),
|
|
|
|
|
label_to_id[record["purpose"]],
|
|
|
|
|
training_weight(record, args.boundary_weight),
|
|
|
|
|
)
|
2026-07-30 03:48:00 -07:00
|
|
|
|
2026-07-30 17:28:44 -07:00
|
|
|
def collate(items: Sequence[tuple[str, int, float]]) -> dict[str, Any]:
|
|
|
|
|
texts, labels, weights = zip(*items)
|
2026-07-30 05:04:34 -07:00
|
|
|
encoded = encode_fixed_shape(tokenizer, list(texts), torch)
|
2026-07-30 03:48:00 -07:00
|
|
|
encoded["labels"] = torch.tensor(labels, dtype=torch.long)
|
2026-07-30 17:28:44 -07:00
|
|
|
encoded["sample_weights"] = torch.tensor(weights, dtype=torch.float)
|
2026-07-30 03:48:00 -07:00
|
|
|
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",
|
|
|
|
|
)
|
2026-07-30 05:04:34 -07:00
|
|
|
validation_scorable = torch.tensor(
|
|
|
|
|
[record.get("slice") != "vague-eval" for record in validation_records],
|
|
|
|
|
dtype=torch.bool,
|
|
|
|
|
)
|
2026-07-30 17:28:44 -07:00
|
|
|
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,
|
|
|
|
|
)
|
2026-07-30 03:48:00 -07:00
|
|
|
|
|
|
|
|
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):
|
2026-07-30 18:39:33 -07:00
|
|
|
epoch_started = time.perf_counter()
|
2026-07-30 03:48:00 -07:00
|
|
|
model.train()
|
|
|
|
|
optimizer.zero_grad(set_to_none=True)
|
|
|
|
|
running_loss = 0.0
|
|
|
|
|
for step, batch in enumerate(train_loader, 1):
|
2026-07-30 17:28:44 -07:00
|
|
|
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
|
2026-07-30 03:48:00 -07:00
|
|
|
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)
|
2026-07-30 18:39:33 -07:00
|
|
|
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,
|
|
|
|
|
)
|
2026-07-30 03:48:00 -07:00
|
|
|
|
|
|
|
|
logits, labels = _evaluate(torch, model, validation_loader, device)
|
|
|
|
|
predictions = logits.argmax(dim=-1).tolist()
|
2026-07-30 05:04:34 -07:00
|
|
|
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)
|
2026-07-30 03:48:00 -07:00
|
|
|
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,
|
|
|
|
|
)
|
2026-07-30 17:28:44 -07:00
|
|
|
improvement = metrics["accuracy"] - best_accuracy
|
|
|
|
|
if improvement > args.minimum_improvement:
|
2026-07-30 03:48:00 -07:00
|
|
|
best_accuracy = metrics["accuracy"]
|
2026-07-30 17:28:44 -07:00
|
|
|
epochs_without_improvement = 0
|
2026-07-30 03:48:00 -07:00
|
|
|
model.save_pretrained(best_dir, safe_serialization=True)
|
|
|
|
|
tokenizer.save_pretrained(best_dir)
|
2026-07-30 17:28:44 -07:00
|
|
|
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
|
2026-07-30 03:48:00 -07:00
|
|
|
|
|
|
|
|
model = AutoModelForSequenceClassification.from_pretrained(best_dir).to(device)
|
|
|
|
|
logits, labels = _evaluate(torch, model, validation_loader, device)
|
2026-07-30 05:04:34 -07:00
|
|
|
temperature = _fit_temperature(
|
|
|
|
|
torch,
|
|
|
|
|
logits[validation_scorable],
|
|
|
|
|
labels[validation_scorable],
|
|
|
|
|
)
|
2026-07-30 03:48:00 -07:00
|
|
|
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 = [
|
2026-07-30 05:04:34 -07:00
|
|
|
prediction == actual and scorable
|
|
|
|
|
for prediction, actual, scorable in zip(
|
|
|
|
|
predictions,
|
|
|
|
|
labels.tolist(),
|
|
|
|
|
validation_scorable.tolist(),
|
|
|
|
|
)
|
2026-07-30 03:48:00 -07:00
|
|
|
]
|
|
|
|
|
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,
|
2026-07-30 17:28:44 -07:00
|
|
|
"baseModelRevision": model_revision or "local-checkpoint",
|
2026-07-30 03:48:00 -07:00
|
|
|
"fixedInputShape": [1, MAX_LENGTH],
|
2026-07-30 05:04:34 -07:00
|
|
|
"truncation": {
|
|
|
|
|
"strategy": "head-tail-pair",
|
|
|
|
|
"headTokens": HEAD_TOKENS,
|
|
|
|
|
"tailTokens": TAIL_TOKENS,
|
|
|
|
|
},
|
2026-07-30 03:48:00 -07:00
|
|
|
"device": str(device),
|
|
|
|
|
"trainingSeconds": time.perf_counter() - started,
|
|
|
|
|
"trainRecords": len(train_records),
|
2026-07-30 17:28:44 -07:00
|
|
|
"boundaryTrainingWeight": args.boundary_weight,
|
2026-07-30 18:39:33 -07:00
|
|
|
"quantizationAwareTraining": args.quantization_aware,
|
|
|
|
|
"quantizationAwareModules": qat_modules,
|
2026-07-30 03:48:00 -07:00
|
|
|
"validationRecords": len(validation_records),
|
2026-07-30 05:04:34 -07:00
|
|
|
"scoredValidationRecords": int(validation_scorable.sum().item()),
|
|
|
|
|
"vagueAbstentionValidationRecords": int(
|
|
|
|
|
(~validation_scorable).sum().item()
|
|
|
|
|
),
|
2026-07-30 03:48:00 -07:00
|
|
|
"bestValidationAccuracy": best_accuracy,
|
2026-07-30 17:28:44 -07:00
|
|
|
"initialValidation": initial_metrics,
|
|
|
|
|
"epochsCompleted": len(history),
|
|
|
|
|
"stoppedEarly": stopped_early,
|
2026-07-30 05:04:34 -07:00
|
|
|
"bestValidation": classification_metrics(
|
|
|
|
|
labels[validation_scorable].tolist(),
|
|
|
|
|
[
|
|
|
|
|
prediction
|
|
|
|
|
for prediction, scorable in zip(
|
|
|
|
|
predictions, validation_scorable.tolist()
|
|
|
|
|
)
|
|
|
|
|
if scorable
|
|
|
|
|
],
|
|
|
|
|
),
|
2026-07-30 03:48:00 -07:00
|
|
|
"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)
|
2026-07-30 18:39:33 -07:00
|
|
|
parser.add_argument("--progress-steps", type=int, default=50)
|
2026-07-30 17:28:44 -07:00
|
|
|
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)
|
2026-07-30 18:39:33 -07:00
|
|
|
parser.add_argument("--quantization-aware", action="store_true")
|
2026-07-30 03:48:00 -07:00
|
|
|
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",
|
2026-07-30 17:28:44 -07:00
|
|
|
"early_stopping_patience",
|
2026-07-30 03:48:00 -07:00
|
|
|
):
|
|
|
|
|
_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)")
|
2026-07-30 17:28:44 -07:00
|
|
|
if args.minimum_improvement < 0.0:
|
|
|
|
|
parser.error("--minimum-improvement must be non-negative")
|
2026-07-30 18:39:33 -07:00
|
|
|
if args.progress_steps < 0:
|
|
|
|
|
parser.error("--progress-steps must be non-negative")
|
2026-07-30 17:28:44 -07:00
|
|
|
if args.boundary_weight <= 0.0:
|
|
|
|
|
parser.error("--boundary-weight must be positive")
|
2026-07-30 03:48:00 -07:00
|
|
|
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())
|