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
+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"