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
|
iteration should incorporate reviewed boundary data and be selected on a revised
|
||||||
validation/frozen dataset version.
|
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:
|
For a wiring smoke test, use a small deterministic prefix:
|
||||||
|
|
||||||
```bash
|
```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
|
completion rule are recorded in `data/curation-review-v1.json`; the semantic report is
|
||||||
versioned as `data/semantic-audit-v1.json`.
|
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
|
The dataset owner accepted the curated generated labels and difficulty metadata as-is on
|
||||||
running the embedding audit again:
|
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
|
```bash
|
||||||
ml/purpose-classifier/.venv/bin/python ml/purpose-classifier/review_data.py
|
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": {
|
"humanLabelAndDifficultyReview": {
|
||||||
"status": "planned",
|
"status": "accepted-as-generated",
|
||||||
"populationRecords": 12193,
|
"populationRecords": 12193,
|
||||||
"sampleFraction": 0.1,
|
"sampleFraction": 0.1,
|
||||||
"sampleRecords": 1219,
|
"sampleRecords": 1219,
|
||||||
@@ -80,6 +80,8 @@
|
|||||||
"notes"
|
"notes"
|
||||||
],
|
],
|
||||||
"generatedArtifact": "ml/purpose-classifier/.artifacts/human-review-v1.csv",
|
"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,
|
"nearDuplicateThreshold": 0.92,
|
||||||
"retainedRecords": 12193,
|
"retainedRecords": 12193,
|
||||||
"reviewPath": "ml/purpose-classifier/data/curation-review-v1.json",
|
"reviewPath": "ml/purpose-classifier/data/curation-review-v1.json",
|
||||||
"reviewSha256": "625251a98bcdcde0bee074e3bba3c564af4347eb7955497a9c3e018aae00b30f",
|
"reviewSha256": "1721ab77a0f723a4b03f42b14850b5c351ead2cd8fa72bb2d1fe3f0d2bb568da",
|
||||||
"vagueEvalPolicy": "validation/test only"
|
"vagueEvalPolicy": "validation/test only"
|
||||||
},
|
},
|
||||||
"datasetVersion": "purpose-dataset-v1",
|
"datasetVersion": "purpose-dataset-v1",
|
||||||
|
|||||||
@@ -65,6 +65,17 @@ def reviewed_population(args: argparse.Namespace) -> list[SourceRecord]:
|
|||||||
).records
|
).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:
|
def build_parser() -> argparse.ArgumentParser:
|
||||||
parser = argparse.ArgumentParser(description=__doc__)
|
parser = argparse.ArgumentParser(description=__doc__)
|
||||||
parser.add_argument("--source", action="append", type=Path)
|
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)
|
args = build_parser().parse_args(argv)
|
||||||
try:
|
try:
|
||||||
population = reviewed_population(args)
|
population = reviewed_population(args)
|
||||||
|
policy_status = review_policy_status(args.curation_review.resolve())
|
||||||
sample = stratified_review_sample(
|
sample = stratified_review_sample(
|
||||||
population,
|
population,
|
||||||
fraction=args.review_fraction,
|
fraction=args.review_fraction,
|
||||||
@@ -153,6 +165,7 @@ def main(argv: Sequence[str] | None = None) -> int:
|
|||||||
f"Human review: {progress.completed}/{progress.records} complete "
|
f"Human review: {progress.completed}/{progress.records} complete "
|
||||||
f"({progress.accepted} accept, {progress.relabeled} relabel, "
|
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
|
return 0
|
||||||
|
|
||||||
|
|||||||
@@ -12,6 +12,19 @@ import train
|
|||||||
|
|
||||||
|
|
||||||
class MetricsTests(unittest.TestCase):
|
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):
|
def test_qat_replacements_keep_checkpoint_keys_and_gradients(self):
|
||||||
model = torch.nn.Sequential(
|
model = torch.nn.Sequential(
|
||||||
torch.nn.Embedding(16, 8),
|
torch.nn.Embedding(16, 8),
|
||||||
|
|||||||
@@ -12,7 +12,14 @@ import time
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Sequence
|
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
|
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
|
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]:
|
def enable_quantization_aware_training(torch: Any, model: Any) -> dict[str, int]:
|
||||||
"""Mirror the export graph's int8 policy with straight-through fake quantization.
|
"""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())
|
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()
|
model.eval()
|
||||||
all_logits = []
|
all_logits = []
|
||||||
all_labels = []
|
all_labels = []
|
||||||
|
started = time.perf_counter()
|
||||||
with torch.inference_mode():
|
with torch.inference_mode():
|
||||||
for batch in loader:
|
for step, batch in enumerate(loader, 1):
|
||||||
labels = batch.pop("labels")
|
labels = batch.pop("labels")
|
||||||
batch.pop("sample_weights", None)
|
batch.pop("sample_weights", None)
|
||||||
|
batch.pop("teacher_logits", None)
|
||||||
inputs = {key: value.to(device) for key, value in batch.items()}
|
inputs = {key: value.to(device) for key, value in batch.items()}
|
||||||
logits = model(**inputs).logits.cpu()
|
logits = model(**inputs).logits.cpu()
|
||||||
all_logits.append(logits)
|
all_logits.append(logits)
|
||||||
all_labels.append(labels)
|
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)
|
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
|
output_dir: Path = args.output_dir
|
||||||
local_model = Path(args.model).expanduser()
|
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:
|
try:
|
||||||
local_model.resolve().relative_to(output_dir.resolve())
|
local_path.resolve().relative_to(output_dir.resolve())
|
||||||
except ValueError:
|
except ValueError:
|
||||||
pass
|
pass
|
||||||
else:
|
else:
|
||||||
raise DataError(
|
raise DataError(
|
||||||
"local --model must not be inside --output-dir; overwrite could "
|
f"local {option_name} must not be inside --output-dir; overwrite "
|
||||||
"destroy the continuation checkpoint"
|
"could destroy the continuation checkpoint"
|
||||||
)
|
)
|
||||||
if output_dir.exists() and any(output_dir.iterdir()):
|
if output_dir.exists() and any(output_dir.iterdir()):
|
||||||
if not args.overwrite_output:
|
if not args.overwrite_output:
|
||||||
@@ -505,31 +576,88 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
model.to(device)
|
model.to(device)
|
||||||
|
|
||||||
class PromptDataset(Dataset):
|
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.records = records
|
||||||
|
self.teacher_logits = teacher_logits
|
||||||
|
|
||||||
def __len__(self) -> int:
|
def __len__(self) -> int:
|
||||||
return len(self.records)
|
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]
|
record = self.records[index]
|
||||||
return (
|
return (
|
||||||
prepare_text(record["prompt"]),
|
prepare_text(record["prompt"]),
|
||||||
label_to_id[record["purpose"]],
|
label_to_id[record["purpose"]],
|
||||||
training_weight(record, args.boundary_weight),
|
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]:
|
def collate(
|
||||||
texts, labels, weights = zip(*items)
|
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 = encode_fixed_shape(tokenizer, list(texts), torch)
|
||||||
encoded["labels"] = torch.tensor(labels, dtype=torch.long)
|
encoded["labels"] = torch.tensor(labels, dtype=torch.long)
|
||||||
encoded["sample_weights"] = torch.tensor(weights, dtype=torch.float)
|
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
|
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 = torch.Generator()
|
||||||
generator.manual_seed(args.seed)
|
generator.manual_seed(args.seed)
|
||||||
train_loader = DataLoader(
|
train_loader = DataLoader(
|
||||||
PromptDataset(train_records),
|
PromptDataset(train_records, teacher_train_logits),
|
||||||
batch_size=args.batch_size,
|
batch_size=args.batch_size,
|
||||||
shuffle=True,
|
shuffle=True,
|
||||||
generator=generator,
|
generator=generator,
|
||||||
@@ -567,7 +695,42 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
if scorable
|
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_accuracy = initial_metrics["accuracy"]
|
||||||
|
best_selection_score = initial_selection_score
|
||||||
epochs_without_improvement = 0
|
epochs_without_improvement = 0
|
||||||
stopped_early = False
|
stopped_early = False
|
||||||
history = []
|
history = []
|
||||||
@@ -576,7 +739,13 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
tokenizer.save_pretrained(best_dir)
|
tokenizer.save_pretrained(best_dir)
|
||||||
print(
|
print(
|
||||||
f"epoch 0: validation_accuracy={initial_metrics['accuracy']:.4%} "
|
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,
|
flush=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -599,20 +768,46 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
model.train()
|
model.train()
|
||||||
optimizer.zero_grad(set_to_none=True)
|
optimizer.zero_grad(set_to_none=True)
|
||||||
running_loss = 0.0
|
running_loss = 0.0
|
||||||
|
running_label_loss = 0.0
|
||||||
|
running_distillation_loss = 0.0
|
||||||
for step, batch in enumerate(train_loader, 1):
|
for step, batch in enumerate(train_loader, 1):
|
||||||
labels = batch.pop("labels").to(device)
|
labels = batch.pop("labels").to(device)
|
||||||
sample_weights = batch.pop("sample_weights").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()}
|
inputs = {key: value.to(device) for key, value in batch.items()}
|
||||||
per_record_loss = torch.nn.functional.cross_entropy(
|
student_logits = model(**inputs).logits
|
||||||
model(**inputs).logits,
|
label_loss = torch.nn.functional.cross_entropy(
|
||||||
|
student_logits,
|
||||||
labels,
|
labels,
|
||||||
reduction="none",
|
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 = (
|
loss = (
|
||||||
(per_record_loss * sample_weights).sum() / sample_weights.sum()
|
(per_record_loss * sample_weights).sum() / sample_weights.sum()
|
||||||
) / args.gradient_accumulation_steps
|
) / args.gradient_accumulation_steps
|
||||||
loss.backward()
|
loss.backward()
|
||||||
running_loss += float(loss.item()) * args.gradient_accumulation_steps
|
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 = (
|
should_update = (
|
||||||
step % args.gradient_accumulation_steps == 0
|
step % args.gradient_accumulation_steps == 0
|
||||||
or step == len(train_loader)
|
or step == len(train_loader)
|
||||||
@@ -628,6 +823,8 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
print(
|
print(
|
||||||
f"epoch {epoch} step {step}/{len(train_loader)} "
|
f"epoch {epoch} step {step}/{len(train_loader)} "
|
||||||
f"mean_loss={running_loss / step:.4f} "
|
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",
|
f"elapsed={time.perf_counter() - epoch_started:.1f}s",
|
||||||
flush=True,
|
flush=True,
|
||||||
)
|
)
|
||||||
@@ -643,18 +840,34 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
if scorable
|
if scorable
|
||||||
]
|
]
|
||||||
metrics = classification_metrics(scored_labels, scored_predictions)
|
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["epoch"] = epoch
|
||||||
metrics["meanTrainingLoss"] = running_loss / len(train_loader)
|
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)
|
history.append(metrics)
|
||||||
print(
|
print(
|
||||||
f"epoch {epoch}: loss={metrics['meanTrainingLoss']:.4f} "
|
f"epoch {epoch}: loss={metrics['meanTrainingLoss']:.4f} "
|
||||||
f"validation_accuracy={metrics['accuracy']:.4%} "
|
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,
|
flush=True,
|
||||||
)
|
)
|
||||||
improvement = metrics["accuracy"] - best_accuracy
|
improvement = candidate_selection_score - best_selection_score
|
||||||
if improvement > args.minimum_improvement:
|
if improvement > args.minimum_improvement:
|
||||||
best_accuracy = metrics["accuracy"]
|
best_accuracy = metrics["accuracy"]
|
||||||
|
best_selection_score = candidate_selection_score
|
||||||
epochs_without_improvement = 0
|
epochs_without_improvement = 0
|
||||||
model.save_pretrained(best_dir, safe_serialization=True)
|
model.save_pretrained(best_dir, safe_serialization=True)
|
||||||
tokenizer.save_pretrained(best_dir)
|
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:
|
if epochs_without_improvement >= args.early_stopping_patience:
|
||||||
stopped_early = True
|
stopped_early = True
|
||||||
print(
|
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"greater than {args.minimum_improvement:.4%} for "
|
||||||
f"{args.early_stopping_patience} epoch(s)",
|
f"{args.early_stopping_patience} epoch(s)",
|
||||||
flush=True,
|
flush=True,
|
||||||
@@ -721,12 +934,19 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
|
|||||||
"boundaryTrainingWeight": args.boundary_weight,
|
"boundaryTrainingWeight": args.boundary_weight,
|
||||||
"quantizationAwareTraining": args.quantization_aware,
|
"quantizationAwareTraining": args.quantization_aware,
|
||||||
"quantizationAwareModules": qat_modules,
|
"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),
|
"validationRecords": len(validation_records),
|
||||||
"scoredValidationRecords": int(validation_scorable.sum().item()),
|
"scoredValidationRecords": int(validation_scorable.sum().item()),
|
||||||
"vagueAbstentionValidationRecords": int(
|
"vagueAbstentionValidationRecords": int(
|
||||||
(~validation_scorable).sum().item()
|
(~validation_scorable).sum().item()
|
||||||
),
|
),
|
||||||
"bestValidationAccuracy": best_accuracy,
|
"bestValidationAccuracy": best_accuracy,
|
||||||
|
"bestValidationSelectionScore": best_selection_score,
|
||||||
"initialValidation": initial_metrics,
|
"initialValidation": initial_metrics,
|
||||||
"epochsCompleted": len(history),
|
"epochsCompleted": len(history),
|
||||||
"stoppedEarly": stopped_early,
|
"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("--minimum-improvement", type=float, default=0.0005)
|
||||||
parser.add_argument("--boundary-weight", type=float, default=1.0)
|
parser.add_argument("--boundary-weight", type=float, default=1.0)
|
||||||
parser.add_argument("--quantization-aware", action="store_true")
|
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("--high-precision", type=float, default=0.98)
|
||||||
parser.add_argument("--accepted-precision", type=float, default=0.95)
|
parser.add_argument("--accepted-precision", type=float, default=0.95)
|
||||||
parser.add_argument("--max-train-records", type=int)
|
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")
|
parser.error("--progress-steps must be non-negative")
|
||||||
if args.boundary_weight <= 0.0:
|
if args.boundary_weight <= 0.0:
|
||||||
parser.error("--boundary-weight must be positive")
|
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:
|
if not 0.0 < args.accepted_precision <= args.high_precision <= 1.0:
|
||||||
parser.error(
|
parser.error(
|
||||||
"precision targets must satisfy 0 < accepted <= high <= 1"
|
"precision targets must satisfy 0 < accepted <= high <= 1"
|
||||||
|
|||||||
Reference in New Issue
Block a user