* Fix resume training crash recovery and MLX checkpoints * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Address review: preserve interrupted stop-and-save output_dir, verify MLX checkpoint - finish_run: add clear_output_dir flag; preserve output_dir for stopped/error unless cancel explicitly clears it (fixes pump finalization wiping persisted path). - training pump: pass interrupted stop-and-save context into finalize_run_in_db. - MLX stop-and-save: verify resumable checkpoint exists before sending complete; return bool from _write_mlx_stop_checkpoint and add regression tests. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Address Codex review: MLX current-step checkpoint and cancel error finalize - Only skip MLX stop checkpoint write when checkpoint-{current_step} exists; stale periodic checkpoints no longer mask missing stop saves. - Pass clear_output_dir through error-event finalization so Stop-without-save cannot leave a persisted output_dir that still offers Resume. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Address review * Address more reviews * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * more reviews * clear in-memory output_dir on interrupted cancel * allow resuming errored runs at the final step * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Clear persisted output_dir in cancel watchdog path * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Write MLX stop checkpoint in stop path, keep output_dir on crash finalize * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * fix(studio): harden resumable run finalization * fix(studio): defer safetensors checkpoint import * fix(studio): reject stale training cancellation * fix(studio): replay null resume targets * fix(studio): serialize terminal cancellation * Harden resume checkpoint validation and fix stop-save cleanup - Reject unrecognized shard formats and keep indexed shard paths inside the checkpoint dir - Require a non-empty tensor record when validating .pt/.bin optimizer and model state - Always finalize TensorBoard and W&B on stop-save-failure exits - Refuse writing an MLX stop checkpoint through a symlinked directory - Clarify the resume rejection message to cover errored runs * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Tighten resume/checkpoint comments * Recover resumability when a valid stop checkpoint landed - Re-validate the current-step checkpoint in the dead-worker and error finalization paths so a stop-and-save that actually wrote a valid checkpoint is not wrongly marked error/resume_blocked - Accept a valid tensor-free optimizer state (e.g. SGD without momentum); the model-state check still requires real tensors - Include errored runs in the frontend resume rejection message * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> Co-authored-by: Lyxot <longyixing331@gmail.com> Co-authored-by: Lee Jackson <130007945+Imagineer99@users.noreply.github.com> Co-authored-by: danielhanchen <danielhanchen@gmail.com>
137 lines
4.5 KiB
Python
137 lines
4.5 KiB
Python
# SPDX-License-Identifier: AGPL-3.0-only
|
|
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
|
|
|
"""Regression tests for MLX stop-and-save checkpoint handling."""
|
|
|
|
import importlib.util
|
|
import json
|
|
import sys
|
|
import types
|
|
from pathlib import Path
|
|
|
|
import numpy as np
|
|
from safetensors.numpy import save_file
|
|
|
|
|
|
_BACKEND = Path(__file__).resolve().parents[1]
|
|
|
|
|
|
def _load_worker_module():
|
|
spec = importlib.util.spec_from_file_location(
|
|
"training_worker_under_test",
|
|
_BACKEND / "core" / "training" / "worker.py",
|
|
)
|
|
module = importlib.util.module_from_spec(spec)
|
|
assert spec.loader is not None
|
|
spec.loader.exec_module(module)
|
|
return module
|
|
|
|
|
|
worker = _load_worker_module()
|
|
|
|
|
|
class _FakeTrainer:
|
|
def __init__(self, step: int):
|
|
self._global_step = step
|
|
self._train_loss_history = []
|
|
self.model = object()
|
|
|
|
|
|
def _write_checkpoint(out: Path, step: int) -> Path:
|
|
checkpoint = out / f"checkpoint-{step}"
|
|
checkpoint.mkdir(parents = True, exist_ok = True)
|
|
(checkpoint / "trainer_state.json").write_text(
|
|
json.dumps({"global_step": step}), encoding = "utf-8"
|
|
)
|
|
save_file({"weight": np.ones(1, dtype = np.float32)}, checkpoint / "adapters.safetensors")
|
|
save_file(
|
|
{"state": np.ones(1, dtype = np.float32)},
|
|
checkpoint / "optimizer_state.safetensors",
|
|
)
|
|
return checkpoint
|
|
|
|
|
|
def test_mlx_has_checkpoint_at_step_requires_complete_state(tmp_path):
|
|
out = tmp_path / "outputs" / "run_x"
|
|
_write_checkpoint(out, 5)
|
|
|
|
assert worker._mlx_has_checkpoint_at_step(out, 5) is True
|
|
|
|
|
|
def test_write_mlx_stop_checkpoint_returns_true_when_current_step_checkpoint_exists(tmp_path):
|
|
out = tmp_path / "outputs" / "run_x"
|
|
_write_checkpoint(out, 5)
|
|
|
|
assert worker._write_mlx_stop_checkpoint(_FakeTrainer(step = 5), object(), out) is True
|
|
|
|
|
|
def test_write_mlx_stop_checkpoint_writes_current_step_when_only_older_checkpoint_exists(
|
|
tmp_path, monkeypatch
|
|
):
|
|
out = tmp_path / "outputs" / "run_x"
|
|
_write_checkpoint(out, 5)
|
|
|
|
saved_steps: list[int] = []
|
|
|
|
def _save_state(_value, path, name):
|
|
save_file({"state": np.ones(1, dtype = np.float32)}, Path(path, name))
|
|
|
|
def _save_trainer_state(state, ckpt_dir, **_kwargs):
|
|
Path(ckpt_dir, "trainer_state.json").write_text(json.dumps(state), encoding = "utf-8")
|
|
saved_steps.append(int(state["global_step"]))
|
|
|
|
fake_utils = types.SimpleNamespace(
|
|
save_trainable_adapters = lambda model, path: _save_state(
|
|
model, path, "adapters.safetensors"
|
|
),
|
|
save_optimizer_state = lambda optimizer, path: _save_state(
|
|
optimizer, path, "optimizer_state.safetensors"
|
|
),
|
|
save_trainer_state = _save_trainer_state,
|
|
)
|
|
monkeypatch.setitem(sys.modules, "unsloth_zoo.mlx.utils", fake_utils)
|
|
|
|
assert worker._write_mlx_stop_checkpoint(_FakeTrainer(step = 10), object(), out) is True
|
|
assert saved_steps == [10]
|
|
assert (out / "checkpoint-10" / "trainer_state.json").is_file()
|
|
|
|
|
|
def test_write_mlx_stop_checkpoint_returns_false_without_optimizer(tmp_path):
|
|
out = tmp_path / "outputs" / "run_x"
|
|
out.mkdir(parents = True)
|
|
|
|
assert worker._write_mlx_stop_checkpoint(_FakeTrainer(step = 5), None, out) is False
|
|
|
|
|
|
def test_write_mlx_stop_checkpoint_rejects_incomplete_current_checkpoint(tmp_path):
|
|
out = tmp_path / "outputs" / "run_x"
|
|
ckpt = out / "checkpoint-5"
|
|
ckpt.mkdir(parents = True)
|
|
(ckpt / "trainer_state.json").write_text('{"global_step": 5}', encoding = "utf-8")
|
|
|
|
assert worker._write_mlx_stop_checkpoint(_FakeTrainer(step = 5), None, out) is False
|
|
|
|
|
|
def test_write_mlx_stop_checkpoint_ignores_stale_checkpoint_without_optimizer(tmp_path):
|
|
# An older checkpoint does not cover the current step, so this still fails.
|
|
out = tmp_path / "outputs" / "run_x"
|
|
_write_checkpoint(out, 5)
|
|
|
|
assert worker._write_mlx_stop_checkpoint(_FakeTrainer(step = 10), None, out) is False
|
|
|
|
|
|
def test_write_mlx_stop_checkpoint_returns_false_when_save_fails(tmp_path, monkeypatch):
|
|
out = tmp_path / "outputs" / "run_x"
|
|
out.mkdir(parents = True)
|
|
|
|
def _boom(*_args, **_kwargs):
|
|
raise RuntimeError("save failed")
|
|
|
|
fake_utils = types.SimpleNamespace(
|
|
save_trainable_adapters = _boom,
|
|
save_optimizer_state = lambda *_a, **_k: None,
|
|
save_trainer_state = lambda *_a, **_k: None,
|
|
)
|
|
monkeypatch.setitem(sys.modules, "unsloth_zoo.mlx.utils", fake_utils)
|
|
|
|
assert worker._write_mlx_stop_checkpoint(_FakeTrainer(step = 5), object(), out) is False
|