From 4ed476355717193ac3206ca354a3da17ef0805d4 Mon Sep 17 00:00:00 2001 From: Nucleic Date: Fri, 31 Jul 2026 03:05:29 -0700 Subject: [PATCH] Merge nucleic/sleek-ember-seal-uady into dev --- README.md | 24 +++++++++++++++++++++ deep_model_mlx.py | 41 ++++++++++++++++++++++++++++++++++++ tests/test_deep_model_mlx.py | 35 ++++++++++++++++++++++++++++++ train_deep_mlx.py | 41 +++++++++++++++++++++++++++--------- 4 files changed, 131 insertions(+), 10 deletions(-) diff --git a/README.md b/README.md index 8920876..f09f911 100644 --- a/README.md +++ b/README.md @@ -326,6 +326,30 @@ expected to be extremely slow there. Each improved epoch atomically rewrites `mo and updates `training-state.json`, so progress is visible and an interrupted run retains the last selected checkpoint. +To continue a completed run without discarding its trained task heads, pass its selected +`model/` directory through `--resume-from` and write to a new output directory. +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: + +```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/model \ + --dataset-dir \ + ml/purpose-classifier/.artifacts/dataset-v1-history-first-prompts \ + --epochs 3 \ + --learning-rate 1e-5 \ + --early-stopping-patience 2 \ + --progress-steps 10 \ + --output-dir \ + ml/purpose-classifier/outputs/purpose-deep-v1-base-mlx-cont-3e \ + --overwrite-output +``` + 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_model_mlx.py b/deep_model_mlx.py index cd3d6b3..c0f183a 100644 --- a/deep_model_mlx.py +++ b/deep_model_mlx.py @@ -351,6 +351,47 @@ def load_pretrained_weights( } +def load_checkpoint_weights( + model: ModernBertForPurposeClassification, + checkpoint: Path, +) -> dict[str, int]: + """Strictly restore a trained purpose-deep checkpoint, including task heads.""" + + if not checkpoint.is_file(): + raise DataError(f"{checkpoint}: purpose-deep checkpoint is missing") + weights = mx.load(str(checkpoint)) + parameters = dict(tree_flatten(model.parameters())) + missing = sorted(set(parameters) - set(weights)) + unexpected = sorted(set(weights) - set(parameters)) + if missing or unexpected: + details = [] + if missing: + details.append( + f"missing {len(missing)} tensors ({', '.join(missing[:3])})" + ) + if unexpected: + details.append( + f"has {len(unexpected)} unexpected tensors " + f"({', '.join(unexpected[:3])})" + ) + raise DataError( + f"{checkpoint}: trained checkpoint " + " and ".join(details) + ) + for key, parameter in parameters.items(): + if tuple(weights[key].shape) != tuple(parameter.shape): + raise DataError( + f"purpose-deep tensor {key} has shape {weights[key].shape}; " + f"expected {parameter.shape}" + ) + model.load_weights(list(weights.items()), strict=True) + mx.eval(model.parameters()) + return { + "loaded": len(weights), + "ignored": 0, + "freshTaskHeads": 0, + } + + def save_weights(model: ModernBertForPurposeClassification, checkpoint: Path) -> None: mx.eval(model.parameters()) mx.save_safetensors( diff --git a/tests/test_deep_model_mlx.py b/tests/test_deep_model_mlx.py index fd9278d..e4cabfa 100644 --- a/tests/test_deep_model_mlx.py +++ b/tests/test_deep_model_mlx.py @@ -13,6 +13,7 @@ sys.path.insert(0, str(MODULE_DIR)) from deep_model_mlx import ( ModernBertForPurposeClassification, ModernBertPurposeConfig, + load_checkpoint_weights, load_pretrained_weights, save_weights, ) @@ -114,6 +115,40 @@ class DeepModelTests(unittest.TestCase): 0, float(mx.max(mx.abs(first[key] - second[key])).item()) ) + def test_trained_checkpoint_loader_restores_every_task_head(self): + model = ModernBertForPurposeClassification(tiny_config()) + with tempfile.TemporaryDirectory() as temp: + path = Path(temp) / "model.safetensors" + save_weights(model, path) + restored = ModernBertForPurposeClassification(tiny_config()) + report = load_checkpoint_weights(restored, path) + self.assertEqual(0, report["freshTaskHeads"]) + self.assertEqual(0, report["ignored"]) + ids = mx.array([[1, 3, 4, 2]]) + mask = mx.ones((1, 4), dtype=mx.int32) + first = model(ids, mask) + second = restored(ids, mask) + mx.eval(*first.values(), *second.values()) + for key in first: + with self.subTest(head=key): + self.assertEqual( + 0, float(mx.max(mx.abs(first[key] - second[key])).item()) + ) + + def test_trained_checkpoint_loader_rejects_missing_task_heads(self): + model = ModernBertForPurposeClassification(tiny_config()) + with tempfile.TemporaryDirectory() as temp: + complete = Path(temp) / "complete.safetensors" + partial = Path(temp) / "partial.safetensors" + save_weights(model, complete) + weights = mx.load(str(complete)) + weights.pop("purpose_classifier.weight") + mx.save_safetensors(str(partial), weights) + with self.assertRaisesRegex(DataError, "missing 1 tensors"): + load_checkpoint_weights( + ModernBertForPurposeClassification(tiny_config()), partial + ) + def test_pretrained_loader_rejects_partial_backbone(self): with tempfile.TemporaryDirectory() as temp: path = Path(temp) / "partial.safetensors" diff --git a/train_deep_mlx.py b/train_deep_mlx.py index 602a22d..c4bd3ab 100644 --- a/train_deep_mlx.py +++ b/train_deep_mlx.py @@ -103,7 +103,7 @@ def _prepare_output(path: Path, source: Path, overwrite: bool) -> None: except ValueError: pass else: - raise DataError("--model must not be inside --output-dir") + raise DataError("the input checkpoint must not be inside --output-dir") if path.exists() and any(path.iterdir()): if not overwrite: raise DataError( @@ -345,6 +345,7 @@ def train(args: argparse.Namespace) -> dict[str, Any]: from deep_model_mlx import ( ModernBertForPurposeClassification, ModernBertPurposeConfig, + load_checkpoint_weights, load_pretrained_weights, ) except ImportError as exc: @@ -354,7 +355,7 @@ def train(args: argparse.Namespace) -> dict[str, Any]: ) from exc variant = DEEP_VARIANTS[args.variant] - source = _resolve_source(variant, args.model) + 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" @@ -392,13 +393,23 @@ def train(args: argparse.Namespace) -> dict[str, Any]: source_config, gradient_checkpointing=not args.no_gradient_checkpointing ) model = ModernBertForPurposeClassification(model_config) - load_report = load_pretrained_weights(model, source / "model.safetensors") - print( - f"loaded ModernBERT tensors={load_report['loaded']} " - f"ignored_mlm_tensors={load_report['ignored']} " - f"fresh_task_tensors={load_report['freshTaskHeads']}", - flush=True, - ) + if 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", + flush=True, + ) + else: + load_report = load_pretrained_weights(model, source / "model.safetensors") + print( + f"loaded ModernBERT tensors={load_report['loaded']} " + f"ignored_mlm_tensors={load_report['ignored']} " + f"fresh_task_tensors={load_report['freshTaskHeads']}", + flush=True, + ) batch_size = args.batch_size or (4 if variant.name == "base" else 2) eval_batch_size = args.eval_batch_size or (8 if variant.name == "base" else 4) @@ -599,6 +610,7 @@ 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, "parameterClass": variant.parameter_class, "trainingBackend": "mlx", "device": args.device, @@ -660,11 +672,20 @@ def train(args: argparse.Namespace) -> dict[str, Any]: def build_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--variant", choices=tuple(DEEP_VARIANTS), default="base") - parser.add_argument( + source_group = parser.add_mutually_exclusive_group() + source_group.add_argument( "--model", type=Path, help="local pinned ModernBERT checkpoint (default: download the pinned revision)", ) + source_group.add_argument( + "--resume-from", + type=Path, + help=( + "selected purpose-deep model directory to continue from; restores " + "the backbone and all four task heads with a fresh optimizer" + ), + ) parser.add_argument("--dataset-dir", type=Path, default=DEFAULT_DATASET_DIR) parser.add_argument("--output-dir", type=Path) parser.add_argument(