Localize cur_batch field in Scheduler to avoid field-based state access (#29407)

This commit is contained in:
fzyzcjy
2026-07-10 08:55:06 +08:00
committed by GitHub
parent 69368d7593
commit 5be9c9f7c6
10 changed files with 60 additions and 48 deletions
@@ -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."""