Merge nucleic/sleek-ember-seal-uady into dev

This commit is contained in:
2026-07-31 03:05:29 -07:00
parent 9a1228efbb
commit 4ed4763557
4 changed files with 131 additions and 10 deletions
+24
View File
@@ -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.
+41
View File
@@ -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(
+35
View File
@@ -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"
+24 -3
View File
@@ -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,6 +393,16 @@ 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:
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']} "
@@ -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(