Merge nucleic/sleek-ember-seal-uady into dev
This commit is contained in:
@@ -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.
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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()
|
||||||
+513
-145
@@ -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,12 +537,26 @@ 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 = (
|
||||||
source_config = _load_config(source, variant)
|
_read_resume_state(args.resume_training)
|
||||||
output_dir = args.output_dir or (
|
if args.resume_training is not None
|
||||||
DEFAULT_OUTPUT_ROOT / f"purpose-deep-v1-{variant.name}-mlx"
|
else None
|
||||||
)
|
)
|
||||||
_prepare_output(output_dir, source, args.overwrite_output)
|
source = _resolve_source(
|
||||||
|
variant, args.resume_training or args.resume_from or args.model
|
||||||
|
)
|
||||||
|
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 (
|
||||||
|
DEFAULT_OUTPUT_ROOT / f"purpose-deep-v1-{variant.name}-mlx"
|
||||||
|
)
|
||||||
|
_prepare_output(output_dir, source, args.overwrite_output)
|
||||||
|
|
||||||
train_path = args.dataset_dir / "train.jsonl"
|
train_path = args.dataset_dir / "train.jsonl"
|
||||||
validation_path = args.dataset_dir / "validation.jsonl"
|
validation_path = args.dataset_dir / "validation.jsonl"
|
||||||
@@ -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,166 +740,320 @@ 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
|
||||||
history: list[dict[str, Any]] = []
|
|
||||||
started = time.perf_counter()
|
|
||||||
|
|
||||||
initial_outputs = _evaluate(
|
|
||||||
mx, model, encoded_validation, eval_batch_size
|
|
||||||
)
|
|
||||||
initial_metrics = multitask_metrics(initial_outputs, validation_records)
|
|
||||||
if teacher_validation_logits is not None:
|
|
||||||
initial_metrics["teacherAgreement"] = float(
|
|
||||||
np.mean(
|
|
||||||
initial_outputs["purpose_logits"][validation_scorable].argmax(
|
|
||||||
axis=-1
|
|
||||||
)
|
|
||||||
== teacher_validation_logits[validation_scorable].argmax(axis=-1)
|
|
||||||
)
|
|
||||||
)
|
|
||||||
initial_metrics["epoch"] = 0
|
|
||||||
best_score = float(initial_metrics["selectionScore"])
|
|
||||||
best_metrics: dict[str, Any] = initial_metrics
|
|
||||||
_save_checkpoint(mx, model, tokenizer, best_dir, checkpoint_config)
|
|
||||||
write_json(
|
write_json(
|
||||||
output_dir / "training-state.json",
|
output_dir / "training-config.json",
|
||||||
{
|
{
|
||||||
"bestEpoch": 0,
|
key: str(value) if isinstance(value, Path) else value
|
||||||
"bestSelectionScore": best_score,
|
for key, value in vars(args).items()
|
||||||
"elapsedSeconds": time.perf_counter() - started,
|
|
||||||
"complete": False,
|
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
print(
|
|
||||||
f"epoch 0: primary_accuracy={initial_metrics['primary']['accuracy']:.4%} "
|
|
||||||
f"hard_accuracy={initial_metrics['primaryHardSlice']['accuracy']:.4%} "
|
|
||||||
f"mixed_f1={initial_metrics['mixed']['f1']:.4%} "
|
|
||||||
f"selection_score={best_score:.4%}"
|
|
||||||
+ (
|
|
||||||
f" teacher_agreement={initial_metrics['teacherAgreement']:.4%}"
|
|
||||||
if "teacherAgreement" in initial_metrics
|
|
||||||
else ""
|
|
||||||
),
|
|
||||||
flush=True,
|
|
||||||
)
|
|
||||||
|
|
||||||
for epoch in range(1, args.epochs + 1):
|
if resume_state is not None:
|
||||||
epoch_started = time.perf_counter()
|
try:
|
||||||
model.train()
|
initial_metrics = resume_state["initialMetrics"]
|
||||||
running = np.zeros(6, dtype=np.float64)
|
best_score = float(resume_state["bestScore"])
|
||||||
permutation = rng.permutation(len(train_records))
|
best_metrics = resume_state["bestMetrics"]
|
||||||
for step, indexes in enumerate(
|
epochs_without_improvement = int(
|
||||||
_batch_indexes(
|
resume_state["epochsWithoutImprovement"]
|
||||||
len(train_records), batch_size, permutation=permutation
|
|
||||||
),
|
|
||||||
1,
|
|
||||||
):
|
|
||||||
batch = _mlx_batch(
|
|
||||||
mx, encoded_train, train_targets, sample_weights, indexes
|
|
||||||
)
|
)
|
||||||
teacher_logits = (
|
history = list(resume_state["history"])
|
||||||
mx.array(teacher_train_logits[indexes])
|
start_epoch = int(resume_state["epoch"])
|
||||||
if teacher_train_logits is not None
|
resume_next_step = int(resume_state["nextStep"])
|
||||||
else None
|
resume_permutation = resume_state.get("permutation")
|
||||||
|
resume_running = np.asarray(
|
||||||
|
resume_state["runningLosses"], dtype=np.float64
|
||||||
)
|
)
|
||||||
losses, gradients = loss_and_grad(
|
resume_epoch_elapsed = float(resume_state["epochElapsedSeconds"])
|
||||||
batch["input_ids"],
|
rng.bit_generator.state = resume_state["numpyRngState"]
|
||||||
batch["attention_mask"],
|
started = time.perf_counter() - float(resume_state["elapsedSeconds"])
|
||||||
batch["primary"],
|
except (KeyError, TypeError, ValueError) as exc:
|
||||||
batch["secondary"],
|
raise DataError(f"invalid training loop resume state: {exc}") from exc
|
||||||
batch["secondary_mask"],
|
if not 1 <= resume_next_step <= steps_per_epoch + 1:
|
||||||
batch["mixed"],
|
raise DataError("training resume step is outside the epoch")
|
||||||
batch["difficulty"],
|
if resume_running.shape != (6,):
|
||||||
batch["sample_weights"],
|
raise DataError("training resume loss accumulator has the wrong shape")
|
||||||
teacher_logits,
|
print(
|
||||||
)
|
f"continuing epoch {start_epoch} at step "
|
||||||
gradients, _ = optim.clip_grad_norm(gradients, args.max_grad_norm)
|
f"{resume_next_step}/{steps_per_epoch} after "
|
||||||
optimizer.update(model, gradients)
|
f"{resume_state['elapsedSeconds']:.1f}s of saved training",
|
||||||
mx.eval(model.parameters(), optimizer.state, *losses)
|
flush=True,
|
||||||
running += np.asarray([float(value.item()) for value in losses])
|
)
|
||||||
if args.progress_steps and (
|
else:
|
||||||
step % args.progress_steps == 0 or step == steps_per_epoch
|
epochs_without_improvement = 0
|
||||||
):
|
history: list[dict[str, Any]] = []
|
||||||
mean = running / step
|
started = time.perf_counter()
|
||||||
print(
|
initial_outputs = _evaluate(
|
||||||
f"epoch {epoch} step {step}/{steps_per_epoch} "
|
|
||||||
f"loss={mean[0]:.4f} primary={mean[1]:.4f} "
|
|
||||||
f"secondary={mean[2]:.4f} mixed={mean[3]:.4f} "
|
|
||||||
f"difficulty={mean[4]:.4f}"
|
|
||||||
+ (
|
|
||||||
f" distillation={mean[5]:.4f}"
|
|
||||||
if teacher_train_logits is not None
|
|
||||||
else ""
|
|
||||||
)
|
|
||||||
+ " "
|
|
||||||
f"elapsed={time.perf_counter() - epoch_started:.1f}s",
|
|
||||||
flush=True,
|
|
||||||
)
|
|
||||||
|
|
||||||
outputs = _evaluate(
|
|
||||||
mx, model, encoded_validation, eval_batch_size
|
mx, model, encoded_validation, eval_batch_size
|
||||||
)
|
)
|
||||||
metrics = multitask_metrics(outputs, validation_records)
|
initial_metrics = multitask_metrics(initial_outputs, validation_records)
|
||||||
if teacher_validation_logits is not None:
|
if teacher_validation_logits is not None:
|
||||||
metrics["teacherAgreement"] = float(
|
initial_metrics["teacherAgreement"] = float(
|
||||||
np.mean(
|
np.mean(
|
||||||
outputs["purpose_logits"][validation_scorable].argmax(axis=-1)
|
initial_outputs["purpose_logits"][validation_scorable].argmax(
|
||||||
|
axis=-1
|
||||||
|
)
|
||||||
== teacher_validation_logits[validation_scorable].argmax(axis=-1)
|
== teacher_validation_logits[validation_scorable].argmax(axis=-1)
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
metrics["epoch"] = epoch
|
initial_metrics["epoch"] = 0
|
||||||
metrics["meanTrainingLoss"] = (running / steps_per_epoch).tolist()
|
best_score = float(initial_metrics["selectionScore"])
|
||||||
history.append(metrics)
|
best_metrics: dict[str, Any] = initial_metrics
|
||||||
score = float(metrics["selectionScore"])
|
_save_checkpoint(mx, model, tokenizer, best_dir, checkpoint_config)
|
||||||
secondary_macro = (
|
write_json(
|
||||||
metrics["secondary"]["supportedMacroRecall"]
|
output_dir / "training-state.json",
|
||||||
if metrics["secondary"] is not None
|
{
|
||||||
else 0.0
|
"bestEpoch": 0,
|
||||||
|
"bestSelectionScore": best_score,
|
||||||
|
"elapsedSeconds": time.perf_counter() - started,
|
||||||
|
"complete": False,
|
||||||
|
},
|
||||||
)
|
)
|
||||||
print(
|
print(
|
||||||
f"epoch {epoch}: primary_accuracy={metrics['primary']['accuracy']:.4%} "
|
f"epoch 0: primary_accuracy={initial_metrics['primary']['accuracy']:.4%} "
|
||||||
f"hard_accuracy={metrics['primaryHardSlice']['accuracy']:.4%} "
|
f"hard_accuracy={initial_metrics['primaryHardSlice']['accuracy']:.4%} "
|
||||||
f"secondary_supported_macro_recall={secondary_macro:.4%} "
|
f"mixed_f1={initial_metrics['mixed']['f1']:.4%} "
|
||||||
f"mixed_f1={metrics['mixed']['f1']:.4%} "
|
f"selection_score={best_score:.4%}"
|
||||||
f"difficulty_mae={metrics['difficulty']['mae']:.4f} "
|
|
||||||
f"selection_score={score:.4%}"
|
|
||||||
+ (
|
+ (
|
||||||
f" teacher_agreement={metrics['teacherAgreement']:.4%}"
|
f" teacher_agreement={initial_metrics['teacherAgreement']:.4%}"
|
||||||
if "teacherAgreement" in metrics
|
if "teacherAgreement" in initial_metrics
|
||||||
else ""
|
else ""
|
||||||
),
|
),
|
||||||
flush=True,
|
flush=True,
|
||||||
)
|
)
|
||||||
improvement = score - best_score
|
start_epoch = 1
|
||||||
if improvement > args.minimum_improvement:
|
resume_next_step = 1
|
||||||
best_score = score
|
resume_permutation = None
|
||||||
best_metrics = metrics
|
resume_running = np.zeros(6, dtype=np.float64)
|
||||||
epochs_without_improvement = 0
|
resume_epoch_elapsed = 0.0
|
||||||
_save_checkpoint(
|
|
||||||
mx, model, tokenizer, best_dir, checkpoint_config
|
shutdown = _ShutdownController()
|
||||||
)
|
|
||||||
write_json(
|
def pause_training(
|
||||||
output_dir / "training-state.json",
|
epoch: int,
|
||||||
{
|
next_step: int,
|
||||||
"bestEpoch": epoch,
|
permutation: np.ndarray | None,
|
||||||
"bestSelectionScore": best_score,
|
running: np.ndarray,
|
||||||
"elapsedSeconds": time.perf_counter() - started,
|
epoch_elapsed: float,
|
||||||
"complete": False,
|
) -> None:
|
||||||
},
|
signum = shutdown.signum or signal.SIGINT
|
||||||
)
|
print(
|
||||||
else:
|
f"shutdown requested; saving exact training state after epoch {epoch} "
|
||||||
epochs_without_improvement += 1
|
f"step {max(next_step - 1, 0)}",
|
||||||
if epochs_without_improvement >= args.early_stopping_patience:
|
flush=True,
|
||||||
stopped_early = True
|
)
|
||||||
print(
|
state = {
|
||||||
f"early stopping after epoch {epoch}: no hard-aware "
|
"schemaVersion": RESUME_SCHEMA_VERSION,
|
||||||
f"selection improvement greater than "
|
"status": "paused",
|
||||||
f"{args.minimum_improvement:.4%} for "
|
"signal": signal.Signals(signum).name,
|
||||||
f"{args.early_stopping_patience} epoch(s)",
|
"arguments": _serialized_resume_arguments(args),
|
||||||
flush=True,
|
"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))
|
||||||
)
|
)
|
||||||
break
|
running = resume_running.copy()
|
||||||
|
first_step = resume_next_step
|
||||||
|
epoch_started = time.perf_counter() - resume_epoch_elapsed
|
||||||
|
else:
|
||||||
|
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(
|
||||||
|
_batch_indexes(
|
||||||
|
len(train_records), batch_size, permutation=permutation
|
||||||
|
),
|
||||||
|
1,
|
||||||
|
):
|
||||||
|
if step < first_step:
|
||||||
|
continue
|
||||||
|
batch = _mlx_batch(
|
||||||
|
mx, encoded_train, train_targets, sample_weights, indexes
|
||||||
|
)
|
||||||
|
teacher_logits = (
|
||||||
|
mx.array(teacher_train_logits[indexes])
|
||||||
|
if teacher_train_logits is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
losses, gradients = loss_and_grad(
|
||||||
|
batch["input_ids"],
|
||||||
|
batch["attention_mask"],
|
||||||
|
batch["primary"],
|
||||||
|
batch["secondary"],
|
||||||
|
batch["secondary_mask"],
|
||||||
|
batch["mixed"],
|
||||||
|
batch["difficulty"],
|
||||||
|
batch["sample_weights"],
|
||||||
|
teacher_logits,
|
||||||
|
)
|
||||||
|
gradients, _ = optim.clip_grad_norm(gradients, args.max_grad_norm)
|
||||||
|
optimizer.update(model, gradients)
|
||||||
|
mx.eval(model.parameters(), optimizer.state, *losses)
|
||||||
|
running += np.asarray([float(value.item()) for value in losses])
|
||||||
|
if args.progress_steps and (
|
||||||
|
step % args.progress_steps == 0 or step == steps_per_epoch
|
||||||
|
):
|
||||||
|
mean = running / step
|
||||||
|
print(
|
||||||
|
f"epoch {epoch} step {step}/{steps_per_epoch} "
|
||||||
|
f"loss={mean[0]:.4f} primary={mean[1]:.4f} "
|
||||||
|
f"secondary={mean[2]:.4f} mixed={mean[3]:.4f} "
|
||||||
|
f"difficulty={mean[4]:.4f}"
|
||||||
|
+ (
|
||||||
|
f" distillation={mean[5]:.4f}"
|
||||||
|
if teacher_train_logits is not None
|
||||||
|
else ""
|
||||||
|
)
|
||||||
|
+ " "
|
||||||
|
f"elapsed={time.perf_counter() - epoch_started:.1f}s",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
if shutdown.signum is not None:
|
||||||
|
pause_training(
|
||||||
|
epoch,
|
||||||
|
step + 1,
|
||||||
|
permutation,
|
||||||
|
running,
|
||||||
|
time.perf_counter() - epoch_started,
|
||||||
|
)
|
||||||
|
|
||||||
|
outputs = _evaluate(
|
||||||
|
mx, model, encoded_validation, eval_batch_size
|
||||||
|
)
|
||||||
|
metrics = multitask_metrics(outputs, validation_records)
|
||||||
|
if teacher_validation_logits is not None:
|
||||||
|
metrics["teacherAgreement"] = float(
|
||||||
|
np.mean(
|
||||||
|
outputs["purpose_logits"][validation_scorable].argmax(
|
||||||
|
axis=-1
|
||||||
|
)
|
||||||
|
== teacher_validation_logits[validation_scorable].argmax(
|
||||||
|
axis=-1
|
||||||
|
)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
metrics["epoch"] = epoch
|
||||||
|
metrics["meanTrainingLoss"] = (running / steps_per_epoch).tolist()
|
||||||
|
history.append(metrics)
|
||||||
|
score = float(metrics["selectionScore"])
|
||||||
|
secondary_macro = (
|
||||||
|
metrics["secondary"]["supportedMacroRecall"]
|
||||||
|
if metrics["secondary"] is not None
|
||||||
|
else 0.0
|
||||||
|
)
|
||||||
|
print(
|
||||||
|
f"epoch {epoch}: primary_accuracy={metrics['primary']['accuracy']:.4%} "
|
||||||
|
f"hard_accuracy={metrics['primaryHardSlice']['accuracy']:.4%} "
|
||||||
|
f"secondary_supported_macro_recall={secondary_macro:.4%} "
|
||||||
|
f"mixed_f1={metrics['mixed']['f1']:.4%} "
|
||||||
|
f"difficulty_mae={metrics['difficulty']['mae']:.4f} "
|
||||||
|
f"selection_score={score:.4%}"
|
||||||
|
+ (
|
||||||
|
f" teacher_agreement={metrics['teacherAgreement']:.4%}"
|
||||||
|
if "teacherAgreement" in metrics
|
||||||
|
else ""
|
||||||
|
),
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
improvement = score - best_score
|
||||||
|
if improvement > args.minimum_improvement:
|
||||||
|
best_score = score
|
||||||
|
best_metrics = metrics
|
||||||
|
epochs_without_improvement = 0
|
||||||
|
_save_checkpoint(
|
||||||
|
mx, model, tokenizer, best_dir, checkpoint_config
|
||||||
|
)
|
||||||
|
write_json(
|
||||||
|
output_dir / "training-state.json",
|
||||||
|
{
|
||||||
|
"bestEpoch": epoch,
|
||||||
|
"bestSelectionScore": best_score,
|
||||||
|
"elapsedSeconds": time.perf_counter() - started,
|
||||||
|
"complete": False,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
epochs_without_improvement += 1
|
||||||
|
if epochs_without_improvement >= args.early_stopping_patience:
|
||||||
|
stopped_early = True
|
||||||
|
print(
|
||||||
|
f"early stopping after epoch {epoch}: no hard-aware "
|
||||||
|
f"selection improvement greater than "
|
||||||
|
f"{args.minimum_improvement:.4%} for "
|
||||||
|
f"{args.early_stopping_patience} epoch(s)",
|
||||||
|
flush=True,
|
||||||
|
)
|
||||||
|
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
|
||||||
|
|||||||
Reference in New Issue
Block a user