Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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())
|
||||
@@ -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."
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user