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
+54 -1
View File
@@ -331,7 +331,8 @@ To continue a completed run without discarding its trained task heads, pass its
Continuation restores the backbone and all four heads strictly, then starts a fresh Continuation restores the backbone and all four heads strictly, then starts a fresh
optimizer and learning-rate schedule; `--model` remains reserved for an untrained local optimizer and learning-rate schedule; `--model` remains reserved for an untrained local
upstream checkpoint. The first base run was still improving when its three-epoch schedule upstream checkpoint. The first base run was still improving when its three-epoch schedule
ended, so continue its selected checkpoint conservatively before changing architecture: ended, so its selected checkpoint was continued conservatively before changing
architecture:
```bash ```bash
ml/purpose-classifier/venv/bin/python -u \ ml/purpose-classifier/venv/bin/python -u \
@@ -350,6 +351,58 @@ ml/purpose-classifier/venv/bin/python -u \
--overwrite-output --overwrite-output
``` ```
That continuation reached 80.31% primary and 78.28% hard-slice validation accuracy;
calibrated mixed F1 reached 67.12%. It remained far below purpose-lite, while primary
training loss and validation accuracy were still improving. Do not chain another plain
continuation. The bounded next experiment distills the mature purpose-lite boundary
teacher into the continued deep checkpoint while retaining direct primary labels and all
three auxiliary losses.
Create a teacher cache bound to the history-augmented split:
```bash
ml/purpose-classifier/venv/bin/python -u \
ml/purpose-classifier/cache_teacher.py \
--dataset-dir \
ml/purpose-classifier/.artifacts/dataset-v1-history-first-prompts \
--model \
ml/purpose-classifier/outputs/purpose-lite-v1-boundary-tune/model \
--output \
ml/purpose-classifier/outputs/purpose-lite-v1-history-first-prompts-teacher.pt \
--device mps \
--batch-size 16 \
--progress-steps 25 \
--overwrite-output
```
Then run two validation-selected distilled continuation epochs:
```bash
ml/purpose-classifier/venv/bin/python -u \
ml/purpose-classifier/train_deep_mlx.py \
--variant base \
--resume-from \
ml/purpose-classifier/outputs/purpose-deep-v1-base-mlx-cont-3e/model \
--dataset-dir \
ml/purpose-classifier/.artifacts/dataset-v1-history-first-prompts \
--distillation-cache \
ml/purpose-classifier/outputs/purpose-lite-v1-history-first-prompts-teacher.pt \
--distillation-weight 0.5 \
--distillation-temperature 2 \
--epochs 2 \
--learning-rate 1e-5 \
--early-stopping-patience 1 \
--progress-steps 10 \
--output-dir \
ml/purpose-classifier/outputs/purpose-deep-v1-base-mlx-distilled \
--overwrite-output
```
The trainer records the resumed checkpoint as epoch zero before updating anything, so a
distillation regression cannot overwrite the 80.31% candidate. Teacher agreement is
reported for diagnosis but does not enter deep checkpoint selection; overall and hard
primary label accuracy remain the only selection inputs.
Do not launch the large rung yet. It is justified only after base is evaluated on the Do not launch the large rung yet. It is justified only after base is evaluated on the
frozen set; large must beat base by at least two hard-slice points, while deep itself must frozen set; large must beat base by at least two hard-slice points, while deep itself must
reach 97% scored overall and beat the shipping lite artifact by five hard-slice points. reach 97% scored overall and beat the shipping lite artifact by five hard-slice points.
+14
View File
@@ -274,6 +274,20 @@ def multitask_metrics(
targets.secondary[targets.secondary_mask].tolist(), targets.secondary[targets.secondary_mask].tolist(),
secondary_predictions[targets.secondary_mask].tolist(), secondary_predictions[targets.secondary_mask].tolist(),
) )
supported_labels = [
label
for label, row in zip(
secondary["confusionMatrix"]["labels"],
secondary["confusionMatrix"]["rows"],
)
if sum(row)
]
secondary["supportedLabels"] = supported_labels
secondary["supportedMacroRecall"] = float(
np.mean(
[secondary["perPurposeRecall"][label] for label in supported_labels]
)
)
else: else:
secondary = None secondary = None
+2
View File
@@ -121,6 +121,8 @@ class TargetAndMetricTests(unittest.TestCase):
self.assertEqual(0, metrics["primaryHardSlice"]["accuracy"]) self.assertEqual(0, metrics["primaryHardSlice"]["accuracy"])
self.assertAlmostEqual(1 / 6, metrics["selectionScore"]) self.assertAlmostEqual(1 / 6, metrics["selectionScore"])
self.assertEqual(1, metrics["secondary"]["accuracy"]) self.assertEqual(1, metrics["secondary"]["accuracy"])
self.assertEqual(["writing"], metrics["secondary"]["supportedLabels"])
self.assertEqual(1, metrics["secondary"]["supportedMacroRecall"])
self.assertEqual(1, metrics["mixed"]["f1"]) self.assertEqual(1, metrics["mixed"]["f1"])
def test_mixed_threshold_prefers_higher_threshold_on_f1_tie(self): def test_mixed_threshold_prefers_higher_threshold_on_f1_tie(self):
+18
View File
@@ -18,6 +18,7 @@ from deep_model_mlx import (
save_weights, save_weights,
) )
from purpose_data import DataError, LABELS from purpose_data import DataError, LABELS
from train_deep_mlx import _distillation_loss
def tiny_config(*, checkpointing=False): def tiny_config(*, checkpointing=False):
@@ -97,6 +98,23 @@ class DeepModelTests(unittest.TestCase):
mx.eval(value, gradients) mx.eval(value, gradients)
self.assertTrue(float(value.item()) > 0) self.assertTrue(float(value.item()) > 0)
def test_distillation_loss_matches_teacher_and_backpropagates(self):
teacher = mx.array([[2.0, 0.0, -1.0]])
def loss(student):
return _distillation_loss(
mx,
student,
teacher,
temperature=2.0,
weights=mx.ones((1,)),
)
value, gradient = mx.value_and_grad(loss)(teacher)
mx.eval(value, gradient)
self.assertAlmostEqual(0.0, float(value.item()), places=6)
self.assertEqual(teacher.shape, gradient.shape)
def test_checkpoint_round_trip(self): def test_checkpoint_round_trip(self):
model = ModernBertForPurposeClassification(tiny_config()) model = ModernBertForPurposeClassification(tiny_config())
with tempfile.TemporaryDirectory() as temp: with tempfile.TemporaryDirectory() as temp:
+160 -14
View File
@@ -37,7 +37,7 @@ from train import (
choose_confidence_thresholds, choose_confidence_thresholds,
expected_calibration_error, 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 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) 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( def _calibration(
outputs: dict[str, np.ndarray], outputs: dict[str, np.ndarray],
records: Sequence[dict[str, Any]], records: Sequence[dict[str, Any]],
@@ -379,8 +412,20 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
encoded_validation = _encode_records(tokenizer, validation_records) encoded_validation = _encode_records(tokenizer, validation_records)
train_targets = encode_targets(train_records) train_targets = encode_targets(train_records)
validation_targets = encode_targets(validation_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) sample_weights = _sample_weights(train_records, args.hard_weight)
secondary_class_weights = _secondary_class_weights(train_records) 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()) non_mixed = len(train_records) - int(train_targets.mixed.sum())
mixed_positive_weight = math.sqrt( mixed_positive_weight = math.sqrt(
non_mixed / max(float(train_targets.mixed.sum()), 1.0) non_mixed / max(float(train_targets.mixed.sum()), 1.0)
@@ -439,7 +484,8 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
mixed: Any, mixed: Any,
difficulty: Any, difficulty: Any,
weights: 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) output = model(input_ids=input_ids, attention_mask=attention_mask)
primary_per_record = nn.losses.cross_entropy( primary_per_record = nn.losses.cross_entropy(
output["purpose_logits"], output["purpose_logits"],
@@ -447,7 +493,20 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
label_smoothing=args.label_smoothing, label_smoothing=args.label_smoothing,
reduction="none", 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( secondary_per_record = nn.losses.cross_entropy(
output["secondary_logits"], output["secondary_logits"],
@@ -482,23 +541,67 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
+ args.mixed_loss_weight * mixed_loss + args.mixed_loss_weight * mixed_loss
+ args.difficulty_loss_weight * difficulty_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) loss_and_grad = nn.value_and_grad(model, loss_function)
rng = np.random.default_rng(args.seed) rng = np.random.default_rng(args.seed)
checkpoint_config = _checkpoint_config(source_config, variant) checkpoint_config = _checkpoint_config(source_config, variant)
best_dir = output_dir / "model" best_dir = output_dir / "model"
best_score = float("-inf")
best_metrics: dict[str, Any] | None = None
epochs_without_improvement = 0 epochs_without_improvement = 0
stopped_early = False stopped_early = False
history: list[dict[str, Any]] = [] history: list[dict[str, Any]] = []
started = time.perf_counter() 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): for epoch in range(1, args.epochs + 1):
epoch_started = time.perf_counter() epoch_started = time.perf_counter()
model.train() model.train()
running = np.zeros(5, dtype=np.float64) running = np.zeros(6, dtype=np.float64)
permutation = rng.permutation(len(train_records)) permutation = rng.permutation(len(train_records))
for step, indexes in enumerate( for step, indexes in enumerate(
_batch_indexes( _batch_indexes(
@@ -509,6 +612,11 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
batch = _mlx_batch( batch = _mlx_batch(
mx, encoded_train, train_targets, sample_weights, indexes 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( losses, gradients = loss_and_grad(
batch["input_ids"], batch["input_ids"],
batch["attention_mask"], batch["attention_mask"],
@@ -518,6 +626,7 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
batch["mixed"], batch["mixed"],
batch["difficulty"], batch["difficulty"],
batch["sample_weights"], batch["sample_weights"],
teacher_logits,
) )
gradients, _ = optim.clip_grad_norm(gradients, args.max_grad_norm) gradients, _ = optim.clip_grad_norm(gradients, args.max_grad_norm)
optimizer.update(model, gradients) 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"epoch {epoch} step {step}/{steps_per_epoch} "
f"loss={mean[0]:.4f} primary={mean[1]:.4f} " f"loss={mean[0]:.4f} primary={mean[1]:.4f} "
f"secondary={mean[2]:.4f} mixed={mean[3]:.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", f"elapsed={time.perf_counter() - epoch_started:.1f}s",
flush=True, flush=True,
) )
@@ -540,22 +655,34 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
mx, model, encoded_validation, eval_batch_size mx, model, encoded_validation, eval_batch_size
) )
metrics = multitask_metrics(outputs, validation_records) 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["epoch"] = epoch
metrics["meanTrainingLoss"] = (running / steps_per_epoch).tolist() metrics["meanTrainingLoss"] = (running / steps_per_epoch).tolist()
history.append(metrics) history.append(metrics)
score = float(metrics["selectionScore"]) score = float(metrics["selectionScore"])
secondary_macro = ( secondary_macro = (
metrics["secondary"]["macroRecall"] metrics["secondary"]["supportedMacroRecall"]
if metrics["secondary"] is not None if metrics["secondary"] is not None
else 0.0 else 0.0
) )
print( print(
f"epoch {epoch}: primary_accuracy={metrics['primary']['accuracy']:.4%} " f"epoch {epoch}: primary_accuracy={metrics['primary']['accuracy']:.4%} "
f"hard_accuracy={metrics['primaryHardSlice']['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"mixed_f1={metrics['mixed']['f1']:.4%} "
f"difficulty_mae={metrics['difficulty']['mae']:.4f} " 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, flush=True,
) )
improvement = score - best_score improvement = score - best_score
@@ -588,9 +715,6 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
) )
break 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 # Release the optimizer graph before opening the selected checkpoint; base and
# especially large should never hold two full optimizer states at calibration time. # especially large should never hold two full optimizer states at calibration time.
del optimizer, loss_and_grad, model del optimizer, loss_and_grad, model
@@ -632,6 +756,15 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
"mixed": args.mixed_loss_weight, "mixed": args.mixed_loss_weight,
"difficulty": args.difficulty_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": { "secondaryClassWeights": {
label: float(secondary_class_weights[index]) label: float(secondary_class_weights[index])
for index, label in enumerate(LABELS) for index, label in enumerate(LABELS)
@@ -641,6 +774,7 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
"batchSize": batch_size, "batchSize": batch_size,
"learningRate": learning_rate, "learningRate": learning_rate,
"bestValidationSelectionScore": best_score, "bestValidationSelectionScore": best_score,
"initialValidation": initial_metrics,
"bestValidation": best_metrics, "bestValidation": best_metrics,
"selectedValidation": calibrated, "selectedValidation": calibrated,
"epochsCompleted": len(history), "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("--secondary-loss-weight", type=float, default=0.25)
parser.add_argument("--mixed-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("--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("--progress-steps", type=int, default=25)
parser.add_argument("--early-stopping-patience", type=int, default=1) parser.add_argument("--early-stopping-patience", type=int, default=1)
parser.add_argument("--minimum-improvement", type=float, default=0.0005) 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: if getattr(args, name) < 0:
parser.error(f"--{name.replace('_', '-')} must be non-negative") 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: if not 0 < args.accepted_precision <= args.high_precision <= 1:
parser.error( parser.error(
"confidence precision targets must satisfy 0 < accepted <= high <= 1" "confidence precision targets must satisfy 0 < accepted <= high <= 1"