Studio: force-terminate a stuck training stop after a grace period (#7099)
* Studio: force-terminate a stuck training stop after a grace period A Stop-with-save only signals the worker and waits for it to save and exit; force_terminate() was reachable only from the /reset cancel path. On Windows + ROCm the worker saves the adapter fine but then wedges in post-save GPU/HIP teardown and never exits, so the run stays in "Stopping..." forever, is_training stays true, and /reset returns 409. Add a stop watchdog: when a stop is requested, a daemon escalates to force_terminate() a short grace after the worker's "complete" (save done), or after an absolute cap covering a hang during save. After escalation the parent state is finalized (is_training=False, "Training stopped.") even if the OS never reaps the wedged worker, so the UI leaves "Stopping..." and a new run can start. No behavior change on a clean quick exit. Grace and timeout are configurable via UNSLOTH_STUDIO_TRAINING_STOP_GRACE_S (15) and UNSLOTH_STUDIO_TRAINING_STOP_TIMEOUT_S (120). * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Studio: harden the training stop watchdog per review Address review feedback so a stop can never corrupt a checkpoint or leave the run stuck: - Never force-kill an in-progress save. The absolute cap is now a last-resort backstop: raise the save default to 600s and only kill past that long window; a not-yet-complete save is not treated as a hang. Cancels have nothing to save, so they keep a shorter 120s cap via UNSLOTH_STUDIO_TRAINING_CANCEL_TIMEOUT_S. The save vs cancel path is now explicit and the backstop logs a clear warning. - Always finalize even if force_terminate raises on a wedged child (try/finally), so the watchdog never dies leaving the run in "Stopping...". - Preserve output_dir when the watchdog finalizes so a saved checkpoint is still recorded in run history. - Track the watched process per watchdog: a new run always gets its own watcher, and a stale watchdog on an old proc no longer suppresses it. - Terminate only the captured proc; force_terminate revalidates under the lock that it is still the current worker, so it can never kill a fresh run. - Name the watchdog thread for debuggability. * studio: tighten training-stop watchdog comments Comment-only pass: collapse the watchdog docstrings and inline notes to fewer lines while keeping the rationale. No behavior change. * Studio: make the stop watchdog safe against concurrent runs and the pump Target-scope the escalation finalize so a stale watchdog can never clobber a run that replaced its worker: capture the watched proc and job id, and no-op the finalize (handle, progress, and DB) when a new run has already taken over. Honor a later cancel by tightening an in-flight save watchdog to the shorter cancel cap. Serialize the DB helpers on the lock so the watchdog and pump can no longer double-create, double-finalize, or corrupt the metric buffer when a force-terminate hands off to a still-finalizing pump. Add regression tests: finalize no-ops when superseded, finalize runs for its own worker, a later cancel tightens the cap, finalize is single-winner under concurrency, finalize honors expected_job_id, and concurrent flushes claim each metric exactly once. * Studio: close the remaining stop-watchdog vs start/pump races Guard the escalation finalize by the watched job id in addition to the proc: start_training sets current_job_id before it installs the new _proc, so a stale watchdog entering during that startup window still sees the old dead handle and was not caught by the proc-only guard. Capture the job id when the watchdog starts and require it to still match before touching state. Snapshot the run id and final progress under the finalize lock and thread them through the flush and finish_run calls, so a new run that starts between the finalize claim and the DB writes cannot be flushed or marked stopped under the old run's finalizer. Publish _db_run_created only after create_run commits, gated by a dedicated in-progress flag, so a concurrent finalize can no longer run finish_run against a not-yet-inserted row and leave the run stuck as running. Add regression tests for the startup-window job-id guard, run-id pinned flush, snapshot-based finalize across a new run, and create-not-published-before-insert. * Studio: finalize a force-stopped run by its captured id If a new run starts in the gap after the watchdog clears _proc and marks the backend idle, current_job_id changes, so the previous expected_job_id guard made the finalize skip and left the stopped run recorded as running. Capture the run id, metrics, and final progress under the lock (where current_job_id is still the watched run) and finalize by that captured id via _finish_stopped_run: finish_run is an idempotent UPDATE and insert_metrics_batch upserts, so a concurrent pump finalize of the same run is harmless and a newly started run is never touched. Add a test that the watched run is finalized by id with its buffered metrics, and update the escalation tests to assert finalize goes through _finish_stopped_run. * Studio: keep force-stop finalization retryable and unclaimed until the row exists Only claim _run_finalized in the escalation when the DB row already exists; if an early create failed and the pump is retrying it, claiming would make the pump's later finalize no-op and strand the row as running, so leave the finalize to that create-then-finalize path. On a DB error in _finish_stopped_run (e.g. a transient SQLite lock), unclaim the finalize and requeue the drained metrics when the run is still current, so the pump or a later retry can still record the run stopped instead of leaving history with an active run and lost metrics. A superseded run's state is never touched. Add tests: no claim before the row exists, requeue+unclaim on a DB error, and a superseded run left untouched on error. * Studio: tighten stop-watchdog comments Reduce the wording of the docstrings and inline comments added by this PR without dropping any of the concurrency invariants (dual proc/job-id supersession guard, finalize-by-captured-id, publish-after-commit, snapshot-under-lock, unclaim and requeue on error). Comments and docstrings only; no code change. * Studio: record the stopped run's DB state before dropping _proc A wedged worker still reports alive, so the pump never reaches its own finalize and bails on its _proc-is-None guard once the escalation drops the handle. So the watchdog is the sole finalizer: record the terminal DB state (create the row if a start-time create failed, then finish by captured id) BEFORE dropping _proc. While the handle is held is_training_active() stays true, so no new run can start and current_job_id stays the watched run for the write; _proc is dropped last, guarded on target_proc so a run that did replace the worker keeps its handle. _finish_stopped_run retries a transient DB error a few times (the pump can no longer retry once _proc is gone) and unclaims on final failure only when the run is still current. Add tests for create-then-finalize, retry-then-unclaim, and not dropping a new run's handle. * Studio: job-guard the DB create flags against a racing new run _ensure_db_run_created publishes backend-wide _db_run_created and _db_create_in_progress flags. When the watchdog creates a missing row for an escalated stop, the killed worker lets a new /start proceed mid-create, so the stale create could publish those flags against the new current_job_id, making the new run skip inserting its own row (metric/finalize then target a missing run). Publish the flags only when the captured job id is still current; the row is still created by id, and the new run owns/creates its own. Also reset _db_create_in_progress in start_training so a stale claim can't block a new run. --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
This commit is contained in:
parent
67339b15fd
commit
4beb0a3a5f
2 changed files with 1153 additions and 54 deletions
|
|
@ -39,6 +39,27 @@ from utils.paths import outputs_root
|
|||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
def _env_int(name: str, default: int) -> int:
|
||||
try:
|
||||
raw = (os.environ.get(name) or "").strip()
|
||||
return int(raw) if raw else default
|
||||
except ValueError:
|
||||
return default
|
||||
|
||||
|
||||
# Stop-watchdog escalation timeouts. Primary trigger: a short grace once "complete"
|
||||
# (save done). Absolute cap is a backstop: long for save=True so a slow save is never
|
||||
# killed mid-write, shorter for a cancel that has nothing to save.
|
||||
_STOP_GRACE_S = _env_int("UNSLOTH_STUDIO_TRAINING_STOP_GRACE_S", 15)
|
||||
_STOP_TIMEOUT_S = _env_int("UNSLOTH_STUDIO_TRAINING_STOP_TIMEOUT_S", 600)
|
||||
_CANCEL_TIMEOUT_S = _env_int("UNSLOTH_STUDIO_TRAINING_CANCEL_TIMEOUT_S", 120)
|
||||
|
||||
# Watchdog DB finalize: a few short retries so a transient SQLite lock doesn't lose the
|
||||
# terminal state, since the watchdog is the sole finalizer once _proc is dropped.
|
||||
_DB_FINALIZE_RETRIES = 3
|
||||
_DB_FINALIZE_RETRY_S = 0.5
|
||||
|
||||
_pyplot = None
|
||||
_pyplot_failed = False
|
||||
|
||||
|
|
@ -741,6 +762,13 @@ class TrainingBackend:
|
|||
self._pump_running: bool = False
|
||||
self._lock = threading.Lock()
|
||||
|
||||
# Stop watchdog: after a stop is requested, escalates to force_terminate()
|
||||
# if the worker does not exit on its own within a bounded time. The watched
|
||||
# proc is tracked so a new run always gets its own watcher.
|
||||
self._stop_watchdog: Optional[threading.Thread] = None
|
||||
self._stop_watchdog_proc: Optional[mp.Process] = None
|
||||
self._complete_seen = threading.Event()
|
||||
|
||||
# Progress state (updated by pump thread from subprocess events)
|
||||
self._progress = TrainingProgress()
|
||||
self._should_stop = False
|
||||
|
|
@ -765,6 +793,7 @@ class TrainingBackend:
|
|||
self._metric_buffer: list[dict] = []
|
||||
self._run_finalized: bool = False
|
||||
self._db_run_created: bool = False
|
||||
self._db_create_in_progress: bool = False
|
||||
self._db_total_steps_set: bool = False
|
||||
self._db_config: Optional[dict] = None
|
||||
self._db_started_at: Optional[str] = None
|
||||
|
|
@ -896,6 +925,7 @@ class TrainingBackend:
|
|||
self.current_job_id = job_id
|
||||
self._should_stop = False
|
||||
self._cancel_requested = False
|
||||
self._complete_seen.clear()
|
||||
self._progress = TrainingProgress(
|
||||
is_training = True, status_message = "Initializing training..."
|
||||
)
|
||||
|
|
@ -911,6 +941,7 @@ class TrainingBackend:
|
|||
self._metric_buffer.clear()
|
||||
self._run_finalized = False
|
||||
self._db_run_created = False
|
||||
self._db_create_in_progress = False # a stale watchdog create can't block this run
|
||||
self._db_total_steps_set = False
|
||||
self._db_config = _sanitize_db_config(config)
|
||||
self._db_started_at = datetime.now(timezone.utc).isoformat()
|
||||
|
|
@ -953,15 +984,212 @@ class TrainingBackend:
|
|||
self._progress.status_message = (
|
||||
"Stopping training and saving checkpoint..." if save else "Cancelling training..."
|
||||
)
|
||||
# Guarantee the run finalizes even if the worker wedges after saving.
|
||||
self._start_stop_watchdog(cancel = not save)
|
||||
return True
|
||||
|
||||
def force_terminate(self) -> None:
|
||||
"""Force-kill the training subprocess so state can be reset immediately."""
|
||||
def _start_stop_watchdog(self, cancel: bool) -> None:
|
||||
"""Start a daemon that force-terminates the worker if a requested stop does not
|
||||
exit on its own. No-op if no worker is alive or a live watchdog already watches
|
||||
this proc (a stale watchdog on an old proc never blocks a new run's watcher)."""
|
||||
with self._lock:
|
||||
if self._proc is not None and self._proc.is_alive():
|
||||
logger.info("Force-terminating training subprocess (pid=%s)", self._proc.pid)
|
||||
self._proc.terminate()
|
||||
proc = self._proc
|
||||
if proc is None or not proc.is_alive():
|
||||
return
|
||||
if (
|
||||
self._stop_watchdog is not None
|
||||
and self._stop_watchdog.is_alive()
|
||||
and self._stop_watchdog_proc is proc
|
||||
):
|
||||
return
|
||||
watchdog = threading.Thread(
|
||||
target = self._stop_watchdog_loop,
|
||||
args = (proc, cancel, self.current_job_id),
|
||||
name = f"stop-watchdog-{self.current_job_id or 'unknown'}",
|
||||
daemon = True,
|
||||
)
|
||||
self._stop_watchdog = watchdog
|
||||
self._stop_watchdog_proc = proc
|
||||
watchdog.start()
|
||||
|
||||
def _stop_watchdog_loop(
|
||||
self,
|
||||
target_proc: "mp.Process",
|
||||
cancel: bool,
|
||||
watched_job_id: Optional[str] = None,
|
||||
) -> None:
|
||||
"""Escalate a stuck stop to force_terminate(): grace after "complete", else the
|
||||
absolute backstop (see the module timeouts). No-ops on a clean exit; exits
|
||||
silently if a new run replaces the worker."""
|
||||
started = time.monotonic()
|
||||
complete_at: Optional[float] = None
|
||||
reason = ""
|
||||
while True:
|
||||
with self._lock:
|
||||
superseded = self._proc is not target_proc
|
||||
# A later cancel has nothing to save, so tighten an in-flight save
|
||||
# watchdog to the shorter cancel cap.
|
||||
cancelling = cancel or self._cancel_requested
|
||||
if superseded or not target_proc.is_alive():
|
||||
return
|
||||
now = time.monotonic()
|
||||
abs_timeout = _CANCEL_TIMEOUT_S if cancelling else _STOP_TIMEOUT_S
|
||||
if complete_at is None and self._complete_seen.is_set():
|
||||
complete_at = now
|
||||
if complete_at is not None and now - complete_at >= _STOP_GRACE_S:
|
||||
reason = "worker still alive after save"
|
||||
break
|
||||
if now - started >= abs_timeout:
|
||||
reason = "worker did not exit within the absolute timeout"
|
||||
break
|
||||
time.sleep(0.5)
|
||||
|
||||
with self._lock:
|
||||
superseded = self._proc is not target_proc
|
||||
if superseded or not target_proc.is_alive():
|
||||
return
|
||||
if complete_at is None:
|
||||
# Backstop fired pre-completion: a save may still be in progress.
|
||||
logger.warning(
|
||||
"Stop watchdog: absolute timeout with no completion signal; "
|
||||
"force-terminating a possibly-mid-save worker: %s",
|
||||
reason,
|
||||
)
|
||||
else:
|
||||
logger.warning("Stop watchdog force-terminating stuck training worker: %s", reason)
|
||||
# force_terminate can raise on a wedged child; finalize regardless.
|
||||
try:
|
||||
self.force_terminate(target_proc = target_proc)
|
||||
except Exception:
|
||||
logger.exception("Stop watchdog: force_terminate failed; finalizing anyway")
|
||||
finally:
|
||||
self._finalize_stopped_after_escalation(
|
||||
target_proc = target_proc, watched_job_id = watched_job_id
|
||||
)
|
||||
|
||||
def _finalize_stopped_after_escalation(
|
||||
self,
|
||||
target_proc: "Optional[mp.Process]" = None,
|
||||
watched_job_id: Optional[str] = None,
|
||||
) -> None:
|
||||
"""Finalize parent state after a force-terminate so the UI leaves "Stopping..."
|
||||
even if the worker is wedged in driver teardown; preserves output_dir so a saved
|
||||
checkpoint is kept. No-ops if a new run already replaced the watched worker, so a
|
||||
stale watchdog never marks a fresh run stopped or drops its handle.
|
||||
|
||||
Supersession is checked on both the watched proc and job id: start_training sets
|
||||
current_job_id before it installs the new _proc, so a stale watchdog entering that
|
||||
startup window still sees the old (dead) handle and is caught by the job-id guard.
|
||||
|
||||
The run's terminal DB state is recorded (create-if-needed + finish by captured id)
|
||||
BEFORE _proc is dropped: a wedged worker still reports alive, so the pump never
|
||||
reaches its own finalize and would bail on its _proc-is-None guard once the handle
|
||||
is gone. While the handle is held is_training_active() stays true, so no new run can
|
||||
start and current_job_id stays the watched run for the write. _proc is dropped last,
|
||||
re-guarded on target_proc so a run that did replace the worker keeps its handle."""
|
||||
with self._lock:
|
||||
if target_proc is not None and self._proc is not target_proc:
|
||||
return # a new run replaced the worker; never touch its state
|
||||
if watched_job_id is not None and self.current_job_id != watched_job_id:
|
||||
return # a new run is already starting up; leave its state alone
|
||||
run_id = self.current_job_id # == watched_job_id
|
||||
self._progress.is_training = False
|
||||
self._progress.status_message = "Training stopped."
|
||||
# Create the row if a start-time create failed (no-op otherwise; skips when the pump
|
||||
# is mid-create, in which case its create-then-finalize records the run instead).
|
||||
self._ensure_db_run_created()
|
||||
with self._lock:
|
||||
claim = (
|
||||
bool(run_id)
|
||||
and self.current_job_id == run_id
|
||||
and self._db_run_created
|
||||
and not self._run_finalized
|
||||
)
|
||||
batch: list = []
|
||||
final_step = final_loss = duration = None
|
||||
loss_history: list = []
|
||||
output_dir = self._output_dir
|
||||
if claim:
|
||||
self._run_finalized = True # claim this run's finalize
|
||||
batch = list(self._metric_buffer)
|
||||
del self._metric_buffer[: len(batch)]
|
||||
final_step = self._progress.step
|
||||
final_loss = self._progress.loss
|
||||
if final_loss is not None and not math.isfinite(final_loss):
|
||||
final_loss = None
|
||||
duration = self._progress.elapsed_seconds
|
||||
loss_history = list(self.loss_history)
|
||||
if claim:
|
||||
self._finish_stopped_run(
|
||||
run_id, output_dir, batch, final_step, final_loss, duration, loss_history
|
||||
)
|
||||
with self._lock:
|
||||
if target_proc is None or self._proc is target_proc:
|
||||
self._proc = None # drop only our handle, never a run that replaced it
|
||||
|
||||
def _finish_stopped_run(
|
||||
self,
|
||||
run_id: str,
|
||||
output_dir: Optional[str],
|
||||
batch: list,
|
||||
final_step: Optional[int],
|
||||
final_loss: Optional[float],
|
||||
duration: Optional[float],
|
||||
loss_history: list,
|
||||
) -> None:
|
||||
"""Record a force-stopped run finished by its captured id, from state snapshotted
|
||||
under the lock. insert_metrics_batch upserts and finish_run is an idempotent UPDATE,
|
||||
so a concurrent pump finalize of the same run is harmless and a different current run
|
||||
is never touched. The watchdog is the sole finalizer once _proc is dropped, so a
|
||||
transient DB error (e.g. a SQLite lock) is retried a few times; on final failure the
|
||||
finalize is unclaimed (only if the run is still current) so the row is not left
|
||||
claimed-but-unfinalized."""
|
||||
for attempt in range(_DB_FINALIZE_RETRIES):
|
||||
try:
|
||||
from storage.studio_db import finish_run, insert_metrics_batch
|
||||
from utils.downsample import downsample
|
||||
|
||||
if batch:
|
||||
insert_metrics_batch(run_id, batch)
|
||||
sparkline = downsample(loss_history, 50)
|
||||
finish_run(
|
||||
id = run_id,
|
||||
status = "stopped",
|
||||
ended_at = datetime.now(timezone.utc).isoformat(),
|
||||
final_step = final_step,
|
||||
final_loss = final_loss,
|
||||
duration_seconds = duration,
|
||||
loss_sparkline = _json.dumps(sparkline),
|
||||
output_dir = output_dir,
|
||||
error_message = None,
|
||||
)
|
||||
return
|
||||
except Exception:
|
||||
if attempt + 1 < _DB_FINALIZE_RETRIES:
|
||||
time.sleep(_DB_FINALIZE_RETRY_S)
|
||||
continue
|
||||
logger.warning(
|
||||
"Failed to finalize stopped run %s in DB after %d attempts",
|
||||
run_id,
|
||||
_DB_FINALIZE_RETRIES,
|
||||
exc_info = True,
|
||||
)
|
||||
with self._lock:
|
||||
# Only if still current; a new run's finalize state is never touched.
|
||||
if self.current_job_id == run_id:
|
||||
self._run_finalized = False
|
||||
|
||||
def force_terminate(self, target_proc: "Optional[mp.Process]" = None) -> None:
|
||||
"""Force-kill the training subprocess so state can be reset immediately. With
|
||||
``target_proc``, terminate only that handle and no-op if a new run has replaced
|
||||
it, so the watchdog can never kill a fresh worker."""
|
||||
with self._lock:
|
||||
proc = self._proc
|
||||
if target_proc is not None and proc is not target_proc:
|
||||
return # superseded by a new run; do not touch the new worker
|
||||
if proc is not None and proc.is_alive():
|
||||
logger.info("Force-terminating training subprocess (pid=%s)", proc.pid)
|
||||
proc.terminate()
|
||||
cancelled = self._cancel_requested
|
||||
output_dir = self._output_dir
|
||||
|
||||
|
|
@ -1468,6 +1696,8 @@ class TrainingBackend:
|
|||
"training cancelled",
|
||||
"training stopped",
|
||||
}
|
||||
# Save is done by now; let the stop watchdog start its grace timer.
|
||||
self._complete_seen.set()
|
||||
self._progress.is_training = False
|
||||
self._progress.is_completed = not stopped
|
||||
self._output_dir = event.get("output_dir")
|
||||
|
|
@ -1532,90 +1762,135 @@ class TrainingBackend:
|
|||
self._finalize_run_in_db(**db_action_kwargs)
|
||||
|
||||
def _ensure_db_run_created(self) -> None:
|
||||
"""Create the DB row if it doesn't exist yet. Called outside the lock."""
|
||||
if self._db_run_created or not self.current_job_id or not self._db_config:
|
||||
return
|
||||
"""Create the DB row if it doesn't exist yet. An in-progress flag lets only one
|
||||
caller create at a time, and ``_db_run_created`` is published only after
|
||||
``create_run`` commits, so a concurrent finalize never runs ``finish_run`` against a
|
||||
not-yet-inserted row (a zero-row UPDATE that would leave the run stuck as running)."""
|
||||
with self._lock:
|
||||
if (
|
||||
self._db_run_created
|
||||
or self._db_create_in_progress
|
||||
or not self.current_job_id
|
||||
or not self._db_config
|
||||
):
|
||||
return
|
||||
self._db_create_in_progress = True # only one caller creates
|
||||
job_id = self.current_job_id
|
||||
db_config = self._db_config
|
||||
started_at = self._db_started_at or datetime.now(timezone.utc).isoformat()
|
||||
total_steps = self._progress.total_steps or None
|
||||
created = False
|
||||
try:
|
||||
from storage.studio_db import create_run
|
||||
|
||||
dataset_name = (
|
||||
self._db_config.get("hf_dataset")
|
||||
or next(iter(self._db_config.get("local_datasets") or []), None)
|
||||
or _s3_dataset_name(self._db_config.get("s3_dataset"))
|
||||
db_config.get("hf_dataset")
|
||||
or next(iter(db_config.get("local_datasets") or []), None)
|
||||
or _s3_dataset_name(db_config.get("s3_dataset"))
|
||||
or "unknown"
|
||||
)
|
||||
create_run(
|
||||
id = self.current_job_id,
|
||||
model_name = self._db_config["model_name"],
|
||||
id = job_id,
|
||||
model_name = db_config["model_name"],
|
||||
dataset_name = dataset_name,
|
||||
config_json = _json.dumps(self._db_config),
|
||||
started_at = self._db_started_at or datetime.now(timezone.utc).isoformat(),
|
||||
total_steps = self._progress.total_steps or None,
|
||||
config_json = _json.dumps(db_config),
|
||||
started_at = started_at,
|
||||
total_steps = total_steps,
|
||||
)
|
||||
self._db_run_created = True
|
||||
created = True
|
||||
except Exception:
|
||||
logger.warning("Failed to create DB run record for early failure", exc_info = True)
|
||||
finally:
|
||||
with self._lock:
|
||||
# Publish the flags only if this is still the current run. A killed worker
|
||||
# lets a new /start proceed mid-create, and these flags are backend-wide, so
|
||||
# a stale create for the captured job must not satisfy the new run's DB state
|
||||
# (the row was still created by id; the new run owns/creates its own row).
|
||||
if self.current_job_id == job_id:
|
||||
if created:
|
||||
self._db_run_created = True # publish only after the insert commits
|
||||
self._db_create_in_progress = False
|
||||
|
||||
def _finalize_run_in_db(
|
||||
self,
|
||||
status: str,
|
||||
error_message: Optional[str] = None,
|
||||
output_dir: Optional[str] = None,
|
||||
expected_job_id: Optional[str] = None,
|
||||
) -> None:
|
||||
"""Flush remaining metrics and mark a run as finished in the DB."""
|
||||
if not self.current_job_id or not self._db_run_created or self._run_finalized:
|
||||
return
|
||||
self._flush_metrics_to_db()
|
||||
"""Flush remaining metrics and mark a run finished in the DB. Claims the finalize
|
||||
under the lock so the watchdog and pump can't double-finalize, and no-ops when
|
||||
``expected_job_id`` no longer matches (a new run took over). The run id and final
|
||||
progress are snapshotted under the lock and threaded through the flush/finish calls,
|
||||
so a new run racing between this claim and the DB writes can't be flushed or marked
|
||||
stopped under the old run's finalize."""
|
||||
with self._lock:
|
||||
if expected_job_id is not None and self.current_job_id != expected_job_id:
|
||||
return
|
||||
if not self.current_job_id or not self._db_run_created or self._run_finalized:
|
||||
return
|
||||
self._run_finalized = True
|
||||
run_id = self.current_job_id
|
||||
final_step = self._progress.step
|
||||
final_loss = self._progress.loss
|
||||
if final_loss is not None and not math.isfinite(final_loss):
|
||||
final_loss = None
|
||||
duration = self._progress.elapsed_seconds
|
||||
loss_history = list(self.loss_history)
|
||||
self._flush_metrics_to_db(run_id = run_id)
|
||||
try:
|
||||
from storage.studio_db import finish_run
|
||||
from utils.downsample import downsample
|
||||
|
||||
sparkline = downsample(self.loss_history, 50)
|
||||
sparkline = downsample(loss_history, 50)
|
||||
finish_run(
|
||||
id = self.current_job_id,
|
||||
id = run_id,
|
||||
status = status,
|
||||
ended_at = datetime.now(timezone.utc).isoformat(),
|
||||
final_step = self._progress.step,
|
||||
final_loss = self._progress.loss
|
||||
if (self._progress.loss is not None and math.isfinite(self._progress.loss))
|
||||
else None,
|
||||
duration_seconds = self._progress.elapsed_seconds,
|
||||
final_step = final_step,
|
||||
final_loss = final_loss,
|
||||
duration_seconds = duration,
|
||||
loss_sparkline = _json.dumps(sparkline),
|
||||
output_dir = output_dir,
|
||||
error_message = error_message,
|
||||
)
|
||||
self._run_finalized = True
|
||||
except Exception:
|
||||
with self._lock:
|
||||
self._run_finalized = False # unclaim so a later flush can retry
|
||||
logger.warning("Failed to finalize run in DB (status=%s)", status, exc_info = True)
|
||||
|
||||
def _flush_metrics_to_db(self) -> None:
|
||||
"""Flush buffered metrics to the database and update live progress."""
|
||||
if not self._metric_buffer or not self.current_job_id or not self._db_run_created:
|
||||
return
|
||||
# Cap buffer to bound memory growth.
|
||||
if len(self._metric_buffer) > 500:
|
||||
logger.warning(
|
||||
"Metric buffer exceeded 500 entries (%d) — trimming oldest",
|
||||
len(self._metric_buffer),
|
||||
)
|
||||
self._metric_buffer = self._metric_buffer[-500:]
|
||||
# Snapshot before insert so metrics arriving during the write survive.
|
||||
batch = list(self._metric_buffer)
|
||||
def _flush_metrics_to_db(self, run_id: Optional[str] = None) -> None:
|
||||
"""Flush buffered metrics to the DB and update live progress. The target run id,
|
||||
metric batch, and progress snapshot are all taken under the lock, so a concurrent
|
||||
flush can't double-remove metrics and a racing new run can't redirect the write to
|
||||
a different job. A finalizer passes ``run_id`` to pin the target to its captured run."""
|
||||
with self._lock:
|
||||
target = run_id if run_id is not None else self.current_job_id
|
||||
if not self._metric_buffer or not target or not self._db_run_created:
|
||||
return
|
||||
# Cap buffer to bound memory growth.
|
||||
if len(self._metric_buffer) > 500:
|
||||
logger.warning(
|
||||
"Metric buffer exceeded 500 entries (%d) — trimming oldest",
|
||||
len(self._metric_buffer),
|
||||
)
|
||||
del self._metric_buffer[:-500]
|
||||
# Claim the batch under the lock so a concurrent flush can't re-remove it.
|
||||
batch = list(self._metric_buffer)
|
||||
del self._metric_buffer[: len(batch)]
|
||||
step = self._progress.step
|
||||
loss = self._progress.loss
|
||||
if loss is not None and not math.isfinite(loss):
|
||||
loss = None
|
||||
duration = self._progress.elapsed_seconds
|
||||
try:
|
||||
from storage.studio_db import insert_metrics_batch, update_run_progress
|
||||
|
||||
insert_metrics_batch(self.current_job_id, batch)
|
||||
del self._metric_buffer[: len(batch)]
|
||||
update_run_progress(
|
||||
id = self.current_job_id,
|
||||
step = self._progress.step,
|
||||
loss = self._progress.loss
|
||||
if (self._progress.loss is not None and math.isfinite(self._progress.loss))
|
||||
else None,
|
||||
duration_seconds = self._progress.elapsed_seconds,
|
||||
)
|
||||
insert_metrics_batch(target, batch)
|
||||
update_run_progress(id = target, step = step, loss = loss, duration_seconds = duration)
|
||||
except Exception:
|
||||
# Leave buffer intact for retry on next flush
|
||||
# Re-queue the claimed batch at the front so it retries on the next flush.
|
||||
with self._lock:
|
||||
self._metric_buffer[:0] = batch
|
||||
logger.warning("Failed to flush metrics to DB", exc_info = True)
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
824
studio/backend/tests/test_training_stop_watchdog.py
Normal file
824
studio/backend/tests/test_training_stop_watchdog.py
Normal file
|
|
@ -0,0 +1,824 @@
|
|||
# SPDX-License-Identifier: AGPL-3.0-only
|
||||
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
||||
|
||||
"""Stop-watchdog escalation for a stuck training stop.
|
||||
|
||||
A save-stop signals the worker and waits for it to save and exit. On some platforms the
|
||||
worker saves but then wedges in post-save GPU/driver teardown and never exits, leaving the
|
||||
run stuck in "Stopping..." forever. These tests pin the bounded recovery: the watchdog
|
||||
escalates to force_terminate() a short grace after "complete" (save done) or after an
|
||||
absolute timeout (hang during save), and never force-kills a worker that exits cleanly.
|
||||
Fakes only; no GPU, network, or subprocess.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import logging
|
||||
import queue
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
import types as _types
|
||||
from pathlib import Path
|
||||
|
||||
_BACKEND_DIR = str(Path(__file__).resolve().parent.parent)
|
||||
if _BACKEND_DIR not in sys.path:
|
||||
sys.path.insert(0, _BACKEND_DIR)
|
||||
|
||||
# Stub the heavy module-level imports of core/training/training.py so it imports
|
||||
# under CPU-only/no-network, then restore them (see the restore loop below).
|
||||
_SAVED: dict = {}
|
||||
|
||||
|
||||
def _stub(name, mod):
|
||||
_SAVED[name] = sys.modules.get(name)
|
||||
sys.modules[name] = mod
|
||||
|
||||
|
||||
_lg = _types.ModuleType("loggers")
|
||||
_lg.get_logger = lambda name: logging.getLogger(name)
|
||||
_stub("loggers", _lg)
|
||||
_stub("structlog", _types.ModuleType("structlog"))
|
||||
_mpl = _types.ModuleType("matplotlib")
|
||||
_plt = _types.ModuleType("matplotlib.pyplot")
|
||||
_plt.Figure = type("Figure", (), {}) # referenced in a class-def annotation
|
||||
_mpl.pyplot = _plt
|
||||
_stub("matplotlib", _mpl)
|
||||
_stub("matplotlib.pyplot", _plt)
|
||||
_hw = _types.ModuleType("utils.hardware")
|
||||
_hw.prepare_gpu_selection = lambda *a, **k: (None, None)
|
||||
_stub("utils.hardware", _hw)
|
||||
_npl = _types.ModuleType("utils.native_path_leases")
|
||||
_npl.native_path_secret_removed_for_child_start = lambda: contextlib.nullcontext()
|
||||
_npl.run_without_native_path_secret = lambda fn: fn
|
||||
_stub("utils.native_path_leases", _npl)
|
||||
_pth = _types.ModuleType("utils.paths")
|
||||
_pth.outputs_root = lambda *a, **k: "/tmp/outputs"
|
||||
_stub("utils.paths", _pth)
|
||||
|
||||
# Whether core.training.training was already imported before this file ran; only
|
||||
# evict it below if we were the one to create the (stub-bound) module instance.
|
||||
_TRAINING_PRE_IMPORTED = "core.training.training" in sys.modules
|
||||
|
||||
from core.training.training import TrainingBackend
|
||||
|
||||
# Restore every stubbed module so this file never pollutes the shared session.
|
||||
for _name in (
|
||||
"loggers",
|
||||
"structlog",
|
||||
"matplotlib",
|
||||
"matplotlib.pyplot",
|
||||
"utils.hardware",
|
||||
"utils.native_path_leases",
|
||||
"utils.paths",
|
||||
):
|
||||
_prev = _SAVED.get(_name)
|
||||
if _prev is None:
|
||||
sys.modules.pop(_name, None)
|
||||
else:
|
||||
sys.modules[_name] = _prev
|
||||
|
||||
if not _TRAINING_PRE_IMPORTED:
|
||||
sys.modules.pop("core.training.training", None)
|
||||
sys.modules.pop("core.training", None)
|
||||
|
||||
# The module globals hold the escalation timeouts and are the watchdog's own
|
||||
# namespace; patch them here so tests run in well under a second.
|
||||
_G = TrainingBackend._stop_watchdog_loop.__globals__
|
||||
|
||||
|
||||
class _FakeProc:
|
||||
"""A subprocess handle whose liveness and kill calls the test observes."""
|
||||
|
||||
def __init__(self, alive: bool = True):
|
||||
self._alive = alive
|
||||
self.pid = 4321
|
||||
self.terminated = False
|
||||
self.killed = False
|
||||
|
||||
def is_alive(self):
|
||||
return self._alive
|
||||
|
||||
def terminate(self):
|
||||
self.terminated = True
|
||||
|
||||
def kill(self):
|
||||
self.killed = True
|
||||
|
||||
def join(self, timeout = None):
|
||||
pass
|
||||
|
||||
|
||||
def _wait_until(predicate, timeout = 5.0):
|
||||
deadline = time.time() + timeout
|
||||
while time.time() < deadline:
|
||||
if predicate():
|
||||
return True
|
||||
time.sleep(0.01)
|
||||
return predicate()
|
||||
|
||||
|
||||
def _record_force_terminate(monkeypatch, b):
|
||||
"""Replace force_terminate + escalation finalize with recorders (no DB/OS)."""
|
||||
calls: list = []
|
||||
monkeypatch.setattr(b, "force_terminate", lambda target_proc = None: calls.append("force"))
|
||||
monkeypatch.setattr(
|
||||
b,
|
||||
"_finalize_stopped_after_escalation",
|
||||
lambda target_proc = None, watched_job_id = None: calls.append("final"),
|
||||
)
|
||||
return calls
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------------
|
||||
# (a) Escalate a short grace after "complete" (save done) if still alive.
|
||||
# ----------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_watchdog_escalates_after_grace_once_complete_seen(monkeypatch):
|
||||
monkeypatch.setitem(_G, "_STOP_GRACE_S", 0.05)
|
||||
monkeypatch.setitem(_G, "_STOP_TIMEOUT_S", 100.0) # ensure grace, not timeout, fires
|
||||
b = TrainingBackend()
|
||||
calls = _record_force_terminate(monkeypatch, b)
|
||||
|
||||
proc = _FakeProc(alive = True)
|
||||
b._proc = proc
|
||||
b._complete_seen.set() # worker reported "complete" -> save is done
|
||||
|
||||
b._start_stop_watchdog(cancel = False)
|
||||
assert _wait_until(
|
||||
lambda: calls == ["force", "final"]
|
||||
), "watchdog must force_terminate a worker still alive after the post-save grace"
|
||||
b._stop_watchdog.join(timeout = 5)
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------------
|
||||
# (b) The absolute cap is a last-resort backstop, not a save killer.
|
||||
# ----------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_watchdog_does_not_kill_save_still_saving_within_window(monkeypatch):
|
||||
# save=True, no "complete" yet: a slow save in progress must not be force-killed
|
||||
# inside the (long) absolute window.
|
||||
monkeypatch.setitem(_G, "_STOP_GRACE_S", 100.0)
|
||||
monkeypatch.setitem(_G, "_STOP_TIMEOUT_S", 100.0)
|
||||
b = TrainingBackend()
|
||||
calls = _record_force_terminate(monkeypatch, b)
|
||||
|
||||
proc = _FakeProc(alive = True)
|
||||
b._proc = proc
|
||||
b._start_stop_watchdog(cancel = False)
|
||||
|
||||
time.sleep(0.3)
|
||||
assert calls == [], "an in-progress save must not be killed within the absolute window"
|
||||
assert b._stop_watchdog.is_alive()
|
||||
|
||||
proc._alive = False
|
||||
b._stop_watchdog.join(timeout = 5)
|
||||
|
||||
|
||||
def test_watchdog_backstop_fires_for_save_after_absolute_timeout(monkeypatch):
|
||||
# Past the long save=True cap with no completion: force-terminate as last resort.
|
||||
monkeypatch.setitem(_G, "_STOP_GRACE_S", 100.0) # never trips (no complete)
|
||||
monkeypatch.setitem(_G, "_STOP_TIMEOUT_S", 0.05)
|
||||
b = TrainingBackend()
|
||||
calls = _record_force_terminate(monkeypatch, b)
|
||||
|
||||
b._proc = _FakeProc(alive = True)
|
||||
b._start_stop_watchdog(cancel = False)
|
||||
assert _wait_until(
|
||||
lambda: calls == ["force", "final"]
|
||||
), "the absolute backstop must force_terminate a save that never completes"
|
||||
b._stop_watchdog.join(timeout = 5)
|
||||
|
||||
|
||||
def test_cancel_uses_shorter_absolute_timeout(monkeypatch):
|
||||
# A cancel has nothing to save, so it escalates on the shorter cancel cap even before
|
||||
# the long save cap elapses.
|
||||
monkeypatch.setitem(_G, "_STOP_GRACE_S", 100.0)
|
||||
monkeypatch.setitem(_G, "_STOP_TIMEOUT_S", 100.0) # save cap would not fire
|
||||
monkeypatch.setitem(_G, "_CANCEL_TIMEOUT_S", 0.05)
|
||||
b = TrainingBackend()
|
||||
calls = _record_force_terminate(monkeypatch, b)
|
||||
|
||||
b._proc = _FakeProc(alive = True)
|
||||
b._start_stop_watchdog(cancel = True)
|
||||
assert _wait_until(
|
||||
lambda: calls == ["force", "final"]
|
||||
), "a cancel must escalate on the shorter cancel timeout"
|
||||
b._stop_watchdog.join(timeout = 5)
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------------
|
||||
# (c) No force-kill when the worker exits cleanly and promptly.
|
||||
# ----------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_watchdog_no_op_on_clean_quick_exit(monkeypatch):
|
||||
monkeypatch.setitem(_G, "_STOP_GRACE_S", 5.0)
|
||||
monkeypatch.setitem(_G, "_STOP_TIMEOUT_S", 10.0)
|
||||
b = TrainingBackend()
|
||||
calls = _record_force_terminate(monkeypatch, b)
|
||||
|
||||
proc = _FakeProc(alive = True)
|
||||
b._proc = proc
|
||||
b._complete_seen.set() # save done; worker is about to exit on its own
|
||||
|
||||
b._start_stop_watchdog(cancel = False)
|
||||
# Worker exits promptly, well before the grace period elapses.
|
||||
time.sleep(0.1)
|
||||
proc._alive = False
|
||||
|
||||
b._stop_watchdog.join(timeout = 5)
|
||||
assert not b._stop_watchdog.is_alive()
|
||||
assert calls == [], "a clean quick exit must not trigger force_terminate"
|
||||
|
||||
|
||||
def test_watchdog_no_op_when_worker_superseded(monkeypatch):
|
||||
# A stale watchdog from a prior run must never kill a new run's worker: once
|
||||
# self._proc is replaced, it exits silently.
|
||||
monkeypatch.setitem(_G, "_STOP_GRACE_S", 0.05)
|
||||
monkeypatch.setitem(_G, "_STOP_TIMEOUT_S", 0.05)
|
||||
b = TrainingBackend()
|
||||
calls = _record_force_terminate(monkeypatch, b)
|
||||
|
||||
old_proc = _FakeProc(alive = True)
|
||||
b._proc = old_proc
|
||||
b._complete_seen.set()
|
||||
b._start_stop_watchdog(cancel = False)
|
||||
|
||||
# A new run takes over the handle before the grace elapses.
|
||||
b._proc = _FakeProc(alive = True)
|
||||
|
||||
b._stop_watchdog.join(timeout = 5)
|
||||
assert calls == [], "watchdog must not force_terminate a superseded worker"
|
||||
|
||||
|
||||
def test_new_run_gets_its_own_watchdog(monkeypatch):
|
||||
# A stale watchdog sleeping on an old proc must not stop a new run's stop from
|
||||
# creating its own watcher.
|
||||
monkeypatch.setitem(_G, "_STOP_GRACE_S", 100.0)
|
||||
monkeypatch.setitem(_G, "_STOP_TIMEOUT_S", 100.0)
|
||||
b = TrainingBackend()
|
||||
_record_force_terminate(monkeypatch, b)
|
||||
|
||||
old_proc = _FakeProc(alive = True)
|
||||
b._proc = old_proc
|
||||
b._start_stop_watchdog(cancel = False)
|
||||
first_wd = b._stop_watchdog
|
||||
|
||||
# New run: fresh worker replaces the handle; its stop must get a new watcher
|
||||
# even though the old (superseded) watchdog is still alive.
|
||||
new_proc = _FakeProc(alive = True)
|
||||
b._proc = new_proc
|
||||
b._start_stop_watchdog(cancel = False)
|
||||
second_wd = b._stop_watchdog
|
||||
|
||||
try:
|
||||
assert first_wd.is_alive()
|
||||
assert second_wd is not first_wd, "a new run must get its own watchdog"
|
||||
assert b._stop_watchdog_proc is new_proc
|
||||
finally:
|
||||
old_proc._alive = False
|
||||
new_proc._alive = False
|
||||
first_wd.join(timeout = 5)
|
||||
second_wd.join(timeout = 5)
|
||||
|
||||
|
||||
def test_force_terminate_targets_only_captured_proc():
|
||||
# Superseded: force_terminate(target) must not touch a different current worker.
|
||||
b = TrainingBackend()
|
||||
old_proc = _FakeProc(alive = True)
|
||||
new_proc = _FakeProc(alive = True)
|
||||
b._proc = new_proc
|
||||
b.force_terminate(target_proc = old_proc)
|
||||
assert new_proc.terminated is False, "must not terminate the new run's worker"
|
||||
assert old_proc.terminated is False, "must not terminate a handle that is not current"
|
||||
|
||||
# Matching: the captured handle is the current worker, so it is terminated.
|
||||
p = _FakeProc(alive = True)
|
||||
b._proc = p
|
||||
b.force_terminate(target_proc = p)
|
||||
assert p.terminated is True
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------------
|
||||
# Post-escalation finalize leaves the parent ready for a new run.
|
||||
# ----------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_finalize_runs_even_if_force_terminate_raises(monkeypatch):
|
||||
# A wedged child can make force_terminate() raise; finalize must still run so the
|
||||
# run does not stay stuck in "Stopping...".
|
||||
monkeypatch.setitem(_G, "_STOP_GRACE_S", 0.05)
|
||||
monkeypatch.setitem(_G, "_STOP_TIMEOUT_S", 100.0)
|
||||
b = TrainingBackend()
|
||||
|
||||
def _boom(target_proc = None):
|
||||
raise RuntimeError("kill() failed on wedged child")
|
||||
|
||||
finalized: list = []
|
||||
monkeypatch.setattr(b, "force_terminate", _boom)
|
||||
monkeypatch.setattr(
|
||||
b,
|
||||
"_finalize_stopped_after_escalation",
|
||||
lambda target_proc = None, watched_job_id = None: finalized.append(True),
|
||||
)
|
||||
|
||||
b._proc = _FakeProc(alive = True)
|
||||
b._complete_seen.set()
|
||||
b._start_stop_watchdog(cancel = False)
|
||||
|
||||
assert _wait_until(
|
||||
lambda: finalized == [True]
|
||||
), "finalize must run even when force_terminate raises"
|
||||
b._stop_watchdog.join(timeout = 5)
|
||||
|
||||
|
||||
def test_finalize_after_escalation_clears_state(monkeypatch):
|
||||
# Even if the OS never reaps the wedged worker, the parent must report the run
|
||||
# stopped so the UI leaves "Stopping..." and a new run can start.
|
||||
b = TrainingBackend()
|
||||
finstop: list = []
|
||||
monkeypatch.setattr(b, "_finish_stopped_run", lambda *a: finstop.append(a))
|
||||
|
||||
b._proc = _FakeProc(alive = True) # wedged: still reports alive
|
||||
b._should_stop = True
|
||||
b.current_job_id = "job_c"
|
||||
b._db_run_created = True
|
||||
b._progress.is_training = True
|
||||
|
||||
b._finalize_stopped_after_escalation(watched_job_id = "job_c")
|
||||
|
||||
assert b._proc is None, "the wedged handle must be dropped so is_training_active clears"
|
||||
assert b._progress.is_training is False
|
||||
assert b._progress.status_message == "Training stopped."
|
||||
assert finstop and finstop[0][0] == "job_c", "the captured run must be finalized by id"
|
||||
assert b.is_training_active() is False
|
||||
|
||||
|
||||
def test_finalize_after_escalation_preserves_output_dir(monkeypatch):
|
||||
# A save-stop that already emitted "complete" has the checkpoint dir; run history
|
||||
# must record it even if the watchdog wins the finalize race against the pump.
|
||||
b = TrainingBackend()
|
||||
finstop: list = []
|
||||
monkeypatch.setattr(b, "_finish_stopped_run", lambda *a: finstop.append(a))
|
||||
|
||||
b._proc = _FakeProc(alive = True)
|
||||
b._should_stop = True
|
||||
b.current_job_id = "job_c"
|
||||
b._db_run_created = True
|
||||
b._output_dir = "/tmp/outputs/run-123"
|
||||
|
||||
b._finalize_stopped_after_escalation(watched_job_id = "job_c")
|
||||
|
||||
# _finish_stopped_run(run_id, output_dir, batch, final_step, final_loss, duration, loss_history)
|
||||
assert finstop and finstop[0][0] == "job_c"
|
||||
assert finstop[0][1] == "/tmp/outputs/run-123"
|
||||
|
||||
|
||||
def test_stop_training_starts_watchdog_only_when_worker_alive(monkeypatch):
|
||||
# No worker -> nothing to escalate; the watchdog must not spawn.
|
||||
b = TrainingBackend()
|
||||
b._proc = None
|
||||
assert b.stop_training(save = True) is True
|
||||
assert b._stop_watchdog is None
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------------
|
||||
# (d) A stale watchdog must never clobber a run that replaced its worker.
|
||||
# ----------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_finalize_after_escalation_no_ops_when_superseded(monkeypatch):
|
||||
# A /start can slip in while the watchdog force-terminates the old worker
|
||||
# (is_training_active() is False once _should_stop is set and the old proc is dead).
|
||||
# The escalation finalize must then leave the NEW run untouched, not drop its handle.
|
||||
b = TrainingBackend()
|
||||
finstop: list = []
|
||||
monkeypatch.setattr(b, "_finish_stopped_run", lambda *a: finstop.append(a))
|
||||
|
||||
old_proc = _FakeProc(alive = False) # force-terminated worker we were watching
|
||||
new_proc = _FakeProc(alive = True) # a new run already took over
|
||||
b._proc = new_proc
|
||||
b.current_job_id = "job_new"
|
||||
b._db_run_created = True
|
||||
b._progress.is_training = True
|
||||
|
||||
b._finalize_stopped_after_escalation(target_proc = old_proc)
|
||||
|
||||
assert b._proc is new_proc, "must not drop the new run's handle"
|
||||
assert b._progress.is_training is True, "must not mark the new run stopped"
|
||||
assert finstop == [], "must not finalize the new run in the DB"
|
||||
|
||||
|
||||
def test_finalize_after_escalation_runs_for_its_own_worker(monkeypatch):
|
||||
# Common case: the watched worker is still current, so finalize proceeds and
|
||||
# finalizes the captured run by id.
|
||||
b = TrainingBackend()
|
||||
finstop: list = []
|
||||
monkeypatch.setattr(b, "_finish_stopped_run", lambda *a: finstop.append(a))
|
||||
|
||||
proc = _FakeProc(alive = False)
|
||||
b._proc = proc
|
||||
b.current_job_id = "job_a"
|
||||
b._db_run_created = True
|
||||
b._progress.is_training = True
|
||||
|
||||
b._finalize_stopped_after_escalation(target_proc = proc, watched_job_id = "job_a")
|
||||
|
||||
assert b._proc is None
|
||||
assert b._progress.is_training is False
|
||||
assert finstop and finstop[0][0] == "job_a", "must finalize the captured run by id"
|
||||
|
||||
|
||||
def test_finalize_after_escalation_no_ops_on_job_change_during_startup(monkeypatch):
|
||||
# start_training updates current_job_id BEFORE it installs the new _proc, so a stale
|
||||
# watchdog can enter while _proc is still the old (dead) handle. The job-id guard must
|
||||
# catch this even though the proc-only guard would not.
|
||||
b = TrainingBackend()
|
||||
finstop: list = []
|
||||
monkeypatch.setattr(b, "_finish_stopped_run", lambda *a: finstop.append(a))
|
||||
|
||||
old_proc = _FakeProc(alive = False) # old worker, dead; new _proc not installed yet
|
||||
b._proc = old_proc # still the old handle (== target), so proc guard would pass
|
||||
b.current_job_id = "job_new" # but the new run already claimed the job id
|
||||
b._db_run_created = True
|
||||
b._progress.is_training = True
|
||||
|
||||
b._finalize_stopped_after_escalation(target_proc = old_proc, watched_job_id = "job_old")
|
||||
|
||||
assert b._proc is old_proc, "must not drop the handle during a new run's startup"
|
||||
assert b._progress.is_training is True, "must not mark the starting run stopped"
|
||||
assert finstop == [], "must not finalize while a new run is starting up"
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------------
|
||||
# (e) A later cancel (save=False) tightens an in-flight save watchdog.
|
||||
# ----------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_later_cancel_tightens_watchdog_timeout(monkeypatch):
|
||||
monkeypatch.setitem(_G, "_STOP_GRACE_S", 100.0) # never trips (no complete)
|
||||
monkeypatch.setitem(_G, "_STOP_TIMEOUT_S", 100.0) # save cap would not fire
|
||||
monkeypatch.setitem(_G, "_CANCEL_TIMEOUT_S", 0.05)
|
||||
b = TrainingBackend()
|
||||
calls = _record_force_terminate(monkeypatch, b)
|
||||
|
||||
b._proc = _FakeProc(alive = True)
|
||||
b._start_stop_watchdog(cancel = False) # started as a save-stop with the long cap
|
||||
time.sleep(0.15)
|
||||
assert calls == [], "a save-stop must not escalate on the short cancel cap yet"
|
||||
|
||||
# The user now cancels the in-flight stop: the watchdog must tighten its cap.
|
||||
b._cancel_requested = True
|
||||
assert _wait_until(
|
||||
lambda: calls == ["force", "final"]
|
||||
), "a later cancel must tighten the watchdog to the shorter cancel cap"
|
||||
b._stop_watchdog.join(timeout = 5)
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------------
|
||||
# (f) DB finalize/flush are safe when the watchdog and pump race (see Item 4).
|
||||
# ----------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _install_fake_db(monkeypatch):
|
||||
"""Stub storage.studio_db + utils.downsample so the real DB helpers run without
|
||||
SQLite. Returns the recorder dict."""
|
||||
recs = {"created": [], "finished": [], "inserted": [], "insert_ids": [], "progress_ids": []}
|
||||
fake_storage = _types.ModuleType("storage")
|
||||
fake_db = _types.ModuleType("storage.studio_db")
|
||||
fake_db.create_run = lambda **kw: recs["created"].append(kw)
|
||||
fake_db.finish_run = lambda **kw: recs["finished"].append(kw)
|
||||
fake_db.insert_metrics_batch = lambda job_id, batch: (
|
||||
recs["inserted"].extend(batch),
|
||||
recs["insert_ids"].append(job_id),
|
||||
)
|
||||
fake_db.update_run_progress = lambda **kw: recs["progress_ids"].append(kw.get("id"))
|
||||
fake_storage.studio_db = fake_db
|
||||
monkeypatch.setitem(sys.modules, "storage", fake_storage)
|
||||
monkeypatch.setitem(sys.modules, "storage.studio_db", fake_db)
|
||||
fake_ds = _types.ModuleType("utils.downsample")
|
||||
fake_ds.downsample = lambda seq, n: list(seq)[:n]
|
||||
monkeypatch.setitem(sys.modules, "utils.downsample", fake_ds)
|
||||
return recs
|
||||
|
||||
|
||||
def test_finalize_run_in_db_single_winner_under_concurrency(monkeypatch):
|
||||
# The watchdog and pump can both finalize; only one call may reach finish_run.
|
||||
recs = _install_fake_db(monkeypatch)
|
||||
b = TrainingBackend()
|
||||
b.current_job_id = "job_x"
|
||||
b._db_run_created = True
|
||||
b._run_finalized = False
|
||||
|
||||
start = threading.Barrier(8)
|
||||
|
||||
def worker():
|
||||
start.wait()
|
||||
b._finalize_run_in_db(status = "stopped")
|
||||
|
||||
threads = [threading.Thread(target = worker) for _ in range(8)]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join(timeout = 5)
|
||||
|
||||
assert len(recs["finished"]) == 1, f"finalize must run once, got {len(recs['finished'])}"
|
||||
assert b._run_finalized is True
|
||||
|
||||
|
||||
def test_finalize_run_in_db_no_ops_on_job_mismatch(monkeypatch):
|
||||
# A finalize captured for an old job must not finalize the run that replaced it.
|
||||
recs = _install_fake_db(monkeypatch)
|
||||
b = TrainingBackend()
|
||||
b.current_job_id = "job_new"
|
||||
b._db_run_created = True
|
||||
b._run_finalized = False
|
||||
|
||||
b._finalize_run_in_db(status = "stopped", expected_job_id = "job_old")
|
||||
|
||||
assert recs["finished"] == [], "a superseded job id must not finalize the current run"
|
||||
assert b._run_finalized is False
|
||||
|
||||
|
||||
def test_concurrent_flush_claims_each_metric_once(monkeypatch):
|
||||
# Concurrent flushes (pump periodic flush vs watchdog finalize flush) must not
|
||||
# double-remove or drop buffered metrics.
|
||||
recs = _install_fake_db(monkeypatch)
|
||||
b = TrainingBackend()
|
||||
b.current_job_id = "job_y"
|
||||
b._db_run_created = True
|
||||
b._metric_buffer[:] = [{"step": i} for i in range(200)]
|
||||
|
||||
start = threading.Barrier(6)
|
||||
|
||||
def worker():
|
||||
start.wait()
|
||||
for _ in range(50):
|
||||
b._flush_metrics_to_db()
|
||||
|
||||
threads = [threading.Thread(target = worker) for _ in range(6)]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join(timeout = 5)
|
||||
b._flush_metrics_to_db() # drain any remainder
|
||||
|
||||
steps = sorted(m["step"] for m in recs["inserted"])
|
||||
assert steps == list(range(200)), "each metric must be inserted exactly once"
|
||||
assert b._metric_buffer == [], "the buffer must be fully drained"
|
||||
|
||||
|
||||
def test_flush_pins_to_passed_run_id(monkeypatch):
|
||||
# A finalizer flushes to the run it captured, even if a new /start has already
|
||||
# changed current_job_id.
|
||||
recs = _install_fake_db(monkeypatch)
|
||||
b = TrainingBackend()
|
||||
b.current_job_id = "job_new" # a new run is already live
|
||||
b._db_run_created = True
|
||||
b._metric_buffer[:] = [{"step": 1}, {"step": 2}]
|
||||
|
||||
b._flush_metrics_to_db(run_id = "job_old")
|
||||
|
||||
assert recs["insert_ids"] == ["job_old"], "metrics must go to the captured run, not the new one"
|
||||
assert recs["progress_ids"] == ["job_old"]
|
||||
|
||||
|
||||
def test_finalize_uses_snapshot_run_id_across_new_run(monkeypatch):
|
||||
# If a new /start changes current_job_id after the finalize claim but before the DB
|
||||
# writes, finish_run must still target the run captured under the lock.
|
||||
recs = _install_fake_db(monkeypatch)
|
||||
b = TrainingBackend()
|
||||
b.current_job_id = "job_x"
|
||||
b._db_run_created = True
|
||||
b._run_finalized = False
|
||||
|
||||
def hijack(run_id = None):
|
||||
# Simulate a new run taking over during the flush (after the finalize claim).
|
||||
b.current_job_id = "job_y"
|
||||
|
||||
monkeypatch.setattr(b, "_flush_metrics_to_db", hijack)
|
||||
|
||||
b._finalize_run_in_db(status = "stopped", expected_job_id = "job_x")
|
||||
|
||||
assert [f["id"] for f in recs["finished"]] == [
|
||||
"job_x"
|
||||
], "finish_run must target the captured run, not the run that replaced it"
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------------
|
||||
# (g) DB row creation must not be published before the insert commits.
|
||||
# ----------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_ensure_db_run_created_publishes_only_after_insert(monkeypatch):
|
||||
# _db_run_created must stay False while create_run is in flight, so a concurrent
|
||||
# finalize can't run finish_run (an UPDATE) against a not-yet-inserted row.
|
||||
b = TrainingBackend()
|
||||
b.current_job_id = "job_z"
|
||||
b._db_config = {"model_name": "m"}
|
||||
observed: dict = {}
|
||||
|
||||
fake_storage = _types.ModuleType("storage")
|
||||
fake_db = _types.ModuleType("storage.studio_db")
|
||||
|
||||
def _create(**kw):
|
||||
observed["flag_during_create"] = b._db_run_created
|
||||
observed["in_progress_during_create"] = b._db_create_in_progress
|
||||
|
||||
fake_db.create_run = _create
|
||||
fake_storage.studio_db = fake_db
|
||||
monkeypatch.setitem(sys.modules, "storage", fake_storage)
|
||||
monkeypatch.setitem(sys.modules, "storage.studio_db", fake_db)
|
||||
|
||||
b._ensure_db_run_created()
|
||||
|
||||
assert observed["flag_during_create"] is False, "flag must not be published before insert"
|
||||
assert observed["in_progress_during_create"] is True
|
||||
assert b._db_run_created is True, "flag must be published after a successful insert"
|
||||
assert b._db_create_in_progress is False
|
||||
|
||||
|
||||
def test_ensure_db_run_created_stays_unpublished_on_failure(monkeypatch):
|
||||
# If create_run raises, neither flag stays set, so a later caller can retry.
|
||||
b = TrainingBackend()
|
||||
b.current_job_id = "job_z"
|
||||
b._db_config = {"model_name": "m"}
|
||||
|
||||
fake_storage = _types.ModuleType("storage")
|
||||
fake_db = _types.ModuleType("storage.studio_db")
|
||||
|
||||
def _boom_create(**kw):
|
||||
raise RuntimeError("insert failed")
|
||||
|
||||
fake_db.create_run = _boom_create
|
||||
fake_storage.studio_db = fake_db
|
||||
monkeypatch.setitem(sys.modules, "storage", fake_storage)
|
||||
monkeypatch.setitem(sys.modules, "storage.studio_db", fake_db)
|
||||
|
||||
b._ensure_db_run_created()
|
||||
|
||||
assert b._db_run_created is False, "a failed insert must not publish the row as created"
|
||||
assert b._db_create_in_progress is False, "the in-progress flag must be cleared on failure"
|
||||
|
||||
|
||||
def test_ensure_db_run_created_does_not_publish_for_a_new_run(monkeypatch):
|
||||
# A killed worker lets a new /start proceed while the watchdog is still creating the old
|
||||
# run's row. The stale create must not publish the backend-wide flags against the new
|
||||
# current_job_id, or the new run would skip inserting its own row.
|
||||
b = TrainingBackend()
|
||||
b.current_job_id = "job_old"
|
||||
b._db_config = {"model_name": "m"}
|
||||
b._db_run_created = False
|
||||
b._db_create_in_progress = False
|
||||
|
||||
fake_storage = _types.ModuleType("storage")
|
||||
fake_db = _types.ModuleType("storage.studio_db")
|
||||
|
||||
def _create(**kw):
|
||||
b.current_job_id = "job_new" # a new run takes over during the slow create
|
||||
|
||||
fake_db.create_run = _create
|
||||
fake_storage.studio_db = fake_db
|
||||
monkeypatch.setitem(sys.modules, "storage", fake_storage)
|
||||
monkeypatch.setitem(sys.modules, "storage.studio_db", fake_db)
|
||||
|
||||
b._ensure_db_run_created()
|
||||
|
||||
assert b._db_run_created is False, "must not publish the created flag against the new run"
|
||||
# The stale claim is left for start_training to reset, not satisfied for the new run.
|
||||
assert b._db_create_in_progress is True, "must not clear the claim once the run is not current"
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------------
|
||||
# (h) The escalation finalizes the watched run by id (so it is never left running).
|
||||
# ----------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_escalation_finalizes_watched_run_by_id_end_to_end(monkeypatch):
|
||||
# Exercise the real _finish_stopped_run against a fake DB. The watched run is finalized
|
||||
# by its captured id with its buffered metrics, so a new run that starts in the gap
|
||||
# after the backend goes idle can never leave the stopped run recorded running.
|
||||
recs = _install_fake_db(monkeypatch)
|
||||
b = TrainingBackend()
|
||||
b.current_job_id = "job_old"
|
||||
b._db_run_created = True
|
||||
b._proc = _FakeProc(alive = False)
|
||||
b._progress.is_training = True
|
||||
b._progress.step = 42
|
||||
b._metric_buffer[:] = [{"step": 41}, {"step": 42}]
|
||||
|
||||
b._finalize_stopped_after_escalation(target_proc = b._proc, watched_job_id = "job_old")
|
||||
|
||||
assert [f["id"] for f in recs["finished"]] == ["job_old"], "must finish the captured run by id"
|
||||
assert recs["finished"][0]["status"] == "stopped"
|
||||
assert recs["insert_ids"] == ["job_old"], "buffered metrics must land on the captured run"
|
||||
assert b._metric_buffer == [], "the captured batch must be drained"
|
||||
|
||||
|
||||
def test_escalation_defers_when_row_cannot_be_created_here(monkeypatch):
|
||||
# If the row does not exist and cannot be created here (no db_config, or the pump is
|
||||
# mid-create), the escalation must not claim _run_finalized or call _finish_stopped_run,
|
||||
# so the pump's create-then-finalize records the run. Parent state still clears.
|
||||
b = TrainingBackend()
|
||||
called: list = []
|
||||
monkeypatch.setattr(b, "_finish_stopped_run", lambda *a: called.append(a))
|
||||
|
||||
b._proc = _FakeProc(alive = False)
|
||||
b.current_job_id = "job_q"
|
||||
b._db_run_created = False # row not created yet
|
||||
b._db_config = None # ... and cannot be created here
|
||||
b._run_finalized = False
|
||||
b._progress.is_training = True
|
||||
|
||||
b._finalize_stopped_after_escalation(target_proc = b._proc, watched_job_id = "job_q")
|
||||
|
||||
assert called == [], "must not finalize when the row can't be established here"
|
||||
assert b._run_finalized is False, "must not claim the finalize the pump still owes"
|
||||
assert b._progress.is_training is False, "parent state must still clear so the UI unsticks"
|
||||
assert b._proc is None
|
||||
|
||||
|
||||
def test_escalation_creates_row_then_finalizes_when_start_create_failed(monkeypatch):
|
||||
# A wedged worker's pump can never finalize and would bail once _proc is dropped, so if
|
||||
# the row was never created (start-time create failed) the escalation creates it and
|
||||
# finalizes by id itself, recording the terminal state before dropping the handle.
|
||||
recs = _install_fake_db(monkeypatch)
|
||||
b = TrainingBackend()
|
||||
b.current_job_id = "job_s"
|
||||
b._db_config = {"model_name": "m"} # so _ensure_db_run_created can create the row
|
||||
b._db_run_created = False # start-time create failed
|
||||
b._proc = _FakeProc(alive = True) # wedged: still reports alive
|
||||
b._should_stop = True
|
||||
b._progress.is_training = True
|
||||
|
||||
b._finalize_stopped_after_escalation(target_proc = b._proc, watched_job_id = "job_s")
|
||||
|
||||
assert [c["id"] for c in recs["created"]] == ["job_s"], "must create the missing row"
|
||||
assert [f["id"] for f in recs["finished"]] == ["job_s"], "must finish the created row by id"
|
||||
assert b._proc is None, "handle dropped only after the terminal state is recorded"
|
||||
assert b._db_run_created is True
|
||||
|
||||
|
||||
def test_escalation_does_not_drop_a_new_runs_handle(monkeypatch):
|
||||
# If a run replaces the worker while the finalize DB write is in flight, the final _proc
|
||||
# drop must leave the new run's handle intact (re-guarded on target_proc).
|
||||
b = TrainingBackend()
|
||||
b.current_job_id = "job_old"
|
||||
b._db_run_created = True
|
||||
old_proc = _FakeProc(alive = False)
|
||||
new_proc = _FakeProc(alive = True)
|
||||
b._proc = old_proc
|
||||
|
||||
def hijack(*a):
|
||||
b._proc = new_proc # a new run takes over during the finalize
|
||||
|
||||
monkeypatch.setattr(b, "_finish_stopped_run", hijack)
|
||||
|
||||
b._finalize_stopped_after_escalation(target_proc = old_proc, watched_job_id = "job_old")
|
||||
|
||||
assert b._proc is new_proc, "must not drop the handle a new run installed during finalize"
|
||||
|
||||
|
||||
def _make_finish_raise(monkeypatch, calls):
|
||||
fn = sys.modules["storage.studio_db"]
|
||||
|
||||
def _boom(**kw):
|
||||
calls.append(kw)
|
||||
raise RuntimeError("database is locked")
|
||||
|
||||
fn.finish_run = _boom
|
||||
|
||||
|
||||
def test_finish_stopped_run_retries_then_unclaims_on_db_error(monkeypatch):
|
||||
# The watchdog is the sole finalizer once _proc is dropped, so a transient DB error is
|
||||
# retried a few times; on final failure the finalize is unclaimed (run still current).
|
||||
monkeypatch.setitem(_G, "_DB_FINALIZE_RETRY_S", 0.0)
|
||||
_install_fake_db(monkeypatch)
|
||||
tries: list = []
|
||||
_make_finish_raise(monkeypatch, tries)
|
||||
b = TrainingBackend()
|
||||
b.current_job_id = "job_r"
|
||||
b._run_finalized = True # the caller (escalation) already claimed
|
||||
|
||||
b._finish_stopped_run("job_r", None, [{"step": 1}], 1, None, None, [])
|
||||
|
||||
assert len(tries) == 3, "a transient DB error must be retried before giving up"
|
||||
assert b._run_finalized is False, "a persistent DB error must unclaim the finalize"
|
||||
|
||||
|
||||
def test_finish_stopped_run_error_leaves_new_run_untouched(monkeypatch):
|
||||
# If the watched run was superseded, a DB error must not unclaim the new run's finalize.
|
||||
monkeypatch.setitem(_G, "_DB_FINALIZE_RETRY_S", 0.0)
|
||||
_install_fake_db(monkeypatch)
|
||||
_make_finish_raise(monkeypatch, [])
|
||||
b = TrainingBackend()
|
||||
b.current_job_id = "job_new" # a new run is live
|
||||
b._run_finalized = True # the new run's flag
|
||||
|
||||
b._finish_stopped_run("job_old", None, [{"step": 1}], 1, None, None, [])
|
||||
|
||||
assert b._run_finalized is True, "must not unclaim the new run's finalize"
|
||||
Loading…
Add table
Add a link
Reference in a new issue