Merge nucleic/sleek-ember-seal-uady into dev

This commit is contained in:
2026-07-30 19:21:08 -07:00
parent bb6d53a520
commit 09f98cdd00
7 changed files with 511 additions and 25 deletions
+33 -3
View File
@@ -117,6 +117,34 @@ it is rejected. Do not continue optimizer-only QAT sweeps on this split. The nex
iteration should incorporate reviewed boundary data and be selected on a revised
validation/frozen dataset version.
To target only the remaining float→int8 decision drift, cache the float teacher in a
separate inference process and use its logits for QAT distillation. Keeping teacher and
student models out of the same process avoids doubling peak resident memory:
```bash
ml/purpose-classifier/.venv/bin/python ml/purpose-classifier/cache_teacher.py \
--model ml/purpose-classifier/outputs/purpose-lite-v1-boundary-tune/model \
--output ml/purpose-classifier/outputs/purpose-lite-v1-boundary-teacher.pt \
--overwrite-output
ml/purpose-classifier/.venv/bin/python ml/purpose-classifier/train.py \
--model ml/purpose-classifier/outputs/purpose-lite-v1-boundary-tune/model \
--distillation-cache \
ml/purpose-classifier/outputs/purpose-lite-v1-boundary-teacher.pt \
--distillation-weight 0.9 --distillation-temperature 2 \
--distillation-selection-weight 0.5 --quantization-aware \
--epochs 2 --learning-rate 1e-6 --warmup-ratio 0 --boundary-weight 1 \
--output-dir ml/purpose-classifier/outputs/purpose-lite-v1-distilled-qat \
--overwrite-output
```
The cache binds each logit row to normalized prompt hash plus expected label. Training
fails closed if either split changes. Selection combines label accuracy with float-teacher
agreement, retains the incoming checkpoint as epoch zero, and logs label/distillation loss
separately. A 64-record wiring run exercised cache loading, shuffled row alignment,
backpropagation, selection, and ordinary checkpoint reload. The current shared CPU runtime
then showed severe post-batch throttling, so no full candidate result is claimed from that
canary.
For a wiring smoke test, use a small deterministic prefix:
```bash
@@ -178,10 +206,12 @@ the frozen split automatically. The 18 word-trigram exclusions and the human-rev
completion rule are recorded in `data/curation-review-v1.json`; the semantic report is
versioned as `data/semantic-audit-v1.json`.
## Complete the human review
## Optional human review
The deterministic CSV currently contains 1,219 blank review rows. Check progress without
running the embedding audit again:
The dataset owner accepted the curated generated labels and difficulty metadata as-is on
2026-07-31, so the blank 1,219-row review sample is not a training or rollout blocker. It
remains available as an optional future audit. Check its progress without running the
embedding audit again:
```bash
ml/purpose-classifier/.venv/bin/python ml/purpose-classifier/review_data.py
+186
View File
@@ -0,0 +1,186 @@
#!/usr/bin/env python3
"""Cache float-teacher logits for memory-bounded QAT distillation."""
from __future__ import annotations
import argparse
import sys
import time
from pathlib import Path
from typing import Any, Sequence
from purpose_data import LABELS, DataError, load_jsonl
from train import (
DEFAULT_DATASET_DIR,
DEFAULT_MODEL_REVISION,
_select_device,
_set_seeds,
_validate_split,
distillation_record_keys,
encode_fixed_shape,
prepare_text,
)
SCRIPT_DIR = Path(__file__).resolve().parent
DEFAULT_OUTPUT = SCRIPT_DIR / "outputs" / "purpose-lite-v1-teacher-logits.pt"
def _predict(
torch: Any,
model: Any,
tokenizer: Any,
records: Sequence[dict[str, Any]],
*,
device: Any,
batch_size: int,
progress_steps: int,
label: str,
) -> Any:
rows = []
batches = (len(records) + batch_size - 1) // batch_size
started = time.perf_counter()
model.eval()
with torch.inference_mode():
for batch_index, start in enumerate(range(0, len(records), batch_size), 1):
batch = records[start : start + batch_size]
encoded = encode_fixed_shape(
tokenizer,
[prepare_text(record["prompt"]) for record in batch],
torch,
)
inputs = {key: value.to(device) for key, value in encoded.items()}
rows.append(model(**inputs).logits.cpu())
if progress_steps and (
batch_index % progress_steps == 0 or batch_index == batches
):
print(
f"teacher {label} step {batch_index}/{batches} "
f"elapsed={time.perf_counter() - started:.1f}s",
flush=True,
)
return torch.cat(rows)
def cache_teacher(args: argparse.Namespace) -> dict[str, Any]:
try:
import torch
from transformers import AutoModelForSequenceClassification, AutoTokenizer
except ImportError as exc:
raise DataError(
"teacher-cache dependencies are missing; install requirements.txt"
) 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)
output = args.output.resolve()
local_model = Path(args.model).expanduser()
if local_model.is_dir():
try:
output.relative_to(local_model.resolve())
except ValueError:
pass
else:
raise DataError("teacher cache output must not overwrite the model directory")
if output.exists() and not args.overwrite_output:
raise DataError(f"{output}: cache exists; pass --overwrite-output intentionally")
if output.exists() and output.is_dir():
raise DataError(f"{output}: cache output must be a file path")
output.parent.mkdir(parents=True, exist_ok=True)
_set_seeds(torch, args.seed)
device = _select_device(torch, args.device)
options = (
{"local_files_only": True}
if local_model.exists()
else {"revision": args.model_revision}
)
tokenizer = AutoTokenizer.from_pretrained(args.model, use_fast=True, **options)
model = AutoModelForSequenceClassification.from_pretrained(
args.model,
**options,
).to(device)
teacher_labels = [
model.config.id2label.get(index, model.config.id2label.get(str(index)))
for index in range(len(LABELS))
]
if teacher_labels != list(LABELS):
raise DataError("teacher label order does not match purpose-lite")
train_logits = _predict(
torch,
model,
tokenizer,
train_records,
device=device,
batch_size=args.batch_size,
progress_steps=args.progress_steps,
label="train",
)
validation_logits = _predict(
torch,
model,
tokenizer,
validation_records,
device=device,
batch_size=args.batch_size,
progress_steps=args.progress_steps,
label="validation",
)
artifact = {
"schemaVersion": 1,
"labels": list(LABELS),
"teacher": str(args.model),
"trainRecordKeys": distillation_record_keys(train_records),
"validationRecordKeys": distillation_record_keys(validation_records),
"trainLogits": train_logits,
"validationLogits": validation_logits,
}
torch.save(artifact, output)
return {
"trainRecords": len(train_records),
"validationRecords": len(validation_records),
"output": str(output),
}
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--dataset-dir", type=Path, default=DEFAULT_DATASET_DIR)
parser.add_argument("--model", required=True)
parser.add_argument("--model-revision", default=DEFAULT_MODEL_REVISION)
parser.add_argument("--output", type=Path, default=DEFAULT_OUTPUT)
parser.add_argument("--device", default="auto")
parser.add_argument("--seed", type=int, default=20260730)
parser.add_argument("--batch-size", type=int, default=64)
parser.add_argument("--progress-steps", type=int, default=50)
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)
if args.batch_size <= 0:
parser.error("--batch-size must be positive")
if args.progress_steps < 0:
parser.error("--progress-steps must be non-negative")
try:
metrics = cache_teacher(args)
except (DataError, OSError, ValueError, RuntimeError) as exc:
print(f"error: {exc}", file=sys.stderr)
return 1
print(
f"Cached teacher logits for {metrics['trainRecords']} train and "
f"{metrics['validationRecords']} validation records at {metrics['output']}."
)
return 0
if __name__ == "__main__":
raise SystemExit(main())
+4 -2
View File
@@ -61,7 +61,7 @@
]
},
"humanLabelAndDifficultyReview": {
"status": "planned",
"status": "accepted-as-generated",
"populationRecords": 12193,
"sampleFraction": 0.1,
"sampleRecords": 1219,
@@ -80,6 +80,8 @@
"notes"
],
"generatedArtifact": "ml/purpose-classifier/.artifacts/human-review-v1.csv",
"completionRule": "Every sampled row must be marked accept, relabel, or reject by a human reviewer. review_data.py validates the exact sample and writes a versioned complete ledger; prepare_data.py applies relabel/reject decisions before splitting. The revised split and frozen evaluation must then be intentionally reviewed and versioned before the dataset can be called fully curated."
"decisionDate": "2026-07-31",
"decisionBasis": "The dataset owner explicitly directed the project to assume the generated labels, secondary purposes, difficulties, and slices are correct and validated without completing the row-by-row sample.",
"decision": "Accept the curated generated population as-is. No relabel or reject decisions are inferred, and the blank CSV remains an optional future audit artifact rather than a rollout blocker."
}
}
+1 -1
View File
@@ -9,7 +9,7 @@
"nearDuplicateThreshold": 0.92,
"retainedRecords": 12193,
"reviewPath": "ml/purpose-classifier/data/curation-review-v1.json",
"reviewSha256": "625251a98bcdcde0bee074e3bba3c564af4347eb7955497a9c3e018aae00b30f",
"reviewSha256": "1721ab77a0f723a4b03f42b14850b5c351ead2cd8fa72bb2d1fe3f0d2bb568da",
"vagueEvalPolicy": "validation/test only"
},
"datasetVersion": "purpose-dataset-v1",
+14 -1
View File
@@ -65,6 +65,17 @@ def reviewed_population(args: argparse.Namespace) -> list[SourceRecord]:
).records
def review_policy_status(path: Path) -> str:
try:
value = json.loads(path.read_text(encoding="utf-8"))
status = value["humanLabelAndDifficultyReview"]["status"]
except (OSError, UnicodeError, json.JSONDecodeError, KeyError, TypeError) as exc:
raise DataError(f"{path}: cannot read human-review policy: {exc}") from exc
if not isinstance(status, str) or not status:
raise DataError(f"{path}: invalid human-review policy status")
return status
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--source", action="append", type=Path)
@@ -98,6 +109,7 @@ def main(argv: Sequence[str] | None = None) -> int:
args = build_parser().parse_args(argv)
try:
population = reviewed_population(args)
policy_status = review_policy_status(args.curation_review.resolve())
sample = stratified_review_sample(
population,
fraction=args.review_fraction,
@@ -152,7 +164,8 @@ def main(argv: Sequence[str] | None = None) -> int:
print(
f"Human review: {progress.completed}/{progress.records} complete "
f"({progress.accepted} accept, {progress.relabeled} relabel, "
f"{progress.rejected} reject, {progress.incomplete} remaining)."
f"{progress.rejected} reject, {progress.incomplete} remaining). "
f"Dataset policy: {policy_status}."
)
return 0
+13
View File
@@ -12,6 +12,19 @@ import train
class MetricsTests(unittest.TestCase):
def test_distillation_loss_is_zero_for_matching_logits_and_backpropagates(self):
teacher = torch.tensor([[2.0, 0.0, -1.0]])
student = teacher.clone().requires_grad_(True)
loss = train.knowledge_distillation_loss(
torch,
student,
teacher,
temperature=2.0,
).mean()
self.assertAlmostEqual(0.0, loss.item(), places=6)
loss.backward()
self.assertIsNotNone(student.grad)
def test_qat_replacements_keep_checkpoint_keys_and_gradients(self):
model = torch.nn.Sequential(
torch.nn.Embedding(16, 8),
+260 -18
View File
@@ -12,7 +12,14 @@ import time
from pathlib import Path
from typing import Any, Sequence
from purpose_data import LABELS, DataError, load_jsonl, normalize_prompt, write_json
from purpose_data import (
LABELS,
DataError,
load_jsonl,
normalize_prompt,
prompt_hash,
write_json,
)
SCRIPT_DIR = Path(__file__).resolve().parent
@@ -37,6 +44,42 @@ def training_weight(record: dict[str, Any], boundary_weight: float) -> float:
return boundary_weight if record.get("slice") == "boundary" else 1.0
def knowledge_distillation_loss(
torch: Any,
student_logits: Any,
teacher_logits: Any,
*,
temperature: float,
) -> Any:
"""Return per-record KL loss from a frozen float teacher to the student."""
student_log_probabilities = torch.nn.functional.log_softmax(
student_logits / temperature,
dim=-1,
)
teacher_probabilities = torch.nn.functional.softmax(
teacher_logits / temperature,
dim=-1,
)
return (
torch.nn.functional.kl_div(
student_log_probabilities,
teacher_probabilities,
reduction="none",
).sum(dim=-1)
* temperature
* temperature
)
def distillation_record_keys(records: Sequence[dict[str, Any]]) -> list[str]:
"""Bind cached teacher logits to both normalized prompt and expected label."""
return [
f"{prompt_hash(record['prompt'])}:{record['purpose']}" for record in records
]
def enable_quantization_aware_training(torch: Any, model: Any) -> dict[str, int]:
"""Mirror the export graph's int8 policy with straight-through fake quantization.
@@ -400,18 +443,36 @@ def _fit_temperature(torch: Any, logits: Any, labels: Any) -> float:
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]:
def _evaluate(
torch: Any,
model: Any,
loader: Any,
device: Any,
*,
progress_label: str | None = None,
progress_steps: int = 0,
) -> tuple[Any, Any]:
model.eval()
all_logits = []
all_labels = []
started = time.perf_counter()
with torch.inference_mode():
for batch in loader:
for step, batch in enumerate(loader, 1):
labels = batch.pop("labels")
batch.pop("sample_weights", None)
batch.pop("teacher_logits", 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)
if progress_label and progress_steps and (
step % progress_steps == 0 or step == len(loader)
):
print(
f"{progress_label} step {step}/{len(loader)} "
f"elapsed={time.perf_counter() - started:.1f}s",
flush=True,
)
return torch.cat(all_logits), torch.cat(all_labels)
@@ -442,15 +503,25 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
output_dir: Path = args.output_dir
local_model = Path(args.model).expanduser()
if local_model.exists():
distillation_cache = (
args.distillation_cache.expanduser()
if args.distillation_cache is not None
else None
)
for option_name, local_path in (
("--model", local_model if local_model.exists() else None),
("--distillation-cache", distillation_cache),
):
if local_path is None:
continue
try:
local_model.resolve().relative_to(output_dir.resolve())
local_path.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"
f"local {option_name} 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:
@@ -505,31 +576,88 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
model.to(device)
class PromptDataset(Dataset):
def __init__(self, records: Sequence[dict[str, Any]]) -> None:
def __init__(
self,
records: Sequence[dict[str, Any]],
teacher_logits: Any | None = None,
) -> None:
self.records = records
self.teacher_logits = teacher_logits
def __len__(self) -> int:
return len(self.records)
def __getitem__(self, index: int) -> tuple[str, int, float]:
def __getitem__(self, index: int) -> tuple[str, int, float, Any | None]:
record = self.records[index]
return (
prepare_text(record["prompt"]),
label_to_id[record["purpose"]],
training_weight(record, args.boundary_weight),
(
self.teacher_logits[index]
if self.teacher_logits is not None
else None
),
)
def collate(items: Sequence[tuple[str, int, float]]) -> dict[str, Any]:
texts, labels, weights = zip(*items)
def collate(
items: Sequence[tuple[str, int, float, Any | None]],
) -> dict[str, Any]:
texts, labels, weights, teacher_rows = 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)
if teacher_rows[0] is not None:
encoded["teacher_logits"] = torch.stack(teacher_rows)
return encoded
teacher_train_logits = None
teacher_validation_logits = None
if distillation_cache is not None:
if not distillation_cache.is_file():
raise DataError(f"{distillation_cache}: distillation cache is missing")
try:
cache = torch.load(
distillation_cache,
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")
teacher_train_logits = cache["trainLogits"][
: len(train_records)
].float().clone()
teacher_validation_logits = cache["validationLogits"][
: len(validation_records)
].float().clone()
except DataError:
raise
except (KeyError, TypeError, ValueError, RuntimeError) as exc:
raise DataError(
f"{distillation_cache}: cannot load distillation cache: {exc}"
) from exc
expected_shape = (len(train_records), len(LABELS))
if tuple(teacher_train_logits.shape) != expected_shape:
raise DataError("distillation training logits have the wrong shape")
if tuple(teacher_validation_logits.shape) != (
len(validation_records),
len(LABELS),
):
raise DataError("distillation validation logits have the wrong shape")
del cache
generator = torch.Generator()
generator.manual_seed(args.seed)
train_loader = DataLoader(
PromptDataset(train_records),
PromptDataset(train_records, teacher_train_logits),
batch_size=args.batch_size,
shuffle=True,
generator=generator,
@@ -567,7 +695,42 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
if scorable
],
)
teacher_validation_predictions = (
teacher_validation_logits.argmax(dim=-1).tolist()
if teacher_validation_logits is not None
else None
)
def teacher_agreement(predictions: Sequence[int]) -> float | None:
if teacher_validation_predictions is None:
return None
agreements = [
prediction == teacher_prediction
for prediction, teacher_prediction, scorable in zip(
predictions,
teacher_validation_predictions,
validation_scorable.tolist(),
)
if scorable
]
return sum(agreements) / len(agreements)
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
epochs_without_improvement = 0
stopped_early = False
history = []
@@ -576,7 +739,13 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
tokenizer.save_pretrained(best_dir)
print(
f"epoch 0: validation_accuracy={initial_metrics['accuracy']:.4%} "
f"macro_recall={initial_metrics['macroRecall']:.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,
)
@@ -599,20 +768,46 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
model.train()
optimizer.zero_grad(set_to_none=True)
running_loss = 0.0
running_label_loss = 0.0
running_distillation_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)
teacher_logits = batch.pop("teacher_logits", None)
if teacher_logits is not None:
teacher_logits = teacher_logits.to(device)
inputs = {key: value.to(device) for key, value in batch.items()}
per_record_loss = torch.nn.functional.cross_entropy(
model(**inputs).logits,
student_logits = model(**inputs).logits
label_loss = torch.nn.functional.cross_entropy(
student_logits,
labels,
reduction="none",
)
distillation_loss = torch.zeros_like(label_loss)
if teacher_logits is not None:
distillation_loss = knowledge_distillation_loss(
torch,
student_logits,
teacher_logits,
temperature=args.distillation_temperature,
)
per_record_loss = (
(1.0 - args.distillation_weight) * label_loss
+ args.distillation_weight * distillation_loss
)
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
running_label_loss += float(
(label_loss * sample_weights).sum().item()
/ sample_weights.sum().item()
)
running_distillation_loss += float(
(distillation_loss * sample_weights).sum().item()
/ sample_weights.sum().item()
)
should_update = (
step % args.gradient_accumulation_steps == 0
or step == len(train_loader)
@@ -628,6 +823,8 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
print(
f"epoch {epoch} step {step}/{len(train_loader)} "
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,
)
@@ -643,18 +840,34 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
if scorable
]
metrics = classification_metrics(scored_labels, scored_predictions)
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 / len(train_loader)
metrics["meanLabelLoss"] = running_label_loss / len(train_loader)
metrics["meanDistillationLoss"] = (
running_distillation_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%}",
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 = metrics["accuracy"] - best_accuracy
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
model.save_pretrained(best_dir, safe_serialization=True)
tokenizer.save_pretrained(best_dir)
@@ -663,7 +876,7 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
if epochs_without_improvement >= args.early_stopping_patience:
stopped_early = True
print(
f"early stopping after epoch {epoch}: no validation improvement "
f"early stopping after epoch {epoch}: no selection-score improvement "
f"greater than {args.minimum_improvement:.4%} for "
f"{args.early_stopping_patience} epoch(s)",
flush=True,
@@ -721,12 +934,19 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
"boundaryTrainingWeight": args.boundary_weight,
"quantizationAwareTraining": args.quantization_aware,
"quantizationAwareModules": qat_modules,
"distillation": {
"cache": str(distillation_cache) if distillation_cache is not None else None,
"weight": args.distillation_weight,
"temperature": args.distillation_temperature,
"selectionAgreementWeight": args.distillation_selection_weight,
},
"validationRecords": len(validation_records),
"scoredValidationRecords": int(validation_scorable.sum().item()),
"vagueAbstentionValidationRecords": int(
(~validation_scorable).sum().item()
),
"bestValidationAccuracy": best_accuracy,
"bestValidationSelectionScore": best_selection_score,
"initialValidation": initial_metrics,
"epochsCompleted": len(history),
"stoppedEarly": stopped_early,
@@ -777,6 +997,10 @@ def build_parser() -> argparse.ArgumentParser:
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)
@@ -809,6 +1033,24 @@ def main(argv: Sequence[str] | None = None) -> int:
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.distillation_weight <= 1.0:
parser.error("--distillation-weight must be in [0, 1]")
if args.distillation_temperature <= 0.0:
parser.error("--distillation-temperature must be positive")
if not 0.0 <= args.distillation_selection_weight <= 1.0:
parser.error("--distillation-selection-weight must be in [0, 1]")
if (args.distillation_cache is None) != (args.distillation_weight == 0.0):
parser.error(
"--distillation-cache and a positive --distillation-weight "
"must be supplied together"
)
if (
args.distillation_cache is None
and args.distillation_selection_weight != 0.0
):
parser.error(
"--distillation-selection-weight requires --distillation-cache"
)
if not 0.0 < args.accepted_precision <= args.high_precision <= 1.0:
parser.error(
"precision targets must satisfy 0 < accepted <= high <= 1"