diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 84e5afa83..08808b0da 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -18,7 +18,6 @@ import logging import os import signal import sys -import threading import time from collections import deque from concurrent import futures @@ -146,6 +145,7 @@ from sglang.srt.managers.scheduler_profiler_mixin import SchedulerProfilerMixin from sglang.srt.managers.scheduler_recv_skipper import SchedulerRecvSkipper from sglang.srt.managers.scheduler_runtime_checker_mixin import ( SchedulerRuntimeCheckerMixin, + SchedulerWatchdog, ) from sglang.srt.managers.scheduler_update_weights_mixin import ( SchedulerUpdateWeightsMixin, @@ -524,10 +524,9 @@ class Scheduler( self.new_token_ratio = self.init_new_token_ratio # Init watchdog thread - self.watchdog_timeout = server_args.watchdog_timeout - t = threading.Thread(target=self.watchdog_thread, daemon=True) - t.start() - self.parent_process = psutil.Process().parent() + self.watchdog = SchedulerWatchdog( + self, watchdog_timeout=server_args.watchdog_timeout + ) # Init memory saver, profiler and metric stats self.memory_saver_adapter = TorchMemorySaverAdapter.create( diff --git a/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py b/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py index 757df1cd3..b60bf83e2 100644 --- a/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py +++ b/python/sglang/srt/managers/scheduler_runtime_checker_mixin.py @@ -3,9 +3,12 @@ from __future__ import annotations import logging import signal import sys +import threading import time from typing import TYPE_CHECKING +import psutil + from sglang.srt.disaggregation.utils import DisaggregationMode from sglang.srt.environ import envs from sglang.srt.managers.schedule_batch import ScheduleBatch @@ -302,33 +305,47 @@ class SchedulerRuntimeCheckerMixin: self.new_token_ratio = self.init_new_token_ratio self.maybe_sleep_on_idle() - def watchdog_thread(self: Scheduler): - """A watch dog thread that will try to kill the server itself if one forward batch takes too long.""" - self.watchdog_last_forward_ct = 0 - self.watchdog_last_time = time.perf_counter() + +class SchedulerWatchdog: + """A watch dog that will try to kill the server itself if one forward batch takes too long.""" + + def __init__(self, scheduler: Scheduler, watchdog_timeout: float): + self.scheduler = scheduler + + self.watchdog_timeout = watchdog_timeout + t = threading.Thread(target=self._watchdog_thread, daemon=True) + t.start() + self.parent_process = psutil.Process().parent() + + def _watchdog_thread(self): + watchdog_last_forward_ct = 0 + watchdog_last_time = time.perf_counter() while True: current = time.perf_counter() - if self.cur_batch is not None: - if self.watchdog_last_forward_ct == self.forward_ct: - if current > self.watchdog_last_time + self.watchdog_timeout: + if self.scheduler.cur_batch is not None: + if watchdog_last_forward_ct == self.scheduler.forward_ct: + if current > watchdog_last_time + self.watchdog_timeout: break else: - self.watchdog_last_forward_ct = self.forward_ct - self.watchdog_last_time = current + watchdog_last_forward_ct = self.scheduler.forward_ct + watchdog_last_time = current time.sleep(self.watchdog_timeout // 2) if not disable_request_logging(): + # TODO extract this duplicated logic w/ another place # Print batch size and memory pool info to check whether there are de-sync issues. - if self.is_hybrid_swa: - _, info_msg = self._check_hybrid_memory() - elif self.is_ssm_model and isinstance(self.tree_cache, MambaRadixCache): - _, info_msg = self._check_mamba_memory() + if self.scheduler.is_hybrid_swa: + _, info_msg = self.scheduler._check_hybrid_memory() + elif self.scheduler.is_ssm_model and isinstance( + self.scheduler.tree_cache, MambaRadixCache + ): + _, info_msg = self.scheduler._check_mamba_memory() else: - _, info_msg = self._check_radix_cache_memory() + _, info_msg = self.scheduler._check_radix_cache_memory() logger.error( - f"{self.cur_batch.batch_size()=}\n" - f"{self.cur_batch.reqs=}\n" + f"{self.scheduler.cur_batch.batch_size()=}\n" + f"{self.scheduler.cur_batch.reqs=}\n" f"{info_msg}" )