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
|
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.
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
+31
-10
@@ -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,13 +393,23 @@ 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)
|
||||||
load_report = load_pretrained_weights(model, source / "model.safetensors")
|
if args.resume_from is not None:
|
||||||
print(
|
load_report = load_checkpoint_weights(
|
||||||
f"loaded ModernBERT tensors={load_report['loaded']} "
|
model, source / "model.safetensors"
|
||||||
f"ignored_mlm_tensors={load_report['ignored']} "
|
)
|
||||||
f"fresh_task_tensors={load_report['freshTaskHeads']}",
|
print(
|
||||||
flush=True,
|
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)
|
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)
|
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,
|
"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(
|
||||||
|
|||||||
Reference in New Issue
Block a user