Merge nucleic/sleek-ember-seal-uady into dev

This commit is contained in:
2026-08-02 22:44:24 -07:00
parent ae616c6f85
commit 1a41febf73
4 changed files with 673 additions and 146 deletions
+28
View File
@@ -554,6 +554,34 @@ distillation regression cannot overwrite the 80.31% candidate. Teacher agreement
reported for diagnosis but does not enter deep checkpoint selection; overall and hard reported for diagnosis but does not enter deep checkpoint selection; overall and hard
primary label accuracy remain the only selection inputs. primary label accuracy remain the only selection inputs.
### Pause and resume purpose-deep training
`train_deep_mlx.py` handles both `Ctrl-C` (`SIGINT`) and `SIGTERM` gracefully. Press
`Ctrl-C` once: the current optimizer step finishes, then the trainer atomically saves the
current model, AdamW state and learning-rate cursor, epoch and next batch, exact shuffle
order, accumulated losses, RNG state, best-selection state, and history under
`<output-dir>/resume/`. It exits with the conventional signal-derived status only after
that checkpoint is durable. A second `Ctrl-C` forces an immediate interruption.
The pause message prints the complete resume command. For the distilled run above it is:
```bash
ml/purpose-classifier/venv/bin/python -u \
ml/purpose-classifier/train_deep_mlx.py \
--resume-training \
ml/purpose-classifier/outputs/purpose-deep-v1-base-mlx-distilled/resume
```
`--resume-training` is sufficient by itself: it restores the original dataset,
distillation, schedule, loss, and selection arguments from the checkpoint. It also verifies
hashes of the training split, validation split, and teacher cache before loading. The
checkpoint resumes at the next unprocessed batch—even if stopped after the last training
batch but before epoch evaluation. Normal completion removes the large optimizer resume
artifact while retaining the selected `model/` and reports.
This applies to processes started with the updated script. A process already running older
code cannot acquire signal handling retroactively.
The historical large-rung rule required base to qualify first, large to beat base by at The historical large-rung rule required base to qualify first, large to beat base by at
least two hard-slice points, and deep to reach 97% scored overall while beating lite by least two hard-slice points, and deep to reach 97% scored overall while beating lite by
five hard-slice points. Neither the v1 nor Sol-high v2 base result unlocked that rung. five hard-slice points. Neither the v1 nor Sol-high v2 base result unlocked that rung.
+55 -1
View File
@@ -5,6 +5,8 @@ from pathlib import Path
import mlx.core as mx import mlx.core as mx
import mlx.nn as nn import mlx.nn as nn
import mlx.optimizers as optim
from mlx.utils import tree_flatten
MODULE_DIR = Path(__file__).resolve().parents[1] MODULE_DIR = Path(__file__).resolve().parents[1]
@@ -18,7 +20,11 @@ from deep_model_mlx import (
save_weights, save_weights,
) )
from purpose_data import DataError, LABELS from purpose_data import DataError, LABELS
from train_deep_mlx import _distillation_loss from train_deep_mlx import (
_distillation_loss,
_load_optimizer_state,
_save_training_resume,
)
def tiny_config(*, checkpointing=False): def tiny_config(*, checkpointing=False):
@@ -115,6 +121,54 @@ class DeepModelTests(unittest.TestCase):
self.assertAlmostEqual(0.0, float(value.item()), places=6) self.assertAlmostEqual(0.0, float(value.item()), places=6)
self.assertEqual(teacher.shape, gradient.shape) self.assertEqual(teacher.shape, gradient.shape)
def test_training_resume_round_trips_model_and_optimizer(self):
class Tokenizer:
@staticmethod
def save_pretrained(destination):
(destination / "tokenizer_config.json").write_text("{}\n")
model = ModernBertForPurposeClassification(tiny_config())
optimizer = optim.AdamW(learning_rate=1e-3)
optimizer.init(model.trainable_parameters())
def loss(ids, mask):
return mx.mean(model(ids, mask)["purpose_logits"] ** 2)
value_and_grad = nn.value_and_grad(model, loss)
value, gradients = value_and_grad(
mx.array([[1, 3, 4, 2]]), mx.ones((1, 4), dtype=mx.int32)
)
optimizer.update(model, gradients)
mx.eval(value, model.parameters(), optimizer.state)
with tempfile.TemporaryDirectory() as temp:
checkpoint = _save_training_resume(
mx,
model,
optimizer,
Tokenizer(),
Path(temp),
{},
{"schemaVersion": 1, "status": "paused"},
)
restored_model = ModernBertForPurposeClassification(tiny_config())
restored_model.load_weights(
str(checkpoint / "model.safetensors"), strict=True
)
restored_optimizer = optim.AdamW(learning_rate=1e-3)
restored_optimizer.init(restored_model.trainable_parameters())
_load_optimizer_state(mx, restored_optimizer, checkpoint)
original = tree_flatten(optimizer.state, destination={})
restored = tree_flatten(restored_optimizer.state, destination={})
self.assertEqual(set(original), set(restored))
for key in original:
with self.subTest(optimizer_tensor=key):
self.assertEqual(
0,
float(mx.max(mx.abs(original[key] - restored[key])).item()),
)
def test_checkpoint_round_trip(self): def test_checkpoint_round_trip(self):
model = ModernBertForPurposeClassification(tiny_config()) model = ModernBertForPurposeClassification(tiny_config())
with tempfile.TemporaryDirectory() as temp: with tempfile.TemporaryDirectory() as temp:
+77
View File
@@ -0,0 +1,77 @@
import json
import signal
import sys
import tempfile
import unittest
from pathlib import Path
MODULE_DIR = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(MODULE_DIR))
from train_deep_mlx import (
RESUME_SCHEMA_VERSION,
_ShutdownController,
_read_resume_state,
_restore_resume_arguments,
build_parser,
)
from purpose_data import DataError
class DeepTrainingResumeTests(unittest.TestCase):
def test_saved_arguments_make_resume_command_self_contained(self):
parser = build_parser()
args = parser.parse_args(
[
"--resume-training",
"/tmp/purpose-deep/resume",
]
)
saved = {
"variant": "base",
"dataset_dir": "/datasets/v2",
"epochs": 7,
"distillation_cache": "/datasets/teacher.pt",
"distillation_weight": 0.5,
"progress_steps": 19,
}
checkpoint = Path("/tmp/purpose-deep/resume")
_restore_resume_arguments(
args,
checkpoint,
{"arguments": saved},
)
self.assertEqual(Path("/datasets/v2"), args.dataset_dir)
self.assertEqual(Path("/datasets/teacher.pt"), args.distillation_cache)
self.assertEqual(7, args.epochs)
self.assertEqual(0.5, args.distillation_weight)
self.assertEqual(checkpoint.parent, args.output_dir)
self.assertEqual(checkpoint, args.resume_training)
self.assertIsNone(args.resume_from)
self.assertFalse(args.overwrite_output)
def test_resume_state_fails_closed_on_wrong_schema(self):
with tempfile.TemporaryDirectory() as temp:
checkpoint = Path(temp)
(checkpoint / "resume-state.json").write_text(
json.dumps(
{
"schemaVersion": RESUME_SCHEMA_VERSION + 1,
"status": "paused",
}
)
)
with self.assertRaisesRegex(DataError, "unsupported"):
_read_resume_state(checkpoint)
def test_second_shutdown_request_is_forceful(self):
controller = _ShutdownController()
controller._handle(signal.SIGINT, None)
self.assertEqual(signal.SIGINT, controller.signum)
with self.assertRaises(KeyboardInterrupt):
controller._handle(signal.SIGINT, None)
if __name__ == "__main__":
unittest.main()
+380 -12
View File
@@ -4,11 +4,16 @@
from __future__ import annotations from __future__ import annotations
import argparse import argparse
import hashlib
import json import json
import math import math
import os
import random import random
import shutil import shutil
import shlex
import signal
import sys import sys
import tempfile
import time import time
from collections import Counter from collections import Counter
from pathlib import Path from pathlib import Path
@@ -43,6 +48,95 @@ from train_mlx import _configure_mlx_device, _linear_schedule, _teacher_cache
SCRIPT_DIR = Path(__file__).resolve().parent SCRIPT_DIR = Path(__file__).resolve().parent
DEFAULT_DATASET_DIR = SCRIPT_DIR / ".artifacts" / "dataset-v1" DEFAULT_DATASET_DIR = SCRIPT_DIR / ".artifacts" / "dataset-v1"
DEFAULT_OUTPUT_ROOT = SCRIPT_DIR / "outputs" DEFAULT_OUTPUT_ROOT = SCRIPT_DIR / "outputs"
RESUME_SCHEMA_VERSION = 1
PATH_ARGUMENTS = {"dataset_dir", "distillation_cache"}
class TrainingPaused(Exception):
"""Raised after a signal-requested training checkpoint is durable."""
def __init__(self, checkpoint: Path, signum: int) -> None:
super().__init__(str(checkpoint))
self.checkpoint = checkpoint
self.signum = signum
class _ShutdownController:
def __init__(self) -> None:
self.signum: int | None = None
self._previous: dict[int, Any] = {}
def _handle(self, signum: int, _frame: Any) -> None:
if self.signum is not None:
raise KeyboardInterrupt
self.signum = signum
def install(self) -> None:
for signum in (signal.SIGINT, signal.SIGTERM):
self._previous[signum] = signal.getsignal(signum)
signal.signal(signum, self._handle)
def restore(self) -> None:
for signum, handler in self._previous.items():
signal.signal(signum, handler)
self._previous.clear()
def _sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as source:
for chunk in iter(lambda: source.read(1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
def _read_resume_state(checkpoint: Path) -> dict[str, Any]:
state_path = checkpoint / "resume-state.json"
try:
state = json.loads(state_path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError) as exc:
raise DataError(f"{state_path}: cannot load training resume state: {exc}") from exc
if state.get("schemaVersion") != RESUME_SCHEMA_VERSION:
raise DataError(f"{state_path}: unsupported training resume schema")
if state.get("status") != "paused":
raise DataError(f"{state_path}: checkpoint is not paused training state")
return state
def _serialized_resume_arguments(args: argparse.Namespace) -> dict[str, Any]:
excluded = {
"model",
"resume_from",
"resume_training",
"output_dir",
"overwrite_output",
}
return {
key: (
str(value.expanduser().resolve())
if isinstance(value, Path)
else value
)
for key, value in vars(args).items()
if key not in excluded
}
def _restore_resume_arguments(
args: argparse.Namespace, checkpoint: Path, state: dict[str, Any]
) -> None:
saved = state.get("arguments")
if not isinstance(saved, dict):
raise DataError("training resume state has no saved arguments")
for key, value in saved.items():
if not hasattr(args, key):
raise DataError(f"training resume state has unknown argument {key!r}")
setattr(args, key, Path(value) if key in PATH_ARGUMENTS and value else value)
args.model = None
args.resume_from = None
args.resume_training = checkpoint
args.output_dir = checkpoint.parent
args.overwrite_output = False
def _load_mlx(device: str) -> tuple[Any, Any, Any]: def _load_mlx(device: str) -> tuple[Any, Any, Any]:
@@ -248,6 +342,61 @@ def _save_checkpoint(
mx.eval(model.parameters()) mx.eval(model.parameters())
def _save_training_resume(
mx: Any,
model: Any,
optimizer: Any,
tokenizer: Any,
output_dir: Path,
checkpoint_config: dict[str, Any],
state: dict[str, Any],
) -> Path:
"""Atomically save the current model, optimizer, and loop cursor."""
from mlx.utils import tree_flatten
output_dir.mkdir(parents=True, exist_ok=True)
temporary = Path(
tempfile.mkdtemp(prefix=".training-resume-", dir=output_dir)
)
destination = output_dir / "resume"
backup = output_dir / ".training-resume-backup"
try:
_save_checkpoint(mx, model, tokenizer, temporary, checkpoint_config)
mx.eval(optimizer.state)
optimizer_state = tree_flatten(optimizer.state, destination={})
if not optimizer_state:
raise DataError("optimizer state is empty; refusing an incomplete resume")
mx.save_safetensors(
str(temporary / "optimizer.safetensors"), optimizer_state
)
write_json(temporary / "resume-state.json", state)
if backup.exists():
shutil.rmtree(backup)
if destination.exists():
os.replace(destination, backup)
os.replace(temporary, destination)
if backup.exists():
shutil.rmtree(backup)
except BaseException:
if temporary.exists():
shutil.rmtree(temporary)
if not destination.exists() and backup.exists():
os.replace(backup, destination)
raise
return destination
def _load_optimizer_state(mx: Any, optimizer: Any, checkpoint: Path) -> None:
from mlx.utils import tree_unflatten
path = checkpoint / "optimizer.safetensors"
if not path.is_file():
raise DataError(f"{path}: optimizer resume state is missing")
optimizer.state = tree_unflatten(mx.load(str(path)))
mx.eval(optimizer.state)
def _softmax(values: np.ndarray) -> np.ndarray: def _softmax(values: np.ndarray) -> np.ndarray:
shifted = values - values.max(axis=-1, keepdims=True) shifted = values - values.max(axis=-1, keepdims=True)
exponentials = np.exp(shifted) exponentials = np.exp(shifted)
@@ -388,8 +537,22 @@ 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.resume_from or args.model) resume_state = (
_read_resume_state(args.resume_training)
if args.resume_training is not None
else None
)
source = _resolve_source(
variant, args.resume_training or args.resume_from or args.model
)
source_config = _load_config(source, variant) source_config = _load_config(source, variant)
if args.resume_training is not None:
output_dir = args.resume_training.parent
if args.output_dir is not None and args.output_dir.resolve() != output_dir:
raise DataError("--resume-training must use its original output directory")
if not (output_dir / "model" / "model.safetensors").is_file():
raise DataError(f"{output_dir}: selected model checkpoint is missing")
else:
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"
) )
@@ -406,6 +569,20 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
if args.max_validation_records: if args.max_validation_records:
validation_records = validation_records[: args.max_validation_records] validation_records = validation_records[: args.max_validation_records]
input_hashes = {
"train": _sha256(train_path),
"validation": _sha256(validation_path),
"distillationCache": (
_sha256(args.distillation_cache.expanduser())
if args.distillation_cache is not None
else None
),
}
if resume_state is not None and resume_state.get("inputHashes") != input_hashes:
raise DataError(
"training inputs changed after the pause; refusing a non-deterministic resume"
)
tokenizer = AutoTokenizer.from_pretrained(source, local_files_only=True) tokenizer = AutoTokenizer.from_pretrained(source, local_files_only=True)
print("tokenizing fixed 1x512 train and validation splits", flush=True) print("tokenizing fixed 1x512 train and validation splits", flush=True)
encoded_train = _encode_records(tokenizer, train_records) encoded_train = _encode_records(tokenizer, train_records)
@@ -438,13 +615,18 @@ 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: if args.resume_training is not None or args.resume_from is not None:
load_report = load_checkpoint_weights( load_report = load_checkpoint_weights(
model, source / "model.safetensors" model, source / "model.safetensors"
) )
print( print(
f"resumed purpose-deep tensors={load_report['loaded']} " f"resumed purpose-deep tensors={load_report['loaded']} "
"including all task heads; optimizer state starts fresh", "including all task heads"
+ (
" and exact optimizer/loop state"
if args.resume_training is not None
else "; optimizer state starts fresh"
),
flush=True, flush=True,
) )
else: else:
@@ -473,6 +655,10 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
weight_decay=args.weight_decay, weight_decay=args.weight_decay,
bias_correction=True, bias_correction=True,
) )
optimizer.init(model.trainable_parameters())
mx.eval(optimizer.state)
if args.resume_training is not None:
_load_optimizer_state(mx, optimizer, source)
class_weights_mx = mx.array(secondary_class_weights) class_weights_mx = mx.array(secondary_class_weights)
def loss_function( def loss_function(
@@ -554,11 +740,49 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
rng = np.random.default_rng(args.seed) rng = np.random.default_rng(args.seed)
checkpoint_config = _checkpoint_config(source_config, variant) checkpoint_config = _checkpoint_config(source_config, variant)
best_dir = output_dir / "model" best_dir = output_dir / "model"
epochs_without_improvement = 0
stopped_early = False stopped_early = False
write_json(
output_dir / "training-config.json",
{
key: str(value) if isinstance(value, Path) else value
for key, value in vars(args).items()
},
)
if resume_state is not None:
try:
initial_metrics = resume_state["initialMetrics"]
best_score = float(resume_state["bestScore"])
best_metrics = resume_state["bestMetrics"]
epochs_without_improvement = int(
resume_state["epochsWithoutImprovement"]
)
history = list(resume_state["history"])
start_epoch = int(resume_state["epoch"])
resume_next_step = int(resume_state["nextStep"])
resume_permutation = resume_state.get("permutation")
resume_running = np.asarray(
resume_state["runningLosses"], dtype=np.float64
)
resume_epoch_elapsed = float(resume_state["epochElapsedSeconds"])
rng.bit_generator.state = resume_state["numpyRngState"]
started = time.perf_counter() - float(resume_state["elapsedSeconds"])
except (KeyError, TypeError, ValueError) as exc:
raise DataError(f"invalid training loop resume state: {exc}") from exc
if not 1 <= resume_next_step <= steps_per_epoch + 1:
raise DataError("training resume step is outside the epoch")
if resume_running.shape != (6,):
raise DataError("training resume loss accumulator has the wrong shape")
print(
f"continuing epoch {start_epoch} at step "
f"{resume_next_step}/{steps_per_epoch} after "
f"{resume_state['elapsedSeconds']:.1f}s of saved training",
flush=True,
)
else:
epochs_without_improvement = 0
history: list[dict[str, Any]] = [] history: list[dict[str, Any]] = []
started = time.perf_counter() started = time.perf_counter()
initial_outputs = _evaluate( initial_outputs = _evaluate(
mx, model, encoded_validation, eval_batch_size mx, model, encoded_validation, eval_batch_size
) )
@@ -597,18 +821,111 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
), ),
flush=True, flush=True,
) )
start_epoch = 1
resume_next_step = 1
resume_permutation = None
resume_running = np.zeros(6, dtype=np.float64)
resume_epoch_elapsed = 0.0
for epoch in range(1, args.epochs + 1): shutdown = _ShutdownController()
epoch_started = time.perf_counter()
model.train() def pause_training(
running = np.zeros(6, dtype=np.float64) epoch: int,
next_step: int,
permutation: np.ndarray | None,
running: np.ndarray,
epoch_elapsed: float,
) -> None:
signum = shutdown.signum or signal.SIGINT
print(
f"shutdown requested; saving exact training state after epoch {epoch} "
f"step {max(next_step - 1, 0)}",
flush=True,
)
state = {
"schemaVersion": RESUME_SCHEMA_VERSION,
"status": "paused",
"signal": signal.Signals(signum).name,
"arguments": _serialized_resume_arguments(args),
"inputHashes": input_hashes,
"epoch": epoch,
"nextStep": next_step,
"permutation": permutation.tolist() if permutation is not None else None,
"runningLosses": running.tolist(),
"epochElapsedSeconds": epoch_elapsed,
"elapsedSeconds": time.perf_counter() - started,
"numpyRngState": rng.bit_generator.state,
"initialMetrics": initial_metrics,
"bestScore": best_score,
"bestMetrics": best_metrics,
"epochsWithoutImprovement": epochs_without_improvement,
"history": history,
}
checkpoint = _save_training_resume(
mx,
model,
optimizer,
tokenizer,
output_dir,
checkpoint_config,
state,
)
write_json(
output_dir / "training-state.json",
{
"bestEpoch": int(best_metrics["epoch"]),
"bestSelectionScore": best_score,
"elapsedSeconds": state["elapsedSeconds"],
"complete": False,
"paused": True,
"resumeCheckpoint": str(checkpoint),
},
)
print(
"training paused safely; resume with:\n"
f" {shlex.quote(sys.executable)} -u "
f"{shlex.quote(str(Path(__file__).resolve()))} "
f"--resume-training {shlex.quote(str(checkpoint))}",
flush=True,
)
raise TrainingPaused(checkpoint, signum)
shutdown.install()
try:
for epoch in range(start_epoch, args.epochs + 1):
if resume_state is not None and epoch == start_epoch:
permutation = (
np.asarray(resume_permutation, dtype=np.int64)
if resume_permutation is not None
else rng.permutation(len(train_records))
)
running = resume_running.copy()
first_step = resume_next_step
epoch_started = time.perf_counter() - resume_epoch_elapsed
else:
permutation = rng.permutation(len(train_records)) permutation = rng.permutation(len(train_records))
running = np.zeros(6, dtype=np.float64)
first_step = 1
epoch_started = time.perf_counter()
if permutation.shape != (len(train_records),):
raise DataError("training resume permutation has the wrong shape")
if shutdown.signum is not None:
pause_training(
epoch,
first_step,
permutation,
running,
time.perf_counter() - epoch_started,
)
model.train()
for step, indexes in enumerate( for step, indexes in enumerate(
_batch_indexes( _batch_indexes(
len(train_records), batch_size, permutation=permutation len(train_records), batch_size, permutation=permutation
), ),
1, 1,
): ):
if step < first_step:
continue
batch = _mlx_batch( batch = _mlx_batch(
mx, encoded_train, train_targets, sample_weights, indexes mx, encoded_train, train_targets, sample_weights, indexes
) )
@@ -650,6 +967,14 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
f"elapsed={time.perf_counter() - epoch_started:.1f}s", f"elapsed={time.perf_counter() - epoch_started:.1f}s",
flush=True, flush=True,
) )
if shutdown.signum is not None:
pause_training(
epoch,
step + 1,
permutation,
running,
time.perf_counter() - epoch_started,
)
outputs = _evaluate( outputs = _evaluate(
mx, model, encoded_validation, eval_batch_size mx, model, encoded_validation, eval_batch_size
@@ -658,8 +983,12 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
if teacher_validation_logits is not None: if teacher_validation_logits is not None:
metrics["teacherAgreement"] = float( metrics["teacherAgreement"] = float(
np.mean( np.mean(
outputs["purpose_logits"][validation_scorable].argmax(axis=-1) outputs["purpose_logits"][validation_scorable].argmax(
== teacher_validation_logits[validation_scorable].argmax(axis=-1) axis=-1
)
== teacher_validation_logits[validation_scorable].argmax(
axis=-1
)
) )
) )
metrics["epoch"] = epoch metrics["epoch"] = epoch
@@ -714,6 +1043,17 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
flush=True, flush=True,
) )
break break
resume_state = None
if shutdown.signum is not None:
pause_training(
epoch + 1,
1,
None,
np.zeros(6, dtype=np.float64),
0.0,
)
finally:
shutdown.restore()
# Release the optimizer graph before opening the selected checkpoint; base and # Release the optimizer graph before opening the selected checkpoint; base and
# especially large should never hold two full optimizer states at calibration time. # especially large should never hold two full optimizer states at calibration time.
@@ -734,7 +1074,12 @@ 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, "resumedFrom": (
str(source)
if args.resume_from is not None or args.resume_training is not None
else None
),
"exactTrainingResume": args.resume_training is not None,
"parameterClass": variant.parameter_class, "parameterClass": variant.parameter_class,
"trainingBackend": "mlx", "trainingBackend": "mlx",
"device": args.device, "device": args.device,
@@ -800,6 +1145,12 @@ def train(args: argparse.Namespace) -> dict[str, Any]:
"complete": True, "complete": True,
}, },
) )
for stale_resume in (
output_dir / "resume",
output_dir / ".training-resume-backup",
):
if stale_resume.exists():
shutil.rmtree(stale_resume)
return metrics return metrics
@@ -820,6 +1171,14 @@ def build_parser() -> argparse.ArgumentParser:
"the backbone and all four task heads with a fresh optimizer" "the backbone and all four task heads with a fresh optimizer"
), ),
) )
source_group.add_argument(
"--resume-training",
type=Path,
help=(
"exact signal-created resume checkpoint; restores the saved arguments, "
"model, optimizer, shuffle order, and next batch"
),
)
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(
@@ -864,6 +1223,13 @@ def _positive(parser: argparse.ArgumentParser, name: str, value: Any) -> None:
def main(argv: Sequence[str] | None = None) -> int: def main(argv: Sequence[str] | None = None) -> int:
parser = build_parser() parser = build_parser()
args = parser.parse_args(argv) args = parser.parse_args(argv)
if args.resume_training is not None:
checkpoint = args.resume_training.expanduser().resolve()
try:
resume_state = _read_resume_state(checkpoint)
_restore_resume_arguments(args, checkpoint, resume_state)
except DataError as exc:
parser.error(str(exc))
for name in ( for name in (
"epochs", "epochs",
"batch_size", "batch_size",
@@ -902,6 +1268,8 @@ def main(argv: Sequence[str] | None = None) -> int:
) )
try: try:
metrics = train(args) metrics = train(args)
except TrainingPaused as paused:
return 128 + paused.signum
except (DataError, OSError, RuntimeError, ValueError) as exc: except (DataError, OSError, RuntimeError, ValueError) as exc:
print(f"error: {exc}", file=sys.stderr) print(f"error: {exc}", file=sys.stderr)
return 1 return 1