Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -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