Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
+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