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
+31 -10
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,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(