Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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,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:
|
||||
|
||||
+160
-14
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user