diff --git a/README.md b/README.md index 5595870..0e1ceec 100644 --- a/README.md +++ b/README.md @@ -554,6 +554,34 @@ distillation regression cannot overwrite the 80.31% candidate. Teacher agreement reported for diagnosis but does not enter deep checkpoint selection; overall and hard primary label accuracy remain the only selection inputs. +### Pause and resume purpose-deep training + +`train_deep_mlx.py` handles both `Ctrl-C` (`SIGINT`) and `SIGTERM` gracefully. Press +`Ctrl-C` once: the current optimizer step finishes, then the trainer atomically saves the +current model, AdamW state and learning-rate cursor, epoch and next batch, exact shuffle +order, accumulated losses, RNG state, best-selection state, and history under +`/resume/`. It exits with the conventional signal-derived status only after +that checkpoint is durable. A second `Ctrl-C` forces an immediate interruption. + +The pause message prints the complete resume command. For the distilled run above it is: + +```bash +ml/purpose-classifier/venv/bin/python -u \ + ml/purpose-classifier/train_deep_mlx.py \ + --resume-training \ + ml/purpose-classifier/outputs/purpose-deep-v1-base-mlx-distilled/resume +``` + +`--resume-training` is sufficient by itself: it restores the original dataset, +distillation, schedule, loss, and selection arguments from the checkpoint. It also verifies +hashes of the training split, validation split, and teacher cache before loading. The +checkpoint resumes at the next unprocessed batch—even if stopped after the last training +batch but before epoch evaluation. Normal completion removes the large optimizer resume +artifact while retaining the selected `model/` and reports. + +This applies to processes started with the updated script. A process already running older +code cannot acquire signal handling retroactively. + The historical large-rung rule required base to qualify first, large to beat base by at least two hard-slice points, and deep to reach 97% scored overall while beating lite by five hard-slice points. Neither the v1 nor Sol-high v2 base result unlocked that rung. diff --git a/tests/test_deep_model_mlx.py b/tests/test_deep_model_mlx.py index 7fbc059..4540d90 100644 --- a/tests/test_deep_model_mlx.py +++ b/tests/test_deep_model_mlx.py @@ -5,6 +5,8 @@ from pathlib import Path import mlx.core as mx import mlx.nn as nn +import mlx.optimizers as optim +from mlx.utils import tree_flatten MODULE_DIR = Path(__file__).resolve().parents[1] @@ -18,7 +20,11 @@ from deep_model_mlx import ( save_weights, ) from purpose_data import DataError, LABELS -from train_deep_mlx import _distillation_loss +from train_deep_mlx import ( + _distillation_loss, + _load_optimizer_state, + _save_training_resume, +) def tiny_config(*, checkpointing=False): @@ -115,6 +121,54 @@ class DeepModelTests(unittest.TestCase): self.assertAlmostEqual(0.0, float(value.item()), places=6) self.assertEqual(teacher.shape, gradient.shape) + def test_training_resume_round_trips_model_and_optimizer(self): + class Tokenizer: + @staticmethod + def save_pretrained(destination): + (destination / "tokenizer_config.json").write_text("{}\n") + + model = ModernBertForPurposeClassification(tiny_config()) + optimizer = optim.AdamW(learning_rate=1e-3) + optimizer.init(model.trainable_parameters()) + + def loss(ids, mask): + return mx.mean(model(ids, mask)["purpose_logits"] ** 2) + + value_and_grad = nn.value_and_grad(model, loss) + value, gradients = value_and_grad( + mx.array([[1, 3, 4, 2]]), mx.ones((1, 4), dtype=mx.int32) + ) + optimizer.update(model, gradients) + mx.eval(value, model.parameters(), optimizer.state) + + with tempfile.TemporaryDirectory() as temp: + checkpoint = _save_training_resume( + mx, + model, + optimizer, + Tokenizer(), + Path(temp), + {}, + {"schemaVersion": 1, "status": "paused"}, + ) + restored_model = ModernBertForPurposeClassification(tiny_config()) + restored_model.load_weights( + str(checkpoint / "model.safetensors"), strict=True + ) + restored_optimizer = optim.AdamW(learning_rate=1e-3) + restored_optimizer.init(restored_model.trainable_parameters()) + _load_optimizer_state(mx, restored_optimizer, checkpoint) + + original = tree_flatten(optimizer.state, destination={}) + restored = tree_flatten(restored_optimizer.state, destination={}) + self.assertEqual(set(original), set(restored)) + for key in original: + with self.subTest(optimizer_tensor=key): + self.assertEqual( + 0, + float(mx.max(mx.abs(original[key] - restored[key])).item()), + ) + def test_checkpoint_round_trip(self): model = ModernBertForPurposeClassification(tiny_config()) with tempfile.TemporaryDirectory() as temp: diff --git a/tests/test_train_deep_resume.py b/tests/test_train_deep_resume.py new file mode 100644 index 0000000..f6120c3 --- /dev/null +++ b/tests/test_train_deep_resume.py @@ -0,0 +1,77 @@ +import json +import signal +import sys +import tempfile +import unittest +from pathlib import Path + + +MODULE_DIR = Path(__file__).resolve().parents[1] +sys.path.insert(0, str(MODULE_DIR)) + +from train_deep_mlx import ( + RESUME_SCHEMA_VERSION, + _ShutdownController, + _read_resume_state, + _restore_resume_arguments, + build_parser, +) +from purpose_data import DataError + + +class DeepTrainingResumeTests(unittest.TestCase): + def test_saved_arguments_make_resume_command_self_contained(self): + parser = build_parser() + args = parser.parse_args( + [ + "--resume-training", + "/tmp/purpose-deep/resume", + ] + ) + saved = { + "variant": "base", + "dataset_dir": "/datasets/v2", + "epochs": 7, + "distillation_cache": "/datasets/teacher.pt", + "distillation_weight": 0.5, + "progress_steps": 19, + } + checkpoint = Path("/tmp/purpose-deep/resume") + _restore_resume_arguments( + args, + checkpoint, + {"arguments": saved}, + ) + self.assertEqual(Path("/datasets/v2"), args.dataset_dir) + self.assertEqual(Path("/datasets/teacher.pt"), args.distillation_cache) + self.assertEqual(7, args.epochs) + self.assertEqual(0.5, args.distillation_weight) + self.assertEqual(checkpoint.parent, args.output_dir) + self.assertEqual(checkpoint, args.resume_training) + self.assertIsNone(args.resume_from) + self.assertFalse(args.overwrite_output) + + def test_resume_state_fails_closed_on_wrong_schema(self): + with tempfile.TemporaryDirectory() as temp: + checkpoint = Path(temp) + (checkpoint / "resume-state.json").write_text( + json.dumps( + { + "schemaVersion": RESUME_SCHEMA_VERSION + 1, + "status": "paused", + } + ) + ) + with self.assertRaisesRegex(DataError, "unsupported"): + _read_resume_state(checkpoint) + + def test_second_shutdown_request_is_forceful(self): + controller = _ShutdownController() + controller._handle(signal.SIGINT, None) + self.assertEqual(signal.SIGINT, controller.signum) + with self.assertRaises(KeyboardInterrupt): + controller._handle(signal.SIGINT, None) + + +if __name__ == "__main__": + unittest.main() diff --git a/train_deep_mlx.py b/train_deep_mlx.py index d52b9b2..7b9f8c3 100644 --- a/train_deep_mlx.py +++ b/train_deep_mlx.py @@ -4,11 +4,16 @@ from __future__ import annotations import argparse +import hashlib import json import math +import os import random import shutil +import shlex +import signal import sys +import tempfile import time from collections import Counter from pathlib import Path @@ -43,6 +48,95 @@ from train_mlx import _configure_mlx_device, _linear_schedule, _teacher_cache SCRIPT_DIR = Path(__file__).resolve().parent DEFAULT_DATASET_DIR = SCRIPT_DIR / ".artifacts" / "dataset-v1" DEFAULT_OUTPUT_ROOT = SCRIPT_DIR / "outputs" +RESUME_SCHEMA_VERSION = 1 +PATH_ARGUMENTS = {"dataset_dir", "distillation_cache"} + + +class TrainingPaused(Exception): + """Raised after a signal-requested training checkpoint is durable.""" + + def __init__(self, checkpoint: Path, signum: int) -> None: + super().__init__(str(checkpoint)) + self.checkpoint = checkpoint + self.signum = signum + + +class _ShutdownController: + def __init__(self) -> None: + self.signum: int | None = None + self._previous: dict[int, Any] = {} + + def _handle(self, signum: int, _frame: Any) -> None: + if self.signum is not None: + raise KeyboardInterrupt + self.signum = signum + + def install(self) -> None: + for signum in (signal.SIGINT, signal.SIGTERM): + self._previous[signum] = signal.getsignal(signum) + signal.signal(signum, self._handle) + + def restore(self) -> None: + for signum, handler in self._previous.items(): + signal.signal(signum, handler) + self._previous.clear() + + +def _sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as source: + for chunk in iter(lambda: source.read(1024 * 1024), b""): + digest.update(chunk) + return digest.hexdigest() + + +def _read_resume_state(checkpoint: Path) -> dict[str, Any]: + state_path = checkpoint / "resume-state.json" + try: + state = json.loads(state_path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError) as exc: + raise DataError(f"{state_path}: cannot load training resume state: {exc}") from exc + if state.get("schemaVersion") != RESUME_SCHEMA_VERSION: + raise DataError(f"{state_path}: unsupported training resume schema") + if state.get("status") != "paused": + raise DataError(f"{state_path}: checkpoint is not paused training state") + return state + + +def _serialized_resume_arguments(args: argparse.Namespace) -> dict[str, Any]: + excluded = { + "model", + "resume_from", + "resume_training", + "output_dir", + "overwrite_output", + } + return { + key: ( + str(value.expanduser().resolve()) + if isinstance(value, Path) + else value + ) + for key, value in vars(args).items() + if key not in excluded + } + + +def _restore_resume_arguments( + args: argparse.Namespace, checkpoint: Path, state: dict[str, Any] +) -> None: + saved = state.get("arguments") + if not isinstance(saved, dict): + raise DataError("training resume state has no saved arguments") + for key, value in saved.items(): + if not hasattr(args, key): + raise DataError(f"training resume state has unknown argument {key!r}") + setattr(args, key, Path(value) if key in PATH_ARGUMENTS and value else value) + args.model = None + args.resume_from = None + args.resume_training = checkpoint + args.output_dir = checkpoint.parent + args.overwrite_output = False def _load_mlx(device: str) -> tuple[Any, Any, Any]: @@ -248,6 +342,61 @@ def _save_checkpoint( mx.eval(model.parameters()) +def _save_training_resume( + mx: Any, + model: Any, + optimizer: Any, + tokenizer: Any, + output_dir: Path, + checkpoint_config: dict[str, Any], + state: dict[str, Any], +) -> Path: + """Atomically save the current model, optimizer, and loop cursor.""" + + from mlx.utils import tree_flatten + + output_dir.mkdir(parents=True, exist_ok=True) + temporary = Path( + tempfile.mkdtemp(prefix=".training-resume-", dir=output_dir) + ) + destination = output_dir / "resume" + backup = output_dir / ".training-resume-backup" + try: + _save_checkpoint(mx, model, tokenizer, temporary, checkpoint_config) + mx.eval(optimizer.state) + optimizer_state = tree_flatten(optimizer.state, destination={}) + if not optimizer_state: + raise DataError("optimizer state is empty; refusing an incomplete resume") + mx.save_safetensors( + str(temporary / "optimizer.safetensors"), optimizer_state + ) + write_json(temporary / "resume-state.json", state) + if backup.exists(): + shutil.rmtree(backup) + if destination.exists(): + os.replace(destination, backup) + os.replace(temporary, destination) + if backup.exists(): + shutil.rmtree(backup) + except BaseException: + if temporary.exists(): + shutil.rmtree(temporary) + if not destination.exists() and backup.exists(): + os.replace(backup, destination) + raise + return destination + + +def _load_optimizer_state(mx: Any, optimizer: Any, checkpoint: Path) -> None: + from mlx.utils import tree_unflatten + + path = checkpoint / "optimizer.safetensors" + if not path.is_file(): + raise DataError(f"{path}: optimizer resume state is missing") + optimizer.state = tree_unflatten(mx.load(str(path))) + mx.eval(optimizer.state) + + def _softmax(values: np.ndarray) -> np.ndarray: shifted = values - values.max(axis=-1, keepdims=True) exponentials = np.exp(shifted) @@ -388,12 +537,26 @@ def train(args: argparse.Namespace) -> dict[str, Any]: ) from exc variant = DEEP_VARIANTS[args.variant] - source = _resolve_source(variant, args.resume_from or args.model) - source_config = _load_config(source, variant) - output_dir = args.output_dir or ( - DEFAULT_OUTPUT_ROOT / f"purpose-deep-v1-{variant.name}-mlx" + resume_state = ( + _read_resume_state(args.resume_training) + if args.resume_training is not None + else None ) - _prepare_output(output_dir, source, args.overwrite_output) + source = _resolve_source( + variant, args.resume_training or args.resume_from or args.model + ) + source_config = _load_config(source, variant) + if args.resume_training is not None: + output_dir = args.resume_training.parent + if args.output_dir is not None and args.output_dir.resolve() != output_dir: + raise DataError("--resume-training must use its original output directory") + if not (output_dir / "model" / "model.safetensors").is_file(): + raise DataError(f"{output_dir}: selected model checkpoint is missing") + else: + output_dir = args.output_dir or ( + DEFAULT_OUTPUT_ROOT / f"purpose-deep-v1-{variant.name}-mlx" + ) + _prepare_output(output_dir, source, args.overwrite_output) train_path = args.dataset_dir / "train.jsonl" validation_path = args.dataset_dir / "validation.jsonl" @@ -406,6 +569,20 @@ def train(args: argparse.Namespace) -> dict[str, Any]: if args.max_validation_records: validation_records = validation_records[: args.max_validation_records] + input_hashes = { + "train": _sha256(train_path), + "validation": _sha256(validation_path), + "distillationCache": ( + _sha256(args.distillation_cache.expanduser()) + if args.distillation_cache is not None + else None + ), + } + if resume_state is not None and resume_state.get("inputHashes") != input_hashes: + raise DataError( + "training inputs changed after the pause; refusing a non-deterministic resume" + ) + tokenizer = AutoTokenizer.from_pretrained(source, local_files_only=True) print("tokenizing fixed 1x512 train and validation splits", flush=True) encoded_train = _encode_records(tokenizer, train_records) @@ -438,13 +615,18 @@ def train(args: argparse.Namespace) -> dict[str, Any]: source_config, gradient_checkpointing=not args.no_gradient_checkpointing ) model = ModernBertForPurposeClassification(model_config) - if args.resume_from is not None: + if args.resume_training is not None or args.resume_from is not None: load_report = load_checkpoint_weights( model, source / "model.safetensors" ) print( f"resumed purpose-deep tensors={load_report['loaded']} " - "including all task heads; optimizer state starts fresh", + "including all task heads" + + ( + " and exact optimizer/loop state" + if args.resume_training is not None + else "; optimizer state starts fresh" + ), flush=True, ) else: @@ -473,6 +655,10 @@ def train(args: argparse.Namespace) -> dict[str, Any]: weight_decay=args.weight_decay, bias_correction=True, ) + optimizer.init(model.trainable_parameters()) + mx.eval(optimizer.state) + if args.resume_training is not None: + _load_optimizer_state(mx, optimizer, source) class_weights_mx = mx.array(secondary_class_weights) def loss_function( @@ -554,166 +740,320 @@ def train(args: argparse.Namespace) -> dict[str, Any]: rng = np.random.default_rng(args.seed) checkpoint_config = _checkpoint_config(source_config, variant) best_dir = output_dir / "model" - 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", + output_dir / "training-config.json", { - "bestEpoch": 0, - "bestSelectionScore": best_score, - "elapsedSeconds": time.perf_counter() - started, - "complete": False, + key: str(value) if isinstance(value, Path) else value + for key, value in vars(args).items() }, ) - 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(6, dtype=np.float64) - permutation = rng.permutation(len(train_records)) - for step, indexes in enumerate( - _batch_indexes( - len(train_records), batch_size, permutation=permutation - ), - 1, - ): - batch = _mlx_batch( - mx, encoded_train, train_targets, sample_weights, indexes + if resume_state is not None: + try: + initial_metrics = resume_state["initialMetrics"] + best_score = float(resume_state["bestScore"]) + best_metrics = resume_state["bestMetrics"] + epochs_without_improvement = int( + resume_state["epochsWithoutImprovement"] ) - teacher_logits = ( - mx.array(teacher_train_logits[indexes]) - if teacher_train_logits is not None - else None + history = list(resume_state["history"]) + start_epoch = int(resume_state["epoch"]) + resume_next_step = int(resume_state["nextStep"]) + resume_permutation = resume_state.get("permutation") + resume_running = np.asarray( + resume_state["runningLosses"], dtype=np.float64 ) - losses, gradients = loss_and_grad( - batch["input_ids"], - batch["attention_mask"], - batch["primary"], - batch["secondary"], - batch["secondary_mask"], - batch["mixed"], - batch["difficulty"], - batch["sample_weights"], - teacher_logits, - ) - gradients, _ = optim.clip_grad_norm(gradients, args.max_grad_norm) - optimizer.update(model, gradients) - mx.eval(model.parameters(), optimizer.state, *losses) - running += np.asarray([float(value.item()) for value in losses]) - if args.progress_steps and ( - step % args.progress_steps == 0 or step == steps_per_epoch - ): - mean = running / step - print( - 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" distillation={mean[5]:.4f}" - if teacher_train_logits is not None - else "" - ) - + " " - f"elapsed={time.perf_counter() - epoch_started:.1f}s", - flush=True, - ) - - outputs = _evaluate( + resume_epoch_elapsed = float(resume_state["epochElapsedSeconds"]) + rng.bit_generator.state = resume_state["numpyRngState"] + started = time.perf_counter() - float(resume_state["elapsedSeconds"]) + except (KeyError, TypeError, ValueError) as exc: + raise DataError(f"invalid training loop resume state: {exc}") from exc + if not 1 <= resume_next_step <= steps_per_epoch + 1: + raise DataError("training resume step is outside the epoch") + if resume_running.shape != (6,): + raise DataError("training resume loss accumulator has the wrong shape") + print( + f"continuing epoch {start_epoch} at step " + f"{resume_next_step}/{steps_per_epoch} after " + f"{resume_state['elapsedSeconds']:.1f}s of saved training", + flush=True, + ) + else: + epochs_without_improvement = 0 + history: list[dict[str, Any]] = [] + started = time.perf_counter() + initial_outputs = _evaluate( mx, model, encoded_validation, eval_batch_size ) - metrics = multitask_metrics(outputs, validation_records) + initial_metrics = multitask_metrics(initial_outputs, validation_records) if teacher_validation_logits is not None: - metrics["teacherAgreement"] = float( + initial_metrics["teacherAgreement"] = float( np.mean( - outputs["purpose_logits"][validation_scorable].argmax(axis=-1) + initial_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"]["supportedMacroRecall"] - if metrics["secondary"] is not None - else 0.0 + 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 {epoch}: primary_accuracy={metrics['primary']['accuracy']:.4%} " - f"hard_accuracy={metrics['primaryHardSlice']['accuracy']:.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"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={metrics['teacherAgreement']:.4%}" - if "teacherAgreement" in metrics + f" teacher_agreement={initial_metrics['teacherAgreement']:.4%}" + if "teacherAgreement" in initial_metrics else "" ), flush=True, ) - improvement = score - best_score - if improvement > args.minimum_improvement: - best_score = score - best_metrics = metrics - epochs_without_improvement = 0 - _save_checkpoint( - mx, model, tokenizer, best_dir, checkpoint_config - ) - write_json( - output_dir / "training-state.json", - { - "bestEpoch": epoch, - "bestSelectionScore": best_score, - "elapsedSeconds": time.perf_counter() - started, - "complete": False, - }, - ) - else: - epochs_without_improvement += 1 - if epochs_without_improvement >= args.early_stopping_patience: - stopped_early = True - print( - f"early stopping after epoch {epoch}: no hard-aware " - f"selection improvement greater than " - f"{args.minimum_improvement:.4%} for " - f"{args.early_stopping_patience} epoch(s)", - flush=True, + start_epoch = 1 + resume_next_step = 1 + resume_permutation = None + resume_running = np.zeros(6, dtype=np.float64) + resume_epoch_elapsed = 0.0 + + shutdown = _ShutdownController() + + def pause_training( + epoch: int, + next_step: int, + permutation: np.ndarray | None, + running: np.ndarray, + epoch_elapsed: float, + ) -> None: + signum = shutdown.signum or signal.SIGINT + print( + f"shutdown requested; saving exact training state after epoch {epoch} " + f"step {max(next_step - 1, 0)}", + flush=True, + ) + state = { + "schemaVersion": RESUME_SCHEMA_VERSION, + "status": "paused", + "signal": signal.Signals(signum).name, + "arguments": _serialized_resume_arguments(args), + "inputHashes": input_hashes, + "epoch": epoch, + "nextStep": next_step, + "permutation": permutation.tolist() if permutation is not None else None, + "runningLosses": running.tolist(), + "epochElapsedSeconds": epoch_elapsed, + "elapsedSeconds": time.perf_counter() - started, + "numpyRngState": rng.bit_generator.state, + "initialMetrics": initial_metrics, + "bestScore": best_score, + "bestMetrics": best_metrics, + "epochsWithoutImprovement": epochs_without_improvement, + "history": history, + } + checkpoint = _save_training_resume( + mx, + model, + optimizer, + tokenizer, + output_dir, + checkpoint_config, + state, + ) + write_json( + output_dir / "training-state.json", + { + "bestEpoch": int(best_metrics["epoch"]), + "bestSelectionScore": best_score, + "elapsedSeconds": state["elapsedSeconds"], + "complete": False, + "paused": True, + "resumeCheckpoint": str(checkpoint), + }, + ) + print( + "training paused safely; resume with:\n" + f" {shlex.quote(sys.executable)} -u " + f"{shlex.quote(str(Path(__file__).resolve()))} " + f"--resume-training {shlex.quote(str(checkpoint))}", + flush=True, + ) + raise TrainingPaused(checkpoint, signum) + + shutdown.install() + try: + for epoch in range(start_epoch, args.epochs + 1): + if resume_state is not None and epoch == start_epoch: + permutation = ( + np.asarray(resume_permutation, dtype=np.int64) + if resume_permutation is not None + else rng.permutation(len(train_records)) ) - break + running = resume_running.copy() + first_step = resume_next_step + epoch_started = time.perf_counter() - resume_epoch_elapsed + else: + permutation = rng.permutation(len(train_records)) + running = np.zeros(6, dtype=np.float64) + first_step = 1 + epoch_started = time.perf_counter() + if permutation.shape != (len(train_records),): + raise DataError("training resume permutation has the wrong shape") + if shutdown.signum is not None: + pause_training( + epoch, + first_step, + permutation, + running, + time.perf_counter() - epoch_started, + ) + model.train() + for step, indexes in enumerate( + _batch_indexes( + len(train_records), batch_size, permutation=permutation + ), + 1, + ): + if step < first_step: + continue + 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"], + batch["primary"], + batch["secondary"], + batch["secondary_mask"], + batch["mixed"], + batch["difficulty"], + batch["sample_weights"], + teacher_logits, + ) + gradients, _ = optim.clip_grad_norm(gradients, args.max_grad_norm) + optimizer.update(model, gradients) + mx.eval(model.parameters(), optimizer.state, *losses) + running += np.asarray([float(value.item()) for value in losses]) + if args.progress_steps and ( + step % args.progress_steps == 0 or step == steps_per_epoch + ): + mean = running / step + print( + 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" distillation={mean[5]:.4f}" + if teacher_train_logits is not None + else "" + ) + + " " + f"elapsed={time.perf_counter() - epoch_started:.1f}s", + flush=True, + ) + if shutdown.signum is not None: + pause_training( + epoch, + step + 1, + permutation, + running, + time.perf_counter() - epoch_started, + ) + + outputs = _evaluate( + 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"]["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_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" teacher_agreement={metrics['teacherAgreement']:.4%}" + if "teacherAgreement" in metrics + else "" + ), + flush=True, + ) + improvement = score - best_score + if improvement > args.minimum_improvement: + best_score = score + best_metrics = metrics + epochs_without_improvement = 0 + _save_checkpoint( + mx, model, tokenizer, best_dir, checkpoint_config + ) + write_json( + output_dir / "training-state.json", + { + "bestEpoch": epoch, + "bestSelectionScore": best_score, + "elapsedSeconds": time.perf_counter() - started, + "complete": False, + }, + ) + else: + epochs_without_improvement += 1 + if epochs_without_improvement >= args.early_stopping_patience: + stopped_early = True + print( + f"early stopping after epoch {epoch}: no hard-aware " + f"selection improvement greater than " + f"{args.minimum_improvement:.4%} for " + f"{args.early_stopping_patience} epoch(s)", + flush=True, + ) + break + resume_state = None + if shutdown.signum is not None: + pause_training( + epoch + 1, + 1, + None, + np.zeros(6, dtype=np.float64), + 0.0, + ) + finally: + shutdown.restore() # Release the optimizer graph before opening the selected checkpoint; base and # especially large should never hold two full optimizer states at calibration time. @@ -734,7 +1074,12 @@ def train(args: argparse.Namespace) -> dict[str, Any]: "variant": variant.name, "baseModel": variant.model_id, "baseModelRevision": variant.revision, - "resumedFrom": str(source) if args.resume_from is not None else None, + "resumedFrom": ( + str(source) + if args.resume_from is not None or args.resume_training is not None + else None + ), + "exactTrainingResume": args.resume_training is not None, "parameterClass": variant.parameter_class, "trainingBackend": "mlx", "device": args.device, @@ -800,6 +1145,12 @@ def train(args: argparse.Namespace) -> dict[str, Any]: "complete": True, }, ) + for stale_resume in ( + output_dir / "resume", + output_dir / ".training-resume-backup", + ): + if stale_resume.exists(): + shutil.rmtree(stale_resume) return metrics @@ -820,6 +1171,14 @@ def build_parser() -> argparse.ArgumentParser: "the backbone and all four task heads with a fresh optimizer" ), ) + source_group.add_argument( + "--resume-training", + type=Path, + help=( + "exact signal-created resume checkpoint; restores the saved arguments, " + "model, optimizer, shuffle order, and next batch" + ), + ) parser.add_argument("--dataset-dir", type=Path, default=DEFAULT_DATASET_DIR) parser.add_argument("--output-dir", type=Path) parser.add_argument( @@ -864,6 +1223,13 @@ def _positive(parser: argparse.ArgumentParser, name: str, value: Any) -> None: def main(argv: Sequence[str] | None = None) -> int: parser = build_parser() args = parser.parse_args(argv) + if args.resume_training is not None: + checkpoint = args.resume_training.expanduser().resolve() + try: + resume_state = _read_resume_state(checkpoint) + _restore_resume_arguments(args, checkpoint, resume_state) + except DataError as exc: + parser.error(str(exc)) for name in ( "epochs", "batch_size", @@ -902,6 +1268,8 @@ def main(argv: Sequence[str] | None = None) -> int: ) try: metrics = train(args) + except TrainingPaused as paused: + return 128 + paused.signum except (DataError, OSError, RuntimeError, ValueError) as exc: print(f"error: {exc}", file=sys.stderr) return 1