Studio: close switch/cancel races during model load (#6918)
Fix six race conditions when a user switches or cancels a model while a previous load or generation is still in flight, across the inference orchestrator and the /load and /unload routes: - Cancel an in-flight generation on a safetensors/MLX model switch and serialize unload with load under the inference lifecycle gate. - Cancel an in-flight load off the lifecycle gate so a Stop-loading cancel does not wait out the multi-minute load; guard the dispatched mailbox against a racing unload. - Recheck the loading marker after spawn and again after the load response before publishing, so a load cancelled mid-flight is reaped instead of going live. - Discard the loading marker before tearing the subprocess down in cancel_load, closing a spawn-after-cancel window and an orphaned compare-mode dispatcher during unload. - Match the unload target before canceling an in-flight GGUF load and add an off-gate fast path for the still-loading GGUF case. - Run the Unsloth unload off the event loop so a paused SSE stream holding _gen_lock cannot block the loop. Adds studio/backend/tests/test_orchestrator_unload_cancel.py covering the unload/cancel/switch race paths.
This commit is contained in:
parent
9dabe96786
commit
8ba46b566a
4 changed files with 1883 additions and 93 deletions
|
|
@ -45,6 +45,10 @@ _DISPATCH_STOP_TIMEOUT = 5.0
|
|||
_DISPATCH_IDLE_TIMEOUT = 30.0
|
||||
_DISPATCH_DRAIN_TIMEOUT = 5.0
|
||||
|
||||
# Max wait for a cancelled generation to release _gen_lock before unload_model
|
||||
# tears the subprocess down. Only bounds a wedged worker.
|
||||
_UNLOAD_GEN_LOCK_TIMEOUT = 15.0
|
||||
|
||||
|
||||
class InferenceOrchestrator:
|
||||
"""
|
||||
|
|
@ -60,7 +64,13 @@ class InferenceOrchestrator:
|
|||
self._cmd_queue: Any = None
|
||||
self._resp_queue: Any = None
|
||||
self._cancel_event: Any = None # mp.Event — set to cancel generation
|
||||
# Set for the whole unload; the worker never clears it (unlike _cancel_event),
|
||||
# so a generate queued behind the cancelled one is skipped, not run.
|
||||
self._drain_event: Any = None
|
||||
self._gen_lock = threading.Lock() # Serializes generation
|
||||
# Set during a switch so a generation winning the _gen_lock handoff bails
|
||||
# instead of starting on the outgoing model.
|
||||
self._unload_pending = False
|
||||
|
||||
# Dispatcher state for compare mode (adapter-controlled requests):
|
||||
# bypass _gen_lock, send commands directly, read from per-request
|
||||
|
|
@ -159,6 +169,7 @@ class InferenceOrchestrator:
|
|||
self._cmd_queue = _CTX.Queue()
|
||||
self._resp_queue = _CTX.Queue()
|
||||
self._cancel_event = _CTX.Event()
|
||||
self._drain_event = _CTX.Event()
|
||||
|
||||
self._proc = _CTX.Process(
|
||||
target = run_without_native_path_secret,
|
||||
|
|
@ -167,6 +178,7 @@ class InferenceOrchestrator:
|
|||
"cmd_queue": self._cmd_queue,
|
||||
"resp_queue": self._resp_queue,
|
||||
"cancel_event": self._cancel_event,
|
||||
"drain_event": self._drain_event,
|
||||
"config": config,
|
||||
},
|
||||
daemon = True,
|
||||
|
|
@ -228,6 +240,7 @@ class InferenceOrchestrator:
|
|||
self._cmd_queue = None
|
||||
self._resp_queue = None
|
||||
self._cancel_event = None
|
||||
self._drain_event = None
|
||||
logger.info("Inference subprocess shut down")
|
||||
|
||||
def _cleanup(self):
|
||||
|
|
@ -456,7 +469,15 @@ class InferenceOrchestrator:
|
|||
cancel ack from that same source so stale events don't leak into the
|
||||
next request.
|
||||
"""
|
||||
# Latch this stream's subprocess/queue: if a wedged worker is torn down and a
|
||||
# later load spawns a fresh one, bail rather than re-block on the new queue
|
||||
# under _gen_lock (deadlock).
|
||||
initial_proc = self._proc
|
||||
initial_resp_queue = self._resp_queue
|
||||
while True:
|
||||
if self._proc is not initial_proc or self._resp_queue is not initial_resp_queue:
|
||||
yield f"Error: {self._subprocess_crash_message(crash_context)}"
|
||||
return
|
||||
resp = read_one(read_timeout)
|
||||
if resp is None:
|
||||
# Check subprocess health
|
||||
|
|
@ -595,8 +616,24 @@ class InferenceOrchestrator:
|
|||
if not self.active_model_name:
|
||||
yield "Error: No active model"
|
||||
return
|
||||
# Latch the target model so the recheck below can detect a switch that completed
|
||||
# between _start_dispatcher and mailbox registration (mirrors the locked path's
|
||||
# expected_model check).
|
||||
expected_model = self.active_model_name
|
||||
|
||||
# Ensure dispatcher is running
|
||||
# Switch in flight (unload waiting on _gen_lock). This path bypasses the lock,
|
||||
# so without this early-out a compare request would enqueue a generate on the
|
||||
# outgoing model and delay the switch.
|
||||
if self._unload_pending:
|
||||
yield "Error: model is being unloaded"
|
||||
return
|
||||
|
||||
# Ensure the dispatcher runs. Track whether it was already running: if this call
|
||||
# starts it and then bails on a racing unload, it must stop it again (see the
|
||||
# unloading bail below).
|
||||
dispatcher_preexisting = (
|
||||
self._dispatcher_thread is not None and self._dispatcher_thread.is_alive()
|
||||
)
|
||||
self._start_dispatcher()
|
||||
|
||||
request_id = str(uuid.uuid4())
|
||||
|
|
@ -624,10 +661,42 @@ class InferenceOrchestrator:
|
|||
preserve_thinking = preserve_thinking,
|
||||
)
|
||||
|
||||
# Create mailbox BEFORE sending command
|
||||
# Create the mailbox BEFORE sending, rechecking _unload_pending under
|
||||
# _mailbox_lock: an unload sets _unload_pending before _wait_dispatcher_idle
|
||||
# reads _mailboxes under the same lock, so either the idle check sees this
|
||||
# mailbox (and tears the dispatcher down) or we see the unload and bail.
|
||||
# Registering after would orphan the mailbox and hang the compare stream forever.
|
||||
mailbox: queue.Queue = queue.Queue()
|
||||
with self._mailbox_lock:
|
||||
self._mailboxes[request_id] = mailbox
|
||||
# _unload_pending alone is not enough: an unload that ran fully since
|
||||
# _start_dispatcher clears it in its finally and stops the dispatcher, so it
|
||||
# reads False here though the dispatcher is gone and the model swapped. Also
|
||||
# bail when the active model changed or the dispatcher died: a mailbox with no
|
||||
# dispatcher to route gen_done/gen_error hangs the compare stream.
|
||||
dispatcher_alive = (
|
||||
self._dispatcher_thread is not None and self._dispatcher_thread.is_alive()
|
||||
)
|
||||
unloading = (
|
||||
self._unload_pending
|
||||
or self.active_model_name != expected_model
|
||||
or not dispatcher_alive
|
||||
)
|
||||
if not unloading:
|
||||
self._mailboxes[request_id] = mailbox
|
||||
# When bailing without a mailbox, note whether any OTHER compare request still
|
||||
# routes through the dispatcher; if none and this call started it, stop it below.
|
||||
orphaned_dispatcher = unloading and not dispatcher_preexisting and not self._mailboxes
|
||||
if unloading:
|
||||
# A racing unload can pass its _wait_dispatcher_idle() while the dispatcher was
|
||||
# stopped, then set _unload_pending. The one we just started would otherwise
|
||||
# linger with no mailboxes, race unload_model's _wait_response for the "unloaded"
|
||||
# reply off resp_queue, and drop it as unroutable -- hanging the unload 300s. Stop
|
||||
# it here so the unload stays the sole resp_queue reader. Outside _mailbox_lock:
|
||||
# _stop_dispatcher joins the dispatcher, which itself takes that lock.
|
||||
if orphaned_dispatcher:
|
||||
self._stop_dispatcher()
|
||||
yield "Error: model is being unloaded"
|
||||
return
|
||||
|
||||
try:
|
||||
self._send_cmd(cmd)
|
||||
|
|
@ -676,14 +745,18 @@ class InferenceOrchestrator:
|
|||
return
|
||||
logger.warning("Timed out draining mailbox after cancel")
|
||||
|
||||
def _wait_dispatcher_idle(self) -> None:
|
||||
def _wait_dispatcher_idle(self) -> bool:
|
||||
"""Wait for all dispatched requests to complete, then stop dispatcher.
|
||||
|
||||
Called by _generate_inner before the _gen_lock path so the dispatcher
|
||||
thread isn't competing for resp_queue reads.
|
||||
Returns True if the dispatcher was stopped (all mailboxes drained, or no
|
||||
dispatcher was running), and False if it was left running because compare
|
||||
requests were still active after _DISPATCH_IDLE_TIMEOUT.
|
||||
|
||||
Called before the _gen_lock path so the dispatcher thread isn't competing
|
||||
for resp_queue reads.
|
||||
"""
|
||||
if self._dispatcher_thread is None or not self._dispatcher_thread.is_alive():
|
||||
return
|
||||
return True
|
||||
|
||||
# Wait for all mailboxes to be emptied (dispatched requests complete)
|
||||
deadline = time.monotonic() + _DISPATCH_IDLE_TIMEOUT
|
||||
|
|
@ -704,8 +777,9 @@ class InferenceOrchestrator:
|
|||
"leaving dispatcher running for compare requests",
|
||||
len(self._mailboxes),
|
||||
)
|
||||
else:
|
||||
self._stop_dispatcher()
|
||||
return False
|
||||
self._stop_dispatcher()
|
||||
return True
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Public API — same interface as InferenceBackend
|
||||
|
|
@ -772,6 +846,19 @@ class InferenceOrchestrator:
|
|||
)
|
||||
|
||||
for attempt in range(2):
|
||||
# Stop-loading (/unload -> cancel_load) aborts a load by discarding this
|
||||
# model's loading marker. cancel_load only kills a live child; if the cancel
|
||||
# lands before any child exists (GPU placement, or between retries) there is
|
||||
# nothing to kill, and without this check the loop would spawn a worker and
|
||||
# load the model after /unload reported it unloaded. Observe removal and stop.
|
||||
if model_name not in self.loading_models:
|
||||
logger.info(
|
||||
"Load for '%s' was cancelled before spawn; not starting a worker",
|
||||
model_name,
|
||||
)
|
||||
self.active_model_name = None
|
||||
self.models.clear()
|
||||
return False
|
||||
logger.info(
|
||||
"Spawning fresh inference subprocess for '%s' "
|
||||
"(transformers %s.x, attempt %d/2%s)",
|
||||
|
|
@ -783,6 +870,22 @@ class InferenceOrchestrator:
|
|||
sub_config["disable_xet"] = disable_xet
|
||||
self._spawn_subprocess(sub_config)
|
||||
|
||||
# A cancel can land after the pre-spawn recheck but while _spawn_subprocess
|
||||
# is still creating the queues/process. cancel_load runs off the lifecycle
|
||||
# gate, so its _shutdown_subprocess can see _proc still None and no-op,
|
||||
# orphaning this fresh worker; the load would then wait for "loaded" and
|
||||
# publish a model /unload reported unloaded, over a live subprocess nothing
|
||||
# reaps. Recheck now the child exists and tear it down before publishing.
|
||||
if model_name not in self.loading_models:
|
||||
logger.info(
|
||||
"Load for '%s' was cancelled during spawn; tearing the worker down",
|
||||
model_name,
|
||||
)
|
||||
self._shutdown_subprocess(timeout = 5)
|
||||
self.active_model_name = None
|
||||
self.models.clear()
|
||||
return False
|
||||
|
||||
try:
|
||||
resp = self._wait_response("loaded")
|
||||
except DownloadStallError:
|
||||
|
|
@ -803,8 +906,31 @@ class InferenceOrchestrator:
|
|||
)
|
||||
|
||||
if resp.get("success"):
|
||||
# A cancel can land while we were parked in _wait_response above.
|
||||
# cancel_load (off the lifecycle gate) discards this model's loading
|
||||
# marker BEFORE its teardown, so a Stop-loading that fired after the
|
||||
# worker queued "loaded" (which we can still consume during cancel_load's
|
||||
# shutdown window) shows up here only as the marker's removal. Without
|
||||
# this recheck we would publish active_model_name/models for a model
|
||||
# /unload reported cancelled, over a subprocess cancel_load just killed;
|
||||
# its post-teardown re-clear cannot undo a publish that lands after it
|
||||
# returns. Observe the removal and abort; cancel_load owns teardown.
|
||||
if model_name not in self.loading_models:
|
||||
logger.info(
|
||||
"Load for '%s' was cancelled while waiting for 'loaded'; "
|
||||
"not publishing the cancelled model",
|
||||
model_name,
|
||||
)
|
||||
self.active_model_name = None
|
||||
self.models.clear()
|
||||
return False
|
||||
model_info = resp.get("model_info", {})
|
||||
self.active_model_name = model_info.get("identifier", model_name)
|
||||
# A load always spawns a fresh subprocess holding only this model, so
|
||||
# mirror that. A lingering stale name would pass unload_model's "not in
|
||||
# self.models" guard, and the worker's absent-name fallback would unload
|
||||
# its *active* model, not the already-gone one.
|
||||
self.models = {}
|
||||
self.models[self.active_model_name] = {
|
||||
"is_vision": model_info.get("is_vision", False),
|
||||
"is_lora": model_info.get("is_lora", False),
|
||||
|
|
@ -837,17 +963,65 @@ class InferenceOrchestrator:
|
|||
self.models.clear()
|
||||
raise
|
||||
|
||||
def unload_model(self, model_name: str) -> bool:
|
||||
"""Unload a model from the subprocess."""
|
||||
if model_name in self.loading_models:
|
||||
logger.info(
|
||||
"Cancelling in-flight load for model '%s' by terminating subprocess",
|
||||
def cancel_load(self, model_name: str) -> bool:
|
||||
"""Abort an in-flight load by terminating its subprocess.
|
||||
|
||||
Returns True if a load for ``model_name`` (matched case-insensitively) was
|
||||
cancelled, False if nothing was loading under that name. This only tears the
|
||||
loading subprocess down -- it sends no command to a worker -- so, unlike the
|
||||
rest of ``unload_model``, it is safe to run WITHOUT the inference lifecycle
|
||||
gate. ``/unload`` calls it off-gate so the "stop loading" button can interrupt
|
||||
a safetensors load that holds the gate for its whole (multi-minute) duration;
|
||||
a gated cancel could never preempt that load.
|
||||
"""
|
||||
target = model_name
|
||||
if target not in self.loading_models:
|
||||
target = next(
|
||||
(m for m in self.loading_models if m.lower() == model_name.lower()),
|
||||
model_name,
|
||||
)
|
||||
self._shutdown_subprocess(timeout = 0.5)
|
||||
self.loading_models.discard(model_name)
|
||||
self.active_model_name = None
|
||||
self.models.clear()
|
||||
if target not in self.loading_models:
|
||||
return False
|
||||
logger.info(
|
||||
"Cancelling in-flight load for model '%s' by terminating subprocess",
|
||||
target,
|
||||
)
|
||||
# Discard the loading marker (and clear local state) BEFORE the teardown, not
|
||||
# after. cancel_load runs off the lifecycle gate, alongside a load_model that
|
||||
# rechecks this marker before each spawn. But _shutdown_subprocess can block (~1s
|
||||
# tearing a live child down and joining the dispatcher), so clearing only after
|
||||
# leaves a window where load_model reads the marker still set, passes its pre-spawn
|
||||
# recheck, and loads the model after /unload reported it cancelled. Clear first.
|
||||
self.loading_models.discard(target)
|
||||
self.active_model_name = None
|
||||
self.models.clear()
|
||||
self._shutdown_subprocess(timeout = 0.5)
|
||||
# Clear the local mirrors again AFTER the teardown. A racing off-gate load_model
|
||||
# may still be parked in _wait_response("loaded"): its worker already queued a
|
||||
# "loaded" reply, so during the shutdown window above (the 0.5s settle before the
|
||||
# response queue is drained and nulled) that thread can consume it and repopulate
|
||||
# active_model_name/models, undoing the pre-teardown clear. _shutdown_subprocess
|
||||
# nulls the queue but not the mirrors, so without this second clear /unload reports
|
||||
# success while the backend still advertises a killed model. The nulled queue lets
|
||||
# no further "loaded" through, so re-clearing here wipes any repopulation.
|
||||
self.active_model_name = None
|
||||
self.models.clear()
|
||||
return True
|
||||
|
||||
def unload_model(self, model_name: str) -> bool:
|
||||
"""Unload a model from the subprocess."""
|
||||
# active_model_name can differ in case from the client's raw /unload name (the
|
||||
# load path canonicalizes casing). Match case-insensitively and use the canonical
|
||||
# spelling so the guard, unload command, and cleanup below hit the loaded model.
|
||||
if (
|
||||
self.active_model_name is not None
|
||||
and model_name != self.active_model_name
|
||||
and model_name.lower() == self.active_model_name.lower()
|
||||
):
|
||||
model_name = self.active_model_name
|
||||
# In-flight load: tear its subprocess down (shared loading-cancel logic; no
|
||||
# worker command sent).
|
||||
if self.cancel_load(model_name):
|
||||
return True
|
||||
|
||||
if not self._ensure_subprocess_alive():
|
||||
|
|
@ -857,30 +1031,85 @@ class InferenceOrchestrator:
|
|||
self.active_model_name = None
|
||||
return True
|
||||
|
||||
try:
|
||||
self._send_cmd(
|
||||
{
|
||||
"type": "unload",
|
||||
"model_name": model_name,
|
||||
}
|
||||
)
|
||||
resp = self._wait_response("unloaded")
|
||||
|
||||
# Update local state
|
||||
# Nothing loaded under this name: don't unload a stale model. The worker falls
|
||||
# back to unloading its *active* model when the name is absent, so a stale unload
|
||||
# (lost a race to a concurrent load) would hit the wrong one.
|
||||
if model_name != self.active_model_name and model_name not in self.models:
|
||||
self.models.pop(model_name, None)
|
||||
if self.active_model_name == model_name:
|
||||
self.active_model_name = None
|
||||
|
||||
logger.info("Model '%s' unloaded from subprocess", model_name)
|
||||
return True
|
||||
|
||||
except Exception as exc:
|
||||
logger.error("Error unloading model '%s': %s", model_name, exc)
|
||||
# Clear local state anyway
|
||||
self.models.pop(model_name, None)
|
||||
if self.active_model_name == model_name:
|
||||
self.active_model_name = None
|
||||
return False
|
||||
# The subprocess runs commands sequentially, so a bare unload queues behind a
|
||||
# running generate (a 2-3 min hang). Cancel first (via the mp.Event the worker
|
||||
# polls each token), then take _gen_lock as sole resp_queue reader (like GGUF).
|
||||
self._unload_pending = True
|
||||
# Cancelling only the running generation isn't enough: the worker clears
|
||||
# cancel_event at each generate start, so a queued one would clear it and run the
|
||||
# outgoing model to completion. drain_event, never cleared, makes any generate
|
||||
# dequeued during the unload skip.
|
||||
if self._drain_event is not None:
|
||||
self._drain_event.set()
|
||||
try:
|
||||
self._cancel_generation()
|
||||
acquired = self._gen_lock.acquire(timeout = _UNLOAD_GEN_LOCK_TIMEOUT)
|
||||
if not acquired:
|
||||
# Wedged worker: tear the subprocess down to free the GPU (next load respawns).
|
||||
logger.warning(
|
||||
"Unload: generation did not yield %.1fs after cancel; "
|
||||
"shutting the inference subprocess down to free the model",
|
||||
_UNLOAD_GEN_LOCK_TIMEOUT,
|
||||
)
|
||||
self._shutdown_subprocess(timeout = 5)
|
||||
self.models.pop(model_name, None)
|
||||
if self.active_model_name == model_name:
|
||||
self.active_model_name = None
|
||||
return True
|
||||
|
||||
try:
|
||||
# Stop the compare-mode dispatcher so it can't consume the "unloaded" reply
|
||||
# off resp_queue before we do. A dispatched generation bypasses _gen_lock, so
|
||||
# a wedged one slips past the acquire above; if the dispatcher is still active
|
||||
# it owns resp_queue and the queued unload hangs _wait_response behind the
|
||||
# stuck generate. Mirror the wedged locked path: tear the subprocess down.
|
||||
if not self._wait_dispatcher_idle():
|
||||
logger.warning(
|
||||
"Unload: compare-mode dispatcher still active after idle "
|
||||
"wait; shutting the inference subprocess down to free the model"
|
||||
)
|
||||
self._shutdown_subprocess(timeout = 5)
|
||||
self.models.pop(model_name, None)
|
||||
if self.active_model_name == model_name:
|
||||
self.active_model_name = None
|
||||
return True
|
||||
# Drop stale tokens so they can't be read as the unload reply.
|
||||
self._drain_queue()
|
||||
self._send_cmd(
|
||||
{
|
||||
"type": "unload",
|
||||
"model_name": model_name,
|
||||
}
|
||||
)
|
||||
self._wait_response("unloaded")
|
||||
|
||||
self.models.pop(model_name, None)
|
||||
if self.active_model_name == model_name:
|
||||
self.active_model_name = None
|
||||
|
||||
logger.info("Model '%s' unloaded from subprocess", model_name)
|
||||
return True
|
||||
|
||||
except Exception as exc:
|
||||
logger.error("Error unloading model '%s': %s", model_name, exc)
|
||||
# Clear local state anyway
|
||||
self.models.pop(model_name, None)
|
||||
if self.active_model_name == model_name:
|
||||
self.active_model_name = None
|
||||
return False
|
||||
finally:
|
||||
self._gen_lock.release()
|
||||
finally:
|
||||
self._unload_pending = False
|
||||
if self._drain_event is not None:
|
||||
self._drain_event.clear()
|
||||
|
||||
def generate_chat_response(
|
||||
self,
|
||||
|
|
@ -1068,6 +1297,7 @@ class InferenceOrchestrator:
|
|||
if not self.active_model_name:
|
||||
yield "Error: No active model"
|
||||
return
|
||||
expected_model = self.active_model_name
|
||||
|
||||
# Drain any prior compare-mode dispatcher so we can read resp_queue.
|
||||
self._wait_dispatcher_idle()
|
||||
|
|
@ -1076,6 +1306,14 @@ class InferenceOrchestrator:
|
|||
# consume and drop each other's token events. Hold _gen_lock across the
|
||||
# cmd build + send + whole stream so we stay the sole resp_queue reader.
|
||||
with self._gen_lock:
|
||||
# Recheck under the lock: an unload we raced may have cleared/swapped the model.
|
||||
# _unload_pending resets after the lock releases, so it can read False by now;
|
||||
# the active-model check catches that handoff and a reload that swapped models,
|
||||
# so we never generate on the wrong one.
|
||||
if self._unload_pending or self.active_model_name != expected_model:
|
||||
# Won the lock handoff during a switch; don't start on the outgoing model.
|
||||
yield "Error: model is being unloaded"
|
||||
return
|
||||
request_id = str(uuid.uuid4())
|
||||
image_b64 = self._pil_to_base64(image) if image is not None else None
|
||||
cmd = self._build_generate_cmd(
|
||||
|
|
@ -1143,53 +1381,62 @@ class InferenceOrchestrator:
|
|||
raise RuntimeError("Inference subprocess is not running")
|
||||
if not self.active_model_name:
|
||||
raise RuntimeError("No active model")
|
||||
expected_model = self.active_model_name
|
||||
|
||||
request_id = str(uuid.uuid4())
|
||||
# Serialize under _gen_lock (sole resp_queue reader) and refuse to start on the
|
||||
# outgoing model once an unload is pending, like the text and audio-input paths.
|
||||
# Without this a concurrent /audio/generate could run TTS on a model being switched.
|
||||
with self._gen_lock:
|
||||
# Recheck under the lock (see _generate_inner): a raced unload/switch may have
|
||||
# cleared or swapped the model while we waited.
|
||||
if self._unload_pending or self.active_model_name != expected_model:
|
||||
raise RuntimeError("model is being unloaded")
|
||||
|
||||
cmd = {
|
||||
"type": "generate_audio",
|
||||
"request_id": request_id,
|
||||
"text": text,
|
||||
"temperature": temperature,
|
||||
"top_p": top_p,
|
||||
"top_k": top_k,
|
||||
"min_p": min_p,
|
||||
"max_new_tokens": max_new_tokens,
|
||||
"repetition_penalty": repetition_penalty,
|
||||
}
|
||||
if use_adapter is not None:
|
||||
cmd["use_adapter"] = use_adapter
|
||||
request_id = str(uuid.uuid4())
|
||||
|
||||
self._send_cmd(cmd)
|
||||
cmd = {
|
||||
"type": "generate_audio",
|
||||
"request_id": request_id,
|
||||
"text": text,
|
||||
"temperature": temperature,
|
||||
"top_p": top_p,
|
||||
"top_k": top_k,
|
||||
"min_p": min_p,
|
||||
"max_new_tokens": max_new_tokens,
|
||||
"repetition_penalty": repetition_penalty,
|
||||
}
|
||||
if use_adapter is not None:
|
||||
cmd["use_adapter"] = use_adapter
|
||||
|
||||
# Wait for audio_done or audio_error
|
||||
deadline = time.monotonic() + 120.0
|
||||
while time.monotonic() < deadline:
|
||||
remaining = max(0.1, deadline - time.monotonic())
|
||||
resp = self._read_resp(timeout = min(remaining, 1.0))
|
||||
self._send_cmd(cmd)
|
||||
|
||||
if resp is None:
|
||||
if not self._ensure_subprocess_alive():
|
||||
raise RuntimeError(self._subprocess_crash_message("audio generation"))
|
||||
continue
|
||||
deadline = time.monotonic() + 120.0
|
||||
while time.monotonic() < deadline:
|
||||
remaining = max(0.1, deadline - time.monotonic())
|
||||
resp = self._read_resp(timeout = min(remaining, 1.0))
|
||||
|
||||
rtype = resp.get("type", "")
|
||||
if resp is None:
|
||||
if not self._ensure_subprocess_alive():
|
||||
raise RuntimeError(self._subprocess_crash_message("audio generation"))
|
||||
continue
|
||||
|
||||
if rtype == "audio_done":
|
||||
wav_bytes = base64.b64decode(resp["wav_base64"])
|
||||
sample_rate = resp["sample_rate"]
|
||||
return wav_bytes, sample_rate
|
||||
rtype = resp.get("type", "")
|
||||
|
||||
if rtype == "audio_error":
|
||||
raise RuntimeError(resp.get("error", "Audio generation failed"))
|
||||
if rtype == "audio_done":
|
||||
wav_bytes = base64.b64decode(resp["wav_base64"])
|
||||
sample_rate = resp["sample_rate"]
|
||||
return wav_bytes, sample_rate
|
||||
|
||||
if rtype == "error":
|
||||
raise RuntimeError(resp.get("error", "Unknown error"))
|
||||
if rtype == "audio_error":
|
||||
raise RuntimeError(resp.get("error", "Audio generation failed"))
|
||||
|
||||
if rtype == "status":
|
||||
continue
|
||||
if rtype == "error":
|
||||
raise RuntimeError(resp.get("error", "Unknown error"))
|
||||
|
||||
raise RuntimeError("Timeout waiting for audio generation (120s)")
|
||||
if rtype == "status":
|
||||
continue
|
||||
|
||||
raise RuntimeError("Timeout waiting for audio generation (120s)")
|
||||
|
||||
def generate_whisper_response(
|
||||
self,
|
||||
|
|
@ -1254,8 +1501,15 @@ class InferenceOrchestrator:
|
|||
if not self.active_model_name:
|
||||
yield "Error: No active model"
|
||||
return
|
||||
expected_model = self.active_model_name
|
||||
|
||||
with self._gen_lock:
|
||||
# Recheck under the lock (see _generate_inner): a raced unload/switch may have
|
||||
# cleared or swapped the model while we waited.
|
||||
if self._unload_pending or self.active_model_name != expected_model:
|
||||
# Won the lock handoff during a switch; don't start on the outgoing model.
|
||||
yield "Error: model is being unloaded"
|
||||
return
|
||||
request_id = str(uuid.uuid4())
|
||||
|
||||
# numpy array -> list for mp.Queue serialization
|
||||
|
|
|
|||
|
|
@ -406,6 +406,32 @@ def _handle_load(backend, config: dict, resp_queue: Any) -> None:
|
|||
)
|
||||
|
||||
|
||||
def _drain_skip_generate(cmd: dict, resp_queue: Any, drain_event) -> bool:
|
||||
"""Skip a generate queued behind a cancelled one during an unload.
|
||||
|
||||
The parent sets ``drain_event`` for the whole unload. Because the parent's
|
||||
per-token ``cancel_event`` is cleared at the start of every generate, a cancel
|
||||
set while this generate was still queued would otherwise be lost when it is
|
||||
dequeued. If the drain is in effect, emit an immediate (empty) ``gen_done`` so
|
||||
the parent's stream/mailbox drains fast and the switch stays fast, and report
|
||||
the generate was skipped so the caller does not clear the cancel or run it.
|
||||
"""
|
||||
if drain_event is None or not drain_event.is_set():
|
||||
return False
|
||||
request_id = cmd.get("request_id", "")
|
||||
logger.info("Skipping generate for request %s: unload draining", request_id)
|
||||
_send_response(
|
||||
resp_queue,
|
||||
{
|
||||
"type": "gen_done",
|
||||
"request_id": request_id,
|
||||
"cancelled": True,
|
||||
"stats": None,
|
||||
},
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
def _handle_generate(backend, cmd: dict, resp_queue: Any, cancel_event) -> None:
|
||||
"""Handle a generate command: stream tokens back via resp_queue.
|
||||
|
||||
|
|
@ -632,7 +658,14 @@ def _handle_unload(backend, cmd: dict, resp_queue: Any) -> None:
|
|||
)
|
||||
|
||||
|
||||
def run_inference_process(*, cmd_queue: Any, resp_queue: Any, cancel_event, config: dict) -> None:
|
||||
def run_inference_process(
|
||||
*,
|
||||
cmd_queue: Any,
|
||||
resp_queue: Any,
|
||||
cancel_event,
|
||||
config: dict,
|
||||
drain_event = None,
|
||||
) -> None:
|
||||
"""Subprocess entrypoint. Persistent — runs the command loop until shutdown.
|
||||
|
||||
Args:
|
||||
|
|
@ -640,6 +673,10 @@ def run_inference_process(*, cmd_queue: Any, resp_queue: Any, cancel_event, conf
|
|||
resp_queue: mp.Queue for sending responses to parent.
|
||||
cancel_event: mp.Event the parent sets to cancel generation.
|
||||
config: Initial configuration dict with model info.
|
||||
drain_event: mp.Event the parent sets for the duration of an unload. Unlike
|
||||
cancel_event (cleared at the start of every generate), it is never cleared
|
||||
here, so a generate still queued behind a cancelled one is skipped rather
|
||||
than run — the cancel survives the queue handoff.
|
||||
"""
|
||||
os.environ["TOKENIZERS_PARALLELISM"] = "false"
|
||||
os.environ["PYTHONWARNINGS"] = "ignore" # Suppress warnings at C-level before imports
|
||||
|
|
@ -715,7 +752,16 @@ def run_inference_process(*, cmd_queue: Any, resp_queue: Any, cancel_event, conf
|
|||
cmd_type = cmd.get("type", "")
|
||||
try:
|
||||
if cmd_type == "generate":
|
||||
if _drain_skip_generate(cmd, resp_queue, drain_event):
|
||||
continue
|
||||
cancel_event.clear()
|
||||
# Re-check the drain after clearing: the parent sets drain_event
|
||||
# then cancel_event for an unload, so if that pair landed between
|
||||
# the check above and this clear, the clear just erased the unload's
|
||||
# cancel. Skip here so the outgoing model is not run to completion,
|
||||
# which would stall the switch until the dispatcher idle-timeout.
|
||||
if _drain_skip_generate(cmd, resp_queue, drain_event):
|
||||
continue
|
||||
_handle_generate(backend, cmd, resp_queue, cancel_event)
|
||||
elif cmd_type == "load":
|
||||
if backend.active_model_name:
|
||||
|
|
@ -918,7 +964,16 @@ def run_inference_process(*, cmd_queue: Any, resp_queue: Any, cancel_event, conf
|
|||
|
||||
try:
|
||||
if cmd_type == "generate":
|
||||
if _drain_skip_generate(cmd, resp_queue, drain_event):
|
||||
continue
|
||||
cancel_event.clear()
|
||||
# Re-check the drain after clearing: the parent sets drain_event then
|
||||
# cancel_event for an unload, so if that pair landed between the check
|
||||
# above and this clear, the clear just erased the unload's cancel. Skip
|
||||
# here so the outgoing model is not run to completion, which would stall
|
||||
# the switch until the dispatcher idle-timeout tears the subprocess down.
|
||||
if _drain_skip_generate(cmd, resp_queue, drain_event):
|
||||
continue
|
||||
_handle_generate(backend, cmd, resp_queue, cancel_event)
|
||||
|
||||
elif cmd_type == "load":
|
||||
|
|
|
|||
|
|
@ -3394,12 +3394,15 @@ async def _load_model_impl(request: LoadRequest, fastapi_request: Request, curre
|
|||
llama_backend = get_llama_cpp_backend()
|
||||
unsloth_backend = get_inference_backend()
|
||||
|
||||
# Unload any active Unsloth model to free VRAM
|
||||
# Unload any active Unsloth model to free VRAM (off the event loop:
|
||||
# unload takes _gen_lock and can wait on an in-flight stream).
|
||||
if unsloth_backend.active_model_name:
|
||||
logger.info(
|
||||
f"Unloading Unsloth model '{unsloth_backend.active_model_name}' before loading GGUF"
|
||||
)
|
||||
unsloth_backend.unload_model(unsloth_backend.active_model_name)
|
||||
await asyncio.to_thread(
|
||||
unsloth_backend.unload_model, unsloth_backend.active_model_name
|
||||
)
|
||||
|
||||
# Inherit llama_extra_args from the previous load when the request
|
||||
# omits the field (the chat-settings Apply path doesn't round-trip
|
||||
|
|
@ -4063,28 +4066,76 @@ async def unload_model(request: UnloadRequest, current_subject: str = Depends(ge
|
|||
# A deliberate unload means "stay unloaded": drop any idle reload stash so the
|
||||
# next /v1 request can't resurrect this model. The idle loop unloads via the
|
||||
# backend directly (not this route), so clearing here never fights keep-warm.
|
||||
from core.inference.llama_keepwarm import note_model_unloaded
|
||||
from core.inference.llama_keepwarm import inference_lifecycle_gate, note_model_unloaded
|
||||
try:
|
||||
# Check if the GGUF backend has this model loaded or is loading it.
|
||||
llama_backend = get_llama_cpp_backend()
|
||||
if llama_backend.is_active and (
|
||||
llama_backend.model_identifier == request.model_path
|
||||
or is_registered_native_path_label(llama_backend.model_identifier, request.model_path)
|
||||
or not llama_backend.is_loaded
|
||||
# "Stop loading" (frontend cancelLoading -> /unload) must abort a still-loading
|
||||
# model promptly. /load holds the lifecycle gate for the whole (multi-minute) load,
|
||||
# so gating first would make the cancel wait it out. cancel_load only tears the
|
||||
# loading subprocess down (no unload command), so it is safe off-gate.
|
||||
backend = get_inference_backend()
|
||||
loading = getattr(backend, "get_loading_model", lambda: None)()
|
||||
if (
|
||||
loading is not None
|
||||
and hasattr(backend, "cancel_load")
|
||||
and (request.model_path == loading or request.model_path.lower() == loading.lower())
|
||||
):
|
||||
# A manual unload is a deliberate user action: tear down now even if a
|
||||
# request is mid-stream (only the automatic idle loop defers to it).
|
||||
llama_backend.unload_model()
|
||||
if await asyncio.to_thread(backend.cancel_load, request.model_path):
|
||||
note_model_unloaded()
|
||||
logger.info(f"Cancelled in-flight load: {request.model_path}")
|
||||
return UnloadResponse(status = "unloaded", model = request.model_path)
|
||||
|
||||
# Same "stop loading" fast path for a still-loading GGUF (llama-server spawned,
|
||||
# health check not yet passed). A gated unload would wait out the multi-minute
|
||||
# load; unload_model() sets the cancel_event load_model polls off its own lock and
|
||||
# kills the child, sending no worker command, so it is safe off-gate like
|
||||
# cancel_load. The gated GGUF branch below handles the already-loaded case. Gate on
|
||||
# the loading model (identifier or native label): the single llama-server loads one
|
||||
# GGUF at a time, so an unload for a different model must not cancel this load.
|
||||
llama_backend = get_llama_cpp_backend()
|
||||
if (
|
||||
llama_backend.is_active
|
||||
and not llama_backend.is_loaded
|
||||
and (
|
||||
llama_backend.model_identifier == request.model_path
|
||||
or is_registered_native_path_label(
|
||||
llama_backend.model_identifier, request.model_path
|
||||
)
|
||||
)
|
||||
):
|
||||
await asyncio.to_thread(llama_backend.unload_model)
|
||||
note_model_unloaded()
|
||||
logger.info(f"Unloaded GGUF model: {request.model_path}")
|
||||
logger.info(f"Cancelled in-flight GGUF load: {request.model_path}")
|
||||
return UnloadResponse(status = "unloaded", model = request.model_path)
|
||||
|
||||
# Otherwise, unload from Unsloth backend
|
||||
backend = get_inference_backend()
|
||||
backend.unload_model(request.model_path)
|
||||
note_model_unloaded()
|
||||
logger.info(f"Unloaded model: {request.model_path}")
|
||||
return UnloadResponse(status = "unloaded", model = request.model_path)
|
||||
# Serialize with /load under the same lifecycle gate: the Unsloth unload now runs
|
||||
# off the event loop (asyncio.to_thread), so without this a concurrent /load could
|
||||
# swap in a fresh subprocess mid-unload and the unload command would land on the
|
||||
# new worker. The gate makes load and unload exclusive.
|
||||
async with inference_lifecycle_gate():
|
||||
# Check if the GGUF backend has this model loaded or is loading it.
|
||||
llama_backend = get_llama_cpp_backend()
|
||||
if llama_backend.is_active and (
|
||||
llama_backend.model_identifier == request.model_path
|
||||
or is_registered_native_path_label(
|
||||
llama_backend.model_identifier, request.model_path
|
||||
)
|
||||
or not llama_backend.is_loaded
|
||||
):
|
||||
# A manual unload is a deliberate user action: tear down now even if a
|
||||
# request is mid-stream (only the automatic idle loop defers to it).
|
||||
llama_backend.unload_model()
|
||||
note_model_unloaded()
|
||||
logger.info(f"Unloaded GGUF model: {request.model_path}")
|
||||
return UnloadResponse(status = "unloaded", model = request.model_path)
|
||||
|
||||
# Unload from Unsloth backend off the event loop: unload takes _gen_lock, which
|
||||
# a slow SSE stream paused between tokens still holds, so a sync call would block
|
||||
# the loop that drives the stream's next token and the lock release.
|
||||
backend = get_inference_backend()
|
||||
await asyncio.to_thread(backend.unload_model, request.model_path)
|
||||
note_model_unloaded()
|
||||
logger.info(f"Unloaded model: {request.model_path}")
|
||||
return UnloadResponse(status = "unloaded", model = request.model_path)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error unloading model: {e}", exc_info = True)
|
||||
|
|
|
|||
1430
studio/backend/tests/test_orchestrator_unload_cancel.py
Normal file
1430
studio/backend/tests/test_orchestrator_unload_cancel.py
Normal file
File diff suppressed because it is too large
Load diff
Loading…
Add table
Add a link
Reference in a new issue