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

This commit is contained in:
2026-07-31 14:30:10 -07:00
parent 4ed4763557
commit 3d35f4953f
5 changed files with 248 additions and 15 deletions
+160 -14
View File
@@ -37,7 +37,7 @@ from train import (
choose_confidence_thresholds,
expected_calibration_error,
)
from train_mlx import _configure_mlx_device, _linear_schedule
from train_mlx import _configure_mlx_device, _linear_schedule, _teacher_cache
SCRIPT_DIR = Path(__file__).resolve().parent
@@ -254,6 +254,39 @@ def _softmax(values: np.ndarray) -> np.ndarray:
return exponentials / exponentials.sum(axis=-1, keepdims=True)
def _distillation_loss(
mx: Any,
student_logits: Any,
teacher_logits: Any,
*,
temperature: float,
weights: Any,
) -> Any:
"""Return weighted teacher-to-student KL loss for one MLX batch."""
student_log_probabilities = (
student_logits / temperature
- mx.logsumexp(student_logits / temperature, axis=-1, keepdims=True)
)
teacher_probabilities = mx.softmax(
teacher_logits / temperature,
axis=-1,
)
teacher_log_probabilities = mx.log(
mx.maximum(teacher_probabilities, 1e-12)
)
per_record = (
mx.sum(
teacher_probabilities
* (teacher_log_probabilities - student_log_probabilities),
axis=-1,
)
* temperature
* temperature
)
return mx.sum(per_record * weights) / mx.sum(weights)
def _calibration(
outputs: dict[str, np.ndarray],
records: Sequence[dict[str, Any]],
@@ -379,8 +412,20 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
encoded_validation = _encode_records(tokenizer, validation_records)
train_targets = encode_targets(train_records)
validation_targets = encode_targets(validation_records)
validation_scorable = np.asarray(
[record["slice"] != "vague-eval" for record in validation_records],
dtype=np.bool_,
)
sample_weights = _sample_weights(train_records, args.hard_weight)
secondary_class_weights = _secondary_class_weights(train_records)
teacher_train_logits = None
teacher_validation_logits = None
if args.distillation_cache is not None:
teacher_train_logits, teacher_validation_logits = _teacher_cache(
args.distillation_cache.expanduser(),
train_records,
validation_records,
)
non_mixed = len(train_records) - int(train_targets.mixed.sum())
mixed_positive_weight = math.sqrt(
non_mixed / max(float(train_targets.mixed.sum()), 1.0)
@@ -439,7 +484,8 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
mixed: Any,
difficulty: Any,
weights: Any,
) -> tuple[Any, Any, Any, Any, Any]:
teacher_logits: Any | None,
) -> tuple[Any, Any, Any, Any, Any, Any]:
output = model(input_ids=input_ids, attention_mask=attention_mask)
primary_per_record = nn.losses.cross_entropy(
output["purpose_logits"],
@@ -447,7 +493,20 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
label_smoothing=args.label_smoothing,
reduction="none",
)
primary_loss = mx.sum(primary_per_record * weights) / mx.sum(weights)
primary_label_loss = mx.sum(primary_per_record * weights) / mx.sum(weights)
primary_distillation_loss = mx.zeros_like(primary_label_loss)
if teacher_logits is not None:
primary_distillation_loss = _distillation_loss(
mx,
output["purpose_logits"],
teacher_logits,
temperature=args.distillation_temperature,
weights=weights,
)
primary_loss = (
(1.0 - args.distillation_weight) * primary_label_loss
+ args.distillation_weight * primary_distillation_loss
)
secondary_per_record = nn.losses.cross_entropy(
output["secondary_logits"],
@@ -482,23 +541,67 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
+ args.mixed_loss_weight * mixed_loss
+ args.difficulty_loss_weight * difficulty_loss
)
return total, primary_loss, secondary_loss, mixed_loss, difficulty_loss
return (
total,
primary_label_loss,
secondary_loss,
mixed_loss,
difficulty_loss,
primary_distillation_loss,
)
loss_and_grad = nn.value_and_grad(model, loss_function)
rng = np.random.default_rng(args.seed)
checkpoint_config = _checkpoint_config(source_config, variant)
best_dir = output_dir / "model"
best_score = float("-inf")
best_metrics: dict[str, Any] | None = None
epochs_without_improvement = 0
stopped_early = False
history: list[dict[str, Any]] = []
started = time.perf_counter()
initial_outputs = _evaluate(
mx, model, encoded_validation, eval_batch_size
)
initial_metrics = multitask_metrics(initial_outputs, validation_records)
if teacher_validation_logits is not None:
initial_metrics["teacherAgreement"] = float(
np.mean(
initial_outputs["purpose_logits"][validation_scorable].argmax(
axis=-1
)
== teacher_validation_logits[validation_scorable].argmax(axis=-1)
)
)
initial_metrics["epoch"] = 0
best_score = float(initial_metrics["selectionScore"])
best_metrics: dict[str, Any] = initial_metrics
_save_checkpoint(mx, model, tokenizer, best_dir, checkpoint_config)
write_json(
output_dir / "training-state.json",
{
"bestEpoch": 0,
"bestSelectionScore": best_score,
"elapsedSeconds": time.perf_counter() - started,
"complete": False,
},
)
print(
f"epoch 0: primary_accuracy={initial_metrics['primary']['accuracy']:.4%} "
f"hard_accuracy={initial_metrics['primaryHardSlice']['accuracy']:.4%} "
f"mixed_f1={initial_metrics['mixed']['f1']:.4%} "
f"selection_score={best_score:.4%}"
+ (
f" teacher_agreement={initial_metrics['teacherAgreement']:.4%}"
if "teacherAgreement" in initial_metrics
else ""
),
flush=True,
)
for epoch in range(1, args.epochs + 1):
epoch_started = time.perf_counter()
model.train()
running = np.zeros(5, dtype=np.float64)
running = np.zeros(6, dtype=np.float64)
permutation = rng.permutation(len(train_records))
for step, indexes in enumerate(
_batch_indexes(
@@ -509,6 +612,11 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
batch = _mlx_batch(
mx, encoded_train, train_targets, sample_weights, indexes
)
teacher_logits = (
mx.array(teacher_train_logits[indexes])
if teacher_train_logits is not None
else None
)
losses, gradients = loss_and_grad(
batch["input_ids"],
batch["attention_mask"],
@@ -518,6 +626,7 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
batch["mixed"],
batch["difficulty"],
batch["sample_weights"],
teacher_logits,
)
gradients, _ = optim.clip_grad_norm(gradients, args.max_grad_norm)
optimizer.update(model, gradients)
@@ -531,7 +640,13 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
f"epoch {epoch} step {step}/{steps_per_epoch} "
f"loss={mean[0]:.4f} primary={mean[1]:.4f} "
f"secondary={mean[2]:.4f} mixed={mean[3]:.4f} "
f"difficulty={mean[4]:.4f} "
f"difficulty={mean[4]:.4f}"
+ (
f" distillation={mean[5]:.4f}"
if teacher_train_logits is not None
else ""
)
+ " "
f"elapsed={time.perf_counter() - epoch_started:.1f}s",
flush=True,
)
@@ -540,22 +655,34 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
mx, model, encoded_validation, eval_batch_size
)
metrics = multitask_metrics(outputs, validation_records)
if teacher_validation_logits is not None:
metrics["teacherAgreement"] = float(
np.mean(
outputs["purpose_logits"][validation_scorable].argmax(axis=-1)
== teacher_validation_logits[validation_scorable].argmax(axis=-1)
)
)
metrics["epoch"] = epoch
metrics["meanTrainingLoss"] = (running / steps_per_epoch).tolist()
history.append(metrics)
score = float(metrics["selectionScore"])
secondary_macro = (
metrics["secondary"]["macroRecall"]
metrics["secondary"]["supportedMacroRecall"]
if metrics["secondary"] is not None
else 0.0
)
print(
f"epoch {epoch}: primary_accuracy={metrics['primary']['accuracy']:.4%} "
f"hard_accuracy={metrics['primaryHardSlice']['accuracy']:.4%} "
f"secondary_macro_recall={secondary_macro:.4%} "
f"secondary_supported_macro_recall={secondary_macro:.4%} "
f"mixed_f1={metrics['mixed']['f1']:.4%} "
f"difficulty_mae={metrics['difficulty']['mae']:.4f} "
f"selection_score={score:.4%}",
f"selection_score={score:.4%}"
+ (
f" teacher_agreement={metrics['teacherAgreement']:.4%}"
if "teacherAgreement" in metrics
else ""
),
flush=True,
)
improvement = score - best_score
@@ -588,9 +715,6 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
)
break
if best_metrics is None:
raise DataError("purpose-deep training did not produce a checkpoint")
# Release the optimizer graph before opening the selected checkpoint; base and
# especially large should never hold two full optimizer states at calibration time.
del optimizer, loss_and_grad, model
@@ -632,6 +756,15 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
"mixed": args.mixed_loss_weight,
"difficulty": args.difficulty_loss_weight,
},
"distillation": {
"cache": (
str(args.distillation_cache)
if args.distillation_cache is not None
else None
),
"weight": args.distillation_weight,
"temperature": args.distillation_temperature,
},
"secondaryClassWeights": {
label: float(secondary_class_weights[index])
for index, label in enumerate(LABELS)
@@ -641,6 +774,7 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
"batchSize": batch_size,
"learningRate": learning_rate,
"bestValidationSelectionScore": best_score,
"initialValidation": initial_metrics,
"bestValidation": best_metrics,
"selectedValidation": calibrated,
"epochsCompleted": len(history),
@@ -707,6 +841,9 @@ def build_parser() -> argparse.ArgumentParser:
parser.add_argument("--secondary-loss-weight", type=float, default=0.25)
parser.add_argument("--mixed-loss-weight", type=float, default=0.25)
parser.add_argument("--difficulty-loss-weight", type=float, default=0.10)
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("--progress-steps", type=int, default=25)
parser.add_argument("--early-stopping-patience", type=int, default=1)
parser.add_argument("--minimum-improvement", type=float, default=0.0005)
@@ -750,6 +887,15 @@ def main(argv: Sequence[str] | None = None) -> int:
):
if getattr(args, name) < 0:
parser.error(f"--{name.replace('_', '-')} must be non-negative")
if not 0 <= args.distillation_weight <= 1:
parser.error("--distillation-weight must be in [0, 1]")
if args.distillation_temperature <= 0:
parser.error("--distillation-temperature must be positive")
if (args.distillation_cache is None) != (args.distillation_weight == 0):
parser.error(
"--distillation-cache and a positive --distillation-weight "
"must be supplied together"
)
if not 0 < args.accepted_precision <= args.high_precision <= 1:
parser.error(
"confidence precision targets must satisfy 0 < accepted <= high <= 1"