fix(model): tie abandoned-load verdict to the timed-out worker and clear it when the loader exits (#2394)

Review findings on #2467: a late-failing abandoned load left _load_abandoned set;
a queued caller's timeout could brand another caller's live load as stuck; the
warm-path test had no assertions.

Co-Authored-By: Claude Sonnet 5.5 <noreply@anthropic.com>
This commit is contained in:
Palash Debnath
2026-10-01 18:43:56 +05:30
co-authored by Claude Sonnet 5.5
parent b00147a756
commit e59bb040ee
2 changed files with 171 additions and 29 deletions
+64 -19
View File
@@ -2368,17 +2368,36 @@ class ModelLoadAbandoned(RuntimeError):
"""
#: Set by a cold load for as long as it holds ``_model_load_thread_lock`` —
#: i.e. while it is inside the native loader and cannot be cancelled.
#: #2394: this is what lets a later caller tell "another load is still working,
#: wait your turn" (fine, and it will return the model) apart from "a load was
#: already abandoned and may never finish" (a restart is the only way out).
#: An Event, not a flag: the wait must be able to distinguish the two without
#: polling.
_load_in_progress = threading.Event()
class _LoadTicket:
"""One caller's claim on a cold load, so a timeout blames the RIGHT load.
#2394: ``_load_model_with_timeout()`` must mark a load abandoned only when
*its own* worker is the one stuck in the native loader. A process-global
"a load is in progress" flag cannot tell that apart from "my worker is
still queued behind somebody else's load" — it would brand a healthy,
progressing load as abandoned and make every later cold load fail fast.
``loading`` is set by the worker once it holds the load lock and is about to
enter the native loader. ``gave_up`` is set by the caller when its deadline
passes. Each side sets its own event BEFORE reading the other's, so either
the caller sees ``loading`` (and abandons the running load) or the worker
sees ``gave_up`` (and declines to start one nobody awaits) — never neither.
"""
__slots__ = ("loading", "gave_up", "done")
def __init__(self) -> None:
self.loading = threading.Event()
self.gave_up = threading.Event()
#: Set by the worker as it leaves the critical section, so a caller
#: that times out a hair after the load finished does not brand a
#: finished load as stuck.
self.done = threading.Event()
#: Set when a load is GIVEN UP ON at its deadline while still running (#2394),
#: cleared by that load's own ``finally`` if it ever finishes. Read before
#: cleared by that load's own ``finally`` when it exits — success or failure —
#: because from then on the lock is free and nothing is stuck. Read before
#: waiting on the load lock: a caller that arrives after an abandonment must
#: fail immediately, because the lock it would queue behind is held by a
#: loader nobody is waiting on any more. Without this, a retry re-waits the
@@ -2882,7 +2901,9 @@ def _reset_gpu_pool() -> None:
_gpu_pool_singleton.reset()
def _load_model_exclusive(timeout: float | None = None):
def _load_model_exclusive(
timeout: float | None = None, ticket: _LoadTicket | None = None
):
"""The single cold-load leaf: reclaim, load, publish — under ONE lock.
#2394. Both cold routes reach the native `from_pretrained` through here:
@@ -2947,21 +2968,35 @@ def _load_model_exclusive(timeout: float | None = None):
"Logs → Restart backend) if it does not finish on its own, then try "
"again."
)
# Reaching the lock clears the "stuck" verdict: whoever held it has
# finished, so the backend is healthy again and needs no restart. A load
# that merely looked stuck from outside is not permanently disarmed.
# Reaching the lock means whoever held it has finished, so any "stuck"
# verdict is stale and the backend needs no restart.
_load_abandoned.clear()
_load_in_progress.set()
try:
if ticket is not None:
ticket.loading.set()
if ticket.gave_up.is_set():
# Our caller timed out while we were queued; it already told the
# user and moved on. Starting a native load nobody awaits would
# only pin this lock for the whole load.
raise ModelLoadAbandoned(
"This model load request had already timed out; not "
"starting a load nobody is waiting for."
)
if model is not None:
# A load that finished while we waited. It also clears
# ``_load_in_progress`` on its way out.
# A load that finished while we waited.
return model
_make_room_before_tts_load()
model = _load_model_sync()
return model
finally:
_load_in_progress.clear()
# Whether this load succeeded, failed or was an abandoned one finishing
# late, it is no longer inside the native loader: the lock is about to
# be free, so nothing is stuck any more. Without this a failed
# abandoned load would leave every later cold load refused until the
# backend restarts, despite a free lock.
if ticket is not None:
ticket.done.set()
_load_abandoned.clear()
_model_load_thread_lock.release()
@@ -2990,10 +3025,11 @@ async def _load_model_with_timeout():
"""
loop = asyncio.get_running_loop()
timeout = _model_load_timeout()
ticket = _LoadTicket()
try:
return await asyncio.wait_for(
loop.run_in_executor(
_get_gpu_pool(), _load_model_exclusive, timeout
_get_gpu_pool(), _load_model_exclusive, timeout, ticket
),
timeout=timeout,
)
@@ -3008,8 +3044,17 @@ async def _load_model_with_timeout():
# rather than re-waiting the full budget behind this loader — the
# inline route passes no deadline of its own, so without the flag it
# would sit on the lock indefinitely (#2394).
if _load_in_progress.is_set():
ticket.gave_up.set()
stuck = False
if ticket.loading.is_set():
_load_abandoned.set()
# Re-check after publishing: if the worker finished in between, its
# own clear may already have run and ours would be stale.
if ticket.done.is_set():
_load_abandoned.clear()
else:
stuck = True
if stuck:
logger.error(
"Model load exceeded %ss and is still running in a worker that "
"cannot be interrupted; the load lock stays held until it "
+107 -10
View File
@@ -255,6 +255,19 @@ def test_a_lone_cold_load_still_loads_and_publishes(mm, monkeypatch):
def test_a_warm_model_never_takes_the_load_lock(mm, monkeypatch):
"""The guard must not serialise the hot path behind a load lock."""
probe = _ConcurrentLoadProbe()
resident = object()
monkeypatch.setattr(mm, "model", resident, raising=False)
monkeypatch.setattr(mm, "_load_model_sync", probe, raising=False)
async def _no_heal():
return None
monkeypatch.setattr(mm, "_heal_tts_placement", _no_heal, raising=False)
monkeypatch.setattr(mm, "make_room_before_generate", lambda: None, raising=False)
assert asyncio.run(mm.get_model()) is resident
assert probe.calls == 0, "a resident model must not re-enter the cold load"
def test_a_retry_after_a_timed_out_load_must_not_block_forever(mm, monkeypatch):
@@ -345,17 +358,101 @@ def _run_capture(mm):
except BaseException as exc: # noqa: BLE001 — the failure IS the subject
return exc
probe = _ConcurrentLoadProbe()
resident = object()
monkeypatch.setattr(mm, "model", resident, raising=False)
monkeypatch.setattr(mm, "_load_model_sync", probe, raising=False)
async def _no_heal():
return None
def _wedge_setup(mm, monkeypatch, loader):
from concurrent.futures import ThreadPoolExecutor
monkeypatch.setattr(mm, "_heal_tts_placement", _no_heal, raising=False)
monkeypatch.setattr(mm, "make_room_before_generate", lambda: None, raising=False)
monkeypatch.setattr(mm, "model", None, raising=False)
monkeypatch.setattr(mm, "_model_lock", asyncio.Lock(), raising=False)
monkeypatch.setattr(mm, "_model_load_thread_lock", threading.RLock(), raising=False)
monkeypatch.setattr(mm, "_load_model_sync", loader, raising=False)
monkeypatch.setattr(mm, "_make_room_before_tts_load", lambda: None, raising=False)
monkeypatch.setattr(mm, "_model_load_timeout", lambda: 0.3, raising=False)
mm._load_abandoned.clear()
pool = ThreadPoolExecutor(max_workers=2, thread_name_prefix="test-pool")
monkeypatch.setattr(mm, "_get_gpu_pool", lambda: pool, raising=False)
return pool
assert asyncio.run(mm.get_model()) is resident
assert probe.calls == 0, "a resident model must not re-enter the cold load"
def test_an_abandoned_load_that_later_fails_does_not_strand_the_backend(mm, monkeypatch):
"""Review (Greptile P1): the abandoned verdict must die with its loader.
A timed-out load that eventually FAILS releases the lock but used to leave
``_load_abandoned`` set, so every later cold load was refused until restart
although nothing was stuck any more.
"""
wedged, release = threading.Event(), threading.Event()
calls = []
def _loader():
calls.append(1)
if len(calls) == 1:
wedged.set()
release.wait(30)
raise RuntimeError("late native load failure")
return object()
pool = _wedge_setup(mm, monkeypatch, _loader)
try:
first = _run_capture(mm)
assert isinstance(first, mm.ModelLoadAbandoned), first
assert wedged.is_set()
assert mm._load_abandoned.is_set()
release.set() # the wedged loader now dies on its own
deadline = time.monotonic() + 15
while mm._load_abandoned.is_set() and time.monotonic() < deadline:
time.sleep(0.01)
assert not mm._load_abandoned.is_set(), (
"the abandoned flag outlived its loader; the backend stays refused"
)
monkeypatch.setattr(mm, "_get_gpu_pool", lambda: pool, raising=False)
retry = _run_capture(mm)
assert not isinstance(retry, BaseException), retry
finally:
release.set()
monkeypatch.setattr(mm, "model", None, raising=False)
pool.shutdown(wait=False)
def test_a_caller_queued_behind_another_load_does_not_brand_it_abandoned(mm, monkeypatch):
"""Review (Greptile P1): blame the worker that timed out, not whichever
load happens to be running.
Caller B times out while still queued behind caller A's healthy, in-flight
load. A's load must not be marked abandoned.
"""
in_native, release = threading.Event(), threading.Event()
published = object()
def _loader():
in_native.set()
release.wait(30)
return published
pool = _wedge_setup(mm, monkeypatch, _loader)
monkeypatch.setattr(mm, "_model_load_timeout", lambda: 5.0, raising=False)
try:
# A: an inline-route load that holds the lock with no deadline.
a_result: dict[str, object] = {}
a = threading.Thread(
target=lambda: a_result.setdefault("r", mm._load_model_exclusive()),
daemon=True,
)
a.start()
assert in_native.wait(10)
# B: a pool route whose deadline (shortened) passes while queued.
monkeypatch.setattr(mm, "_model_load_timeout", lambda: 0.3, raising=False)
b = _run_capture(mm)
assert isinstance(b, RuntimeError), b
assert not mm._load_abandoned.is_set(), (
"a queued caller's timeout branded another caller's live load as stuck"
)
release.set()
a.join(timeout=15)
assert a_result.get("r") is published
finally:
release.set()
monkeypatch.setattr(mm, "model", None, raising=False)
pool.shutdown(wait=False)