diff --git a/tensorrt_llm/_torch/pyexecutor/hang_detector.py b/tensorrt_llm/_torch/pyexecutor/hang_detector.py index 9773ff324849..bb191a6f5932 100644 --- a/tensorrt_llm/_torch/pyexecutor/hang_detector.py +++ b/tensorrt_llm/_torch/pyexecutor/hang_detector.py @@ -346,6 +346,9 @@ def __init__( self.active = False self._detected = False self._status_providers: list[Callable[[], str]] = [] + # Monotonic stamp the watcher compares against; ``inf`` means disarmed. + # A plain float store is the entire cost of ``checkpoint()``. + self._deadline = math.inf def start(self): """Enable hang detection.""" @@ -354,18 +357,69 @@ def run_loop(): asyncio.set_event_loop(self.loop) self.loop.run_forever() - self.active = True + with self.lock: + # Locked, not a bare check: concurrent callers could both observe + # ``active`` false and schedule a watcher, and watchers share + # ``_deadline``, so a second one reports the same lapse twice and + # propagates two hard kills. + if self.active: + _best_effort_log_error( + "HangDetector.start() called while already active; ignoring." + ) + return + # Disarmed until the first checkpoint so startup does not lapse. + # Stored before ``active`` is published so a checkpoint racing this + # call cannot have its arm overwritten here. + self._deadline = math.inf + self.active = True + self.loop = asyncio.new_event_loop() self.loop_thread = threading.Thread(target=run_loop, daemon=True, name="hang_detector_loop") self.loop_thread.start() + # One long-lived watcher, scheduled once. The hot path never cancels or + # re-arms it; it only moves ``_deadline``. + self.task = asyncio.run_coroutine_threadsafe(self._watch(), self.loop) def register_status_provider(self, provider: Callable[[], str]) -> None: """Register a nonblocking callable that returns status to dump on hang detection.""" with self.lock: self._status_providers.append(provider) - async def _detect_hang(self) -> None: - await asyncio.sleep(self.timeout) + async def _watch(self) -> None: + """Sleep until the deadline lapses, report, and keep watching. + + Waking early is normal: ``checkpoint()`` pushes ``_deadline`` forward + without touching this task, so each wake-up either finds time left and + sleeps again, or finds the deadline passed and reports. While disarmed + the deadline is ``inf``; the sleep is clamped to ``timeout`` because + ``checkpoint()`` only stores a float and never wakes this loop, so an + unclamped sleep would not notice a later arm. + + This task outlives a report, and outlives a report that raises. A + watchdog that quietly stopped watching would be the exact failure it + exists to catch, and ``on_detected`` is the cross-rank hard kill, which + can itself fail on an already-degraded job. + """ + while self.active: + deadline = self._deadline + remaining = deadline - time.monotonic() + if remaining > 0: + await asyncio.sleep(min(remaining, self.timeout)) + continue + # Disarm only the deadline observed to lapse, so one lapse reports + # once. A checkpoint racing this branch installs a newer deadline; + # clearing that would leave the watcher alive but permanently + # disarmed, which is the failure this watchdog exists to catch. + if self._deadline == deadline: + self._deadline = math.inf + try: + await self._report_hang() + except Exception as error: # noqa: BLE001 - the watcher must survive + _best_effort_log_error( + f"HangDetector: reporting failed with {type(error).__name__}: {error}" + ) + + async def _report_hang(self) -> None: with self.lock: status_providers = tuple(self._status_providers) @@ -399,21 +453,18 @@ def detected(self): def checkpoint(self): """Reset hang detection timer.""" - self.cancel_task() if self.active: - self.task = asyncio.run_coroutine_threadsafe(self._detect_hang(), self.loop) + self._deadline = time.monotonic() + self.timeout def cancel_task(self): - """Cancel the hang detection task.""" - if self.task is not None and not self.task.done(): - self.task.cancel() - self.task = None + """Disarm hang detection until the next checkpoint.""" + self._deadline = math.inf @contextmanager def pause(self): """Pause hang detection in scope.""" + self._deadline = math.inf try: - self.cancel_task() yield finally: self.checkpoint() diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index 00037596ad4d..87295bd20c59 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -1813,9 +1813,8 @@ def profile_step(): # — the events being read have already passed by the time we # read them. Stashing on self lets the /metrics serializer pick # up the values without going through the log line. - should_capture_timing = start_time is not None and ( - self.print_log or self.enable_iter_perf_stats) - if should_capture_timing: + should_capture_timing = self.print_log or self.enable_iter_perf_stats + if should_capture_timing and start_time is not None: end_time = time.time() if it % 2 == 0: end_event_1.record() @@ -1880,14 +1879,15 @@ def profile_step(): calibrator.pre_step(it) start_time = time.time() - if it % 2 == 0: - if start_event_1 is None: - start_event_1 = torch.cuda.Event(enable_timing=True) - start_event_1.record() - else: - if start_event_2 is None: - start_event_2 = torch.cuda.Event(enable_timing=True) - start_event_2.record() + if should_capture_timing: + if it % 2 == 0: + if start_event_1 is None: + start_event_1 = torch.cuda.Event(enable_timing=True) + start_event_1.record() + else: + if start_event_2 is None: + start_event_2 = torch.cuda.Event(enable_timing=True) + start_event_2.record() try: yield profile_step diff --git a/tests/unittest/_torch/executor/test_hang_detector_kill.py b/tests/unittest/_torch/executor/test_hang_detector_kill.py index 210c3bbbda82..8d4f003b04cf 100644 --- a/tests/unittest/_torch/executor/test_hang_detector_kill.py +++ b/tests/unittest/_torch/executor/test_hang_detector_kill.py @@ -16,6 +16,7 @@ import asyncio import contextlib +import math import os import shutil import signal @@ -66,6 +67,106 @@ def test_checkpoint_resets_timer(): assert hd.detected() is False +def test_checkpoint_reuses_one_watcher_task(): + """One watcher task serves every checkpoint, pause and resume. + + The executor loop checkpoints several times per iteration, and each + schedule/cancel of a task wakes the detector's event-loop thread, so the + single-task design is what keeps checkpoint() off that thread entirely. + """ + hd = HangDetector(timeout=30) + with hd: + task = hd.task + assert task is not None + for _ in range(10): + hd.checkpoint() + with hd.pause(): + hd.checkpoint() + hd.checkpoint() + assert hd.task is task + assert not task.done() + + +def test_detector_is_disarmed_until_the_first_checkpoint(): + """start() enables detection; the first checkpoint arms the deadline. + + Callers separate lifecycle start from arming, so the start-to-first- + checkpoint window must not be attributed to the loop as a hang. + """ + fired = [] + hd = HangDetector(timeout=1, on_detected=lambda: fired.append(1)) + with hd: + time.sleep(2.0) # would fire if start() armed the deadline itself + assert fired == [] + assert hd.detected() is False + + +def test_watcher_survives_a_raising_callback(): + """on_detected is the cross-rank hard kill and can fail on a broken job.""" + fired = [] + + def boom(): + fired.append(1) + raise RuntimeError("hard kill failed") + + hd = HangDetector(timeout=1, on_detected=boom) + with hd: + hd.checkpoint() + deadline = time.monotonic() + 5.0 + while time.monotonic() < deadline and len(fired) < 1: + time.sleep(0.05) + assert len(fired) == 1 + + # The watcher is still live and still able to report. + assert not hd.task.done() + hd.checkpoint() + deadline = time.monotonic() + 5.0 + while time.monotonic() < deadline and len(fired) < 2: + time.sleep(0.05) + assert len(fired) == 2 + + +def test_a_checkpoint_racing_the_disarm_is_not_erased(monkeypatch): + """A checkpoint landing as the watcher disarms must survive. + + The watcher reads the lapsed deadline, then clears it. A checkpoint in + between installs a newer deadline; clearing that would leave the watcher + running but permanently disarmed, so the work it just armed could hang + undetected. + """ + hd = HangDetector(timeout=1, on_detected=lambda: None) + real_monotonic = hang_detector_module.time.monotonic + state = {"injected": False, "busy": False} + + def racing_monotonic(): + now = real_monotonic() + if state["injected"] or state["busy"]: + return now + # Only inject from `_watch`'s own lapse computation. asyncio's event + # loop also reads the clock, and injecting from there would land + # outside the window and silently make this test vacuous. + caller = sys._getframe(1) + if caller.f_code.co_name != "_watch" or hd._deadline > now: + return now + # `_watch` has already read `self._deadline` into a local by now, so + # this checkpoint lands exactly between that read and the disarm. + state["busy"] = True + state["injected"] = True + hd.checkpoint() + state["busy"] = False + return now + + monkeypatch.setattr(hang_detector_module.time, "monotonic", racing_monotonic) + with hd: + hd.checkpoint() + deadline = real_monotonic() + 5.0 + while real_monotonic() < deadline and not state["injected"]: + time.sleep(0.05) + assert state["injected"], "the racing checkpoint never landed" + time.sleep(0.2) # let the watcher finish its disarm/report pass + assert hd._deadline != math.inf, "the racing checkpoint's arm was erased" + + def test_pause_suppresses_detection(): fired = [] hd = HangDetector(timeout=1, on_detected=lambda: fired.append(1)) @@ -102,7 +203,7 @@ def failing_provider(): detector.register_status_provider(failing_provider) detector.register_status_provider(lambda: "transceiver status") - asyncio.run(detector._detect_hang()) + asyncio.run(detector._report_hang()) messages = "\n".join(message for kind, message in events if kind == "log") assert "provider failed" in messages