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
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.
+14
View File
@@ -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
+2
View File
@@ -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):
+18
View File
@@ -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:
+159 -13
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)
@@ -532,6 +641,12 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
f"loss={mean[0]:.4f} primary={mean[1]:.4f} "
f"secondary={mean[2]:.4f} mixed={mean[3]:.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"