unsloth/studio/backend/tests/test_mlx_stop_checkpoint.py
Nilay f3c085ad9e
Fix resume training crash recovery and MLX checkpoints (#6796)
* 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>
2026-07-21 02:34:58 -07:00

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