Localize cur_batch field in Scheduler to avoid field-based state access (#29407)
This commit is contained in:
@@ -77,7 +77,10 @@ class TestScriptedPpChunkSweep(ScriptedTestCase):
|
||||
scheduler.chunked_req is None
|
||||
and len(scheduler.waiting_queue) == 0
|
||||
and all(x.is_empty() for x in scheduler.running_mbs)
|
||||
and (scheduler.cur_batch is None or scheduler.cur_batch.is_empty())
|
||||
and (
|
||||
scheduler.cur_batch_for_debug is None
|
||||
or scheduler.cur_batch_for_debug.is_empty()
|
||||
)
|
||||
and (scheduler.last_batch is None or scheduler.last_batch.is_empty())
|
||||
)
|
||||
if in_flight and queues_clear:
|
||||
|
||||
@@ -193,7 +193,7 @@ class TestDecodeQueueCleanup(CustomTestCase):
|
||||
scheduler.dllm_manager = MagicMock()
|
||||
scheduler.dllm_manager.any_staging_reqs.return_value = False
|
||||
scheduler.last_batch = None
|
||||
scheduler.cur_batch = None
|
||||
scheduler.cur_batch_for_debug = None
|
||||
scheduler.enable_overlap = False
|
||||
scheduler.ps = SimpleNamespace(pp_size=1)
|
||||
scheduler.running_mbs = []
|
||||
|
||||
@@ -1135,7 +1135,7 @@ class TestMlxOverlapScheduler(unittest.TestCase):
|
||||
scheduler.waiting_queue = []
|
||||
scheduler.result_queue = deque()
|
||||
scheduler.future_map = SimpleNamespace()
|
||||
scheduler.cur_batch = None
|
||||
scheduler.cur_batch_for_debug = None
|
||||
scheduler.last_batch = None
|
||||
scheduler.tp_worker = SimpleNamespace(
|
||||
async_forward_batch_generation_mlx=fake_forward
|
||||
|
||||
@@ -26,7 +26,7 @@ class TestSchedulerPauseGeneration(unittest.TestCase):
|
||||
scheduler._engine_paused = False
|
||||
scheduler.enable_overlap = False
|
||||
scheduler.last_batch = None
|
||||
scheduler.cur_batch = None
|
||||
scheduler.cur_batch_for_debug = None
|
||||
scheduler.chunked_req = None
|
||||
scheduler.running_batch = MagicMock()
|
||||
scheduler.running_batch.reqs = []
|
||||
@@ -59,11 +59,11 @@ class TestSchedulerPauseGeneration(unittest.TestCase):
|
||||
"""in_place pause should only set _engine_paused and return."""
|
||||
scheduler = self._new_scheduler()
|
||||
scheduler.last_batch = MagicMock()
|
||||
scheduler.cur_batch = MagicMock()
|
||||
scheduler.cur_batch_for_debug = MagicMock()
|
||||
scheduler.chunked_req = MagicMock()
|
||||
|
||||
original_last_batch = scheduler.last_batch
|
||||
original_cur_batch = scheduler.cur_batch
|
||||
original_cur_batch = scheduler.cur_batch_for_debug
|
||||
original_chunked_req = scheduler.chunked_req
|
||||
|
||||
scheduler.pause_generation(PauseGenerationReqInput(mode="in_place"))
|
||||
@@ -71,7 +71,7 @@ class TestSchedulerPauseGeneration(unittest.TestCase):
|
||||
self.assertTrue(scheduler._engine_paused)
|
||||
# All state must be preserved — no mutation
|
||||
self.assertIs(scheduler.last_batch, original_last_batch)
|
||||
self.assertIs(scheduler.cur_batch, original_cur_batch)
|
||||
self.assertIs(scheduler.cur_batch_for_debug, original_cur_batch)
|
||||
self.assertIs(scheduler.chunked_req, original_chunked_req)
|
||||
|
||||
def test_inplace_does_not_drain_overlap_queue(self):
|
||||
@@ -99,17 +99,17 @@ class TestSchedulerPauseGeneration(unittest.TestCase):
|
||||
scheduler.running_batch.merge_batch.assert_not_called()
|
||||
|
||||
def test_abort_clears_state(self):
|
||||
"""abort mode should clear last_batch and cur_batch."""
|
||||
"""abort mode should clear last_batch and cur_batch_for_debug."""
|
||||
scheduler = self._new_scheduler()
|
||||
scheduler.last_batch = MagicMock()
|
||||
scheduler.last_batch.forward_mode.is_extend.return_value = False
|
||||
scheduler.cur_batch = MagicMock()
|
||||
scheduler.cur_batch_for_debug = MagicMock()
|
||||
|
||||
scheduler.pause_generation(PauseGenerationReqInput(mode="abort"))
|
||||
|
||||
self.assertTrue(scheduler._engine_paused)
|
||||
self.assertIsNone(scheduler.last_batch)
|
||||
self.assertIsNone(scheduler.cur_batch)
|
||||
self.assertIsNone(scheduler.cur_batch_for_debug)
|
||||
|
||||
def test_retract_clears_running_batch(self):
|
||||
"""retract mode should retract all requests from running_batch."""
|
||||
|
||||
Reference in New Issue
Block a user