Encapsulate the pending-flush bookkeeping in a small wrapper (#25727)
This commit is contained in:
@@ -8,6 +8,9 @@ maybe_stub_sgl_kernel()
|
||||
|
||||
from sglang.srt.managers.io_struct import FlushCacheReqInput
|
||||
from sglang.srt.managers.scheduler import Scheduler
|
||||
from sglang.srt.managers.scheduler_components.flush_wrapper import (
|
||||
SchedulerFlushWrapper,
|
||||
)
|
||||
|
||||
register_cpu_ci(est_time=14, suite="base-a-test-cpu")
|
||||
|
||||
@@ -15,10 +18,14 @@ register_cpu_ci(est_time=14, suite="base-a-test-cpu")
|
||||
class TestSchedulerFlushCache(unittest.TestCase):
|
||||
def _new_scheduler(self) -> Scheduler:
|
||||
scheduler = Scheduler.__new__(Scheduler)
|
||||
scheduler._pending_flush = None
|
||||
scheduler.ipc_channels = MagicMock()
|
||||
scheduler.flush_cache = MagicMock(return_value=True)
|
||||
scheduler.is_fully_idle = MagicMock(return_value=False)
|
||||
scheduler.flush_wrapper = SchedulerFlushWrapper(
|
||||
flush_cache=scheduler.flush_cache,
|
||||
is_fully_idle=scheduler.is_fully_idle,
|
||||
ipc_channels=scheduler.ipc_channels,
|
||||
)
|
||||
return scheduler
|
||||
|
||||
def test_immediate_flush_no_timeout(self):
|
||||
@@ -26,9 +33,7 @@ class TestSchedulerFlushCache(unittest.TestCase):
|
||||
scheduler = self._new_scheduler()
|
||||
scheduler.flush_cache.return_value = False
|
||||
|
||||
output = Scheduler.flush_cache_wrapped(
|
||||
scheduler, FlushCacheReqInput(timeout_s=None)
|
||||
)
|
||||
output = scheduler.flush_wrapper.handle(FlushCacheReqInput(timeout_s=None))
|
||||
|
||||
self.assertFalse(output.success)
|
||||
scheduler.flush_cache.assert_called_once()
|
||||
@@ -38,9 +43,7 @@ class TestSchedulerFlushCache(unittest.TestCase):
|
||||
scheduler = self._new_scheduler()
|
||||
scheduler.is_fully_idle.return_value = True
|
||||
|
||||
output = Scheduler.flush_cache_wrapped(
|
||||
scheduler, FlushCacheReqInput(timeout_s=5.0)
|
||||
)
|
||||
output = scheduler.flush_wrapper.handle(FlushCacheReqInput(timeout_s=5.0))
|
||||
|
||||
self.assertTrue(output.success)
|
||||
scheduler.flush_cache.assert_called_once()
|
||||
@@ -50,22 +53,25 @@ class TestSchedulerFlushCache(unittest.TestCase):
|
||||
scheduler = self._new_scheduler()
|
||||
req = FlushCacheReqInput(timeout_s=3.0)
|
||||
|
||||
with patch("sglang.srt.managers.scheduler.time.monotonic", return_value=10.0):
|
||||
output = Scheduler.flush_cache_wrapped(scheduler, req)
|
||||
with patch(
|
||||
"sglang.srt.managers.scheduler_components.flush_wrapper.time.monotonic",
|
||||
return_value=10.0,
|
||||
):
|
||||
output = scheduler.flush_wrapper.handle(req)
|
||||
|
||||
self.assertIsNone(output)
|
||||
pending_req, deadline = scheduler._pending_flush
|
||||
pending_req, deadline = scheduler.flush_wrapper._pending
|
||||
self.assertIs(pending_req, req)
|
||||
self.assertEqual(deadline, 13.0)
|
||||
|
||||
def test_rejects_when_already_pending(self):
|
||||
"""Any new request is rejected while another is pending."""
|
||||
scheduler = self._new_scheduler()
|
||||
scheduler._pending_flush = (FlushCacheReqInput(timeout_s=10.0), 999.0)
|
||||
scheduler.flush_wrapper._pending = (FlushCacheReqInput(timeout_s=10.0), 999.0)
|
||||
|
||||
for timeout in [None, 5.0]:
|
||||
output = Scheduler.flush_cache_wrapped(
|
||||
scheduler, FlushCacheReqInput(timeout_s=timeout)
|
||||
output = scheduler.flush_wrapper.handle(
|
||||
FlushCacheReqInput(timeout_s=timeout)
|
||||
)
|
||||
self.assertFalse(output.success)
|
||||
self.assertIn("already in progress", output.message)
|
||||
@@ -76,11 +82,11 @@ class TestSchedulerFlushCache(unittest.TestCase):
|
||||
scheduler = self._new_scheduler()
|
||||
scheduler.is_fully_idle.return_value = True
|
||||
req = FlushCacheReqInput(timeout_s=1.0)
|
||||
scheduler._pending_flush = (req, 111.0)
|
||||
scheduler.flush_wrapper._pending = (req, 111.0)
|
||||
|
||||
Scheduler._check_pending_flush(scheduler)
|
||||
scheduler.flush_wrapper.check_pending()
|
||||
|
||||
self.assertIsNone(scheduler._pending_flush)
|
||||
self.assertIsNone(scheduler.flush_wrapper._pending)
|
||||
scheduler.flush_cache.assert_called_once()
|
||||
out = scheduler.ipc_channels.send_to_tokenizer.send_output.call_args.args[0]
|
||||
self.assertTrue(out.success)
|
||||
@@ -88,12 +94,15 @@ class TestSchedulerFlushCache(unittest.TestCase):
|
||||
def test_pending_flush_expires_on_timeout(self):
|
||||
scheduler = self._new_scheduler()
|
||||
req = FlushCacheReqInput(timeout_s=1.0)
|
||||
scheduler._pending_flush = (req, 99.0)
|
||||
scheduler.flush_wrapper._pending = (req, 99.0)
|
||||
|
||||
with patch("sglang.srt.managers.scheduler.time.monotonic", return_value=100.0):
|
||||
Scheduler._check_pending_flush(scheduler)
|
||||
with patch(
|
||||
"sglang.srt.managers.scheduler_components.flush_wrapper.time.monotonic",
|
||||
return_value=100.0,
|
||||
):
|
||||
scheduler.flush_wrapper.check_pending()
|
||||
|
||||
self.assertIsNone(scheduler._pending_flush)
|
||||
self.assertIsNone(scheduler.flush_wrapper._pending)
|
||||
scheduler.flush_cache.assert_not_called()
|
||||
out = scheduler.ipc_channels.send_to_tokenizer.send_output.call_args.args[0]
|
||||
self.assertFalse(out.success)
|
||||
@@ -101,12 +110,15 @@ class TestSchedulerFlushCache(unittest.TestCase):
|
||||
def test_pending_flush_survives_before_deadline(self):
|
||||
scheduler = self._new_scheduler()
|
||||
req = FlushCacheReqInput(timeout_s=5.0)
|
||||
scheduler._pending_flush = (req, 101.0)
|
||||
scheduler.flush_wrapper._pending = (req, 101.0)
|
||||
|
||||
with patch("sglang.srt.managers.scheduler.time.monotonic", return_value=100.0):
|
||||
Scheduler._check_pending_flush(scheduler)
|
||||
with patch(
|
||||
"sglang.srt.managers.scheduler_components.flush_wrapper.time.monotonic",
|
||||
return_value=100.0,
|
||||
):
|
||||
scheduler.flush_wrapper.check_pending()
|
||||
|
||||
self.assertIsNotNone(scheduler._pending_flush)
|
||||
self.assertIsNotNone(scheduler.flush_wrapper._pending)
|
||||
scheduler.ipc_channels.send_to_tokenizer.send_output.assert_not_called()
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user