diff --git a/README.md b/README.md index f09f911..41b1dd5 100644 --- a/README.md +++ b/README.md @@ -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 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 -ended, so continue its selected checkpoint conservatively before changing architecture: +ended, so its selected checkpoint was continued conservatively before changing +architecture: ```bash ml/purpose-classifier/venv/bin/python -u \ @@ -350,6 +351,58 @@ ml/purpose-classifier/venv/bin/python -u \ --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 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. diff --git a/deep_contract.py b/deep_contract.py index d5bb4f7..1355ebb 100644 --- a/deep_contract.py +++ b/deep_contract.py @@ -274,6 +274,20 @@ def multitask_metrics( targets.secondary[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: secondary = None diff --git a/tests/test_deep_contract.py b/tests/test_deep_contract.py index 7cabcc0..2fecd94 100644 --- a/tests/test_deep_contract.py +++ b/tests/test_deep_contract.py @@ -121,6 +121,8 @@ class TargetAndMetricTests(unittest.TestCase): self.assertEqual(0, metrics["primaryHardSlice"]["accuracy"]) self.assertAlmostEqual(1 / 6, metrics["selectionScore"]) self.assertEqual(1, metrics["secondary"]["accuracy"]) + self.assertEqual(["writing"], metrics["secondary"]["supportedLabels"]) + self.assertEqual(1, metrics["secondary"]["supportedMacroRecall"]) self.assertEqual(1, metrics["mixed"]["f1"]) def test_mixed_threshold_prefers_higher_threshold_on_f1_tie(self): diff --git a/tests/test_deep_model_mlx.py b/tests/test_deep_model_mlx.py index e4cabfa..7fbc059 100644 --- a/tests/test_deep_model_mlx.py +++ b/tests/test_deep_model_mlx.py @@ -18,6 +18,7 @@ from deep_model_mlx import ( save_weights, ) from purpose_data import DataError, LABELS +from train_deep_mlx import _distillation_loss def tiny_config(*, checkpointing=False): @@ -97,6 +98,23 @@ class DeepModelTests(unittest.TestCase): mx.eval(value, gradients) 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): model = ModernBertForPurposeClassification(tiny_config()) with tempfile.TemporaryDirectory() as temp: diff --git a/train_deep_mlx.py b/train_deep_mlx.py index c4bd3ab..d52b9b2 100644 --- a/train_deep_mlx.py +++ b/train_deep_mlx.py @@ -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"