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 and updates `training-state.json`, so progress is visible and an interrupted run retains
the last selected checkpoint. 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 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 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. 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: def save_weights(model: ModernBertForPurposeClassification, checkpoint: Path) -> None:
mx.eval(model.parameters()) mx.eval(model.parameters())
mx.save_safetensors( mx.save_safetensors(
+35
View File
@@ -13,6 +13,7 @@ sys.path.insert(0, str(MODULE_DIR))
from deep_model_mlx import ( from deep_model_mlx import (
ModernBertForPurposeClassification, ModernBertForPurposeClassification,
ModernBertPurposeConfig, ModernBertPurposeConfig,
load_checkpoint_weights,
load_pretrained_weights, load_pretrained_weights,
save_weights, save_weights,
) )
@@ -114,6 +115,40 @@ class DeepModelTests(unittest.TestCase):
0, float(mx.max(mx.abs(first[key] - second[key])).item()) 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): def test_pretrained_loader_rejects_partial_backbone(self):
with tempfile.TemporaryDirectory() as temp: with tempfile.TemporaryDirectory() as temp:
path = Path(temp) / "partial.safetensors" 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: except ValueError:
pass pass
else: 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 path.exists() and any(path.iterdir()):
if not overwrite: if not overwrite:
raise DataError( raise DataError(
@@ -345,6 +345,7 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
from deep_model_mlx import ( from deep_model_mlx import (
ModernBertForPurposeClassification, ModernBertForPurposeClassification,
ModernBertPurposeConfig, ModernBertPurposeConfig,
load_checkpoint_weights,
load_pretrained_weights, load_pretrained_weights,
) )
except ImportError as exc: except ImportError as exc:
@@ -354,7 +355,7 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
) from exc ) from exc
variant = DEEP_VARIANTS[args.variant] 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) source_config = _load_config(source, variant)
output_dir = args.output_dir or ( output_dir = args.output_dir or (
DEFAULT_OUTPUT_ROOT / f"purpose-deep-v1-{variant.name}-mlx" 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 source_config, gradient_checkpointing=not args.no_gradient_checkpointing
) )
model = ModernBertForPurposeClassification(model_config) 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") load_report = load_pretrained_weights(model, source / "model.safetensors")
print( print(
f"loaded ModernBERT tensors={load_report['loaded']} " f"loaded ModernBERT tensors={load_report['loaded']} "
@@ -599,6 +610,7 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
"variant": variant.name, "variant": variant.name,
"baseModel": variant.model_id, "baseModel": variant.model_id,
"baseModelRevision": variant.revision, "baseModelRevision": variant.revision,
"resumedFrom": str(source) if args.resume_from is not None else None,
"parameterClass": variant.parameter_class, "parameterClass": variant.parameter_class,
"trainingBackend": "mlx", "trainingBackend": "mlx",
"device": args.device, "device": args.device,
@@ -660,11 +672,20 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
def build_parser() -> argparse.ArgumentParser: def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description=__doc__) parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--variant", choices=tuple(DEEP_VARIANTS), default="base") 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", "--model",
type=Path, type=Path,
help="local pinned ModernBERT checkpoint (default: download the pinned revision)", 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("--dataset-dir", type=Path, default=DEFAULT_DATASET_DIR)
parser.add_argument("--output-dir", type=Path) parser.add_argument("--output-dir", type=Path)
parser.add_argument( parser.add_argument(