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