From ad788740c2b7a80f22f2a9c2efae4b43cb30d6d0 Mon Sep 17 00:00:00 2001 From: T4t4KAU Date: Thu, 17 Sep 2026 20:45:00 +0800 Subject: [PATCH 1/2] perf: reuse running reservations within each scheduler step --- python/infinilm/llm/scheduler.py | 47 +++++-- test/llm/cpu_modules.py | 36 +++++ test/llm/test_scheduler_admission.py | 197 +++++++++++++++++++++++++++ 3 files changed, 266 insertions(+), 14 deletions(-) create mode 100644 test/llm/cpu_modules.py create mode 100644 test/llm/test_scheduler_admission.py diff --git a/python/infinilm/llm/scheduler.py b/python/infinilm/llm/scheduler.py index c10b55f2f..085837658 100644 --- a/python/infinilm/llm/scheduler.py +++ b/python/infinilm/llm/scheduler.py @@ -133,6 +133,7 @@ def schedule(self) -> Optional[SchedulerOutput]: is_prefill = False current_num_batched_tokens = 0 current_prefill_extra_blocks = 0 + running_required_blocks = None # Process Waiting queue (prefill phase) while ( @@ -199,10 +200,18 @@ def schedule(self) -> Optional[SchedulerOutput]: deferred_requests.append(req) break + if running_required_blocks is None and ( + self.mamba_cache_manager is None + or req.mamba_cache_index is not None + or self.mamba_cache_manager.can_allocate() + ): + # Running requests do not advance during this prefill loop. + running_required_blocks = self._get_running_required_blocks() if not self.can_accept_request( req, num_local_computed_tokens, current_prefill_extra_blocks, + running_required_blocks=running_required_blocks, ): logger.warning( "Insufficient KV cache blocks for request %s, deferring.", @@ -450,12 +459,31 @@ def complete_requests(self, requests: List[InferenceRequest]): # Still running, put back in running queue self.running_queue.sync_q.put(req) + def _get_running_required_blocks(self) -> int: + """Sum decode reservations while preserving running queue order.""" + total_required_blocks = 0 + running_queue_size = self.running_queue.sync_q.qsize() + for _ in range(running_queue_size): + req = self.running_queue.sync_q.get() + remaining_tokens = ( + req.sampling_params.max_tokens - req.get_num_generated_tokens() + ) + num_blocks_needed = ( + remaining_tokens + self.block_size - 1 + ) // self.block_size + total_required_blocks += num_blocks_needed + self.running_queue.sync_q.put(req) + return total_required_blocks + def can_accept_request( self, request: InferenceRequest, num_local_computed_tokens: int, current_prefill_extra_blocks: int = 0, + *, + running_required_blocks: int | None = None, ) -> bool: + """Check capacity, optionally reusing this schedule's running reservations.""" if ( self.mamba_cache_manager is not None and request.mamba_cache_index is None @@ -463,20 +491,11 @@ def can_accept_request( ): return False - total_required_blocks = 0 - - # Calculate blocks needed for running requests - running_queue_size = self.running_queue.sync_q.qsize() - for _ in range(running_queue_size): - req = self.running_queue.sync_q.get() - remaining_tokens = ( - req.sampling_params.max_tokens - req.get_num_generated_tokens() - ) - num_blocks_needed = ( - remaining_tokens + self.block_size - 1 - ) // self.block_size - total_required_blocks += num_blocks_needed - self.running_queue.sync_q.put(req) + total_required_blocks = ( + self._get_running_required_blocks() + if running_required_blocks is None + else running_required_blocks + ) # Calculate blocks needed for the new request total_length = request.get_prompt_length() - num_local_computed_tokens diff --git a/test/llm/cpu_modules.py b/test/llm/cpu_modules.py new file mode 100644 index 000000000..3e65f55ad --- /dev/null +++ b/test/llm/cpu_modules.py @@ -0,0 +1,36 @@ +"""Load cache and scheduler modules without initializing the GPU engine.""" + +import importlib.util +import sys +from pathlib import Path +from types import SimpleNamespace + + +def load_cpu_modules(): + directory = Path(__file__).resolve().parents[2] / "python/infinilm/llm" + missing = object() + originals = {} + modules = {} + try: + for name in ( + "prefix_cache", + "sampling_params", + "request", + "cache_manager", + "scheduler", + ): + key = f"infinilm.llm.{name}" + spec = importlib.util.spec_from_file_location(key, directory / f"{name}.py") + module = importlib.util.module_from_spec(spec) + originals[key] = sys.modules.get(key, missing) + sys.modules[key] = module + spec.loader.exec_module(module) + modules[name] = module + finally: + # Retain imported dependencies, including non-reloadable native extensions. + for key, original in originals.items(): + if original is missing: + sys.modules.pop(key, None) + else: + sys.modules[key] = original + return SimpleNamespace(**modules) diff --git a/test/llm/test_scheduler_admission.py b/test/llm/test_scheduler_admission.py new file mode 100644 index 000000000..bd2c64d1f --- /dev/null +++ b/test/llm/test_scheduler_admission.py @@ -0,0 +1,197 @@ +"""CPU checks for admission reservations within and across scheduler steps.""" + +import unittest +from unittest.mock import patch + +from cpu_modules import load_cpu_modules + +modules = load_cpu_modules() +Scheduler = modules.scheduler.Scheduler +InferenceRequest = modules.request.InferenceRequest +RequestStatus = modules.request.RequestStatus +SamplingParams = modules.sampling_params.SamplingParams + + +def make_request(request_id, prompt_length=1, max_tokens=16): + return InferenceRequest( + request_id, + prompt_token_ids=list(range(prompt_length)), + sampling_params=SamplingParams(max_tokens=max_tokens), + ) + + +def add_running(scheduler, request_id, max_tokens=32): + request = make_request(request_id, max_tokens=max_tokens) + request.initialize_block_hashes(scheduler.block_size, False) + request.block_table, request.slot_mapping = scheduler.cache_manager.allocate_slots( + request.get_prompt_length() + ) + request.num_blocks = len(request.block_table) + request.append_generated_token_id(42) + request.status = RequestStatus.RUNNING + scheduler.complete_requests([request]) + return request + + +def queue_ids(janus_queue): + items = [janus_queue.sync_q.get_nowait() for _ in range(janus_queue.sync_q.qsize())] + for item in items: + janus_queue.sync_q.put(item) + return [item.request_id for item in items] + + +class RemoteConnector: + def get_num_new_matched_tokens(self, request, num_computed_tokens): + if request.request_id.startswith("remote"): + return request.get_prompt_length() - num_computed_tokens, True + return 0, False + + def update_state_after_alloc(self, *args): + pass + + def build_connector_meta(self): + return None + + def request_finished(self, *args): + return False, None + + +class SchedulerAdmissionTest(unittest.TestCase): + def make_scheduler(self, **kwargs): + scheduler = Scheduler(block_size=16, enable_prefix_caching=False, **kwargs) + self.addCleanup(scheduler.waiting_queue.close) + self.addCleanup(scheduler.running_queue.close) + return scheduler + + def test_running_reservations_scanned_once_and_queue_order_preserved(self): + scheduler = self.make_scheduler(max_batch_size=8, num_blocks=128) + for index in range(3): + add_running(scheduler, f"running-{index}") + for index in range(8): + scheduler.add_request(make_request(f"waiting-{index}")) + before = queue_ids(scheduler.running_queue) + with patch.object( + scheduler, + "_get_running_required_blocks", + wraps=scheduler._get_running_required_blocks, + ) as scan: + output = scheduler.schedule() + self.assertTrue(output.is_prefill) + self.assertEqual(output.num_requests, 8) + self.assertEqual(scan.call_count, 1) + self.assertEqual(queue_ids(scheduler.running_queue), before) + + def test_prefill_reservations_and_remaining_capacity_stay_live(self): + scheduler = self.make_scheduler(num_blocks=5) + for index in range(3): + scheduler.add_request(make_request(str(index))) + with self.assertLogs(modules.scheduler.logger, level="WARNING"): + output = scheduler.schedule() + self.assertEqual( + [req.request_id for req in output.scheduled_requests], ["0", "1"] + ) + self.assertEqual(queue_ids(scheduler.waiting_queue), ["2"]) + self.assertEqual(scheduler.cache_manager.get_num_free_blocks(), 3) + + def test_remote_reservations_stay_live_in_same_step(self): + scheduler = self.make_scheduler(num_blocks=5, connector=RemoteConnector()) + scheduler.add_request(make_request("remote-0")) + scheduler.add_request(make_request("local-1")) + scheduler.add_request(make_request("local-2")) + with self.assertLogs(modules.scheduler.logger, level="WARNING"): + output = scheduler.schedule() + self.assertEqual( + [req.request_id for req in output.scheduled_requests], ["local-1"] + ) + self.assertEqual(scheduler.pending_kv_decode_blocks, 1) + self.assertEqual(set(scheduler.remote_kv_requests), {"remote-0"}) + self.assertEqual(queue_ids(scheduler.waiting_queue), ["local-2"]) + + def test_snapshot_is_recomputed_on_next_step(self): + scheduler = self.make_scheduler(num_blocks=64) + running = add_running(scheduler, "running") + scheduler.add_request(make_request("first")) + with patch.object( + scheduler, "can_accept_request", wraps=scheduler.can_accept_request + ) as admission: + output = scheduler.schedule() + self.assertEqual(admission.call_args.kwargs["running_required_blocks"], 2) + output.scheduled_requests[0].status = RequestStatus.FINISHED + scheduler.complete_requests(output.scheduled_requests) + scheduler.cache_manager.append_slots(running.block_table, 2, 15) + for _ in range(15): + running.append_generated_token_id(42) + scheduler.add_request(make_request("second")) + scheduler.schedule() + self.assertEqual(admission.call_args.kwargs["running_required_blocks"], 1) + + def test_standalone_admission_recomputes_running_reservations(self): + scheduler = self.make_scheduler(num_blocks=5) + running = add_running(scheduler, "running") + candidate = make_request("candidate") + scheduler.pending_kv_decode_blocks = 1 + self.assertFalse(scheduler.can_accept_request(candidate, 0)) + for _ in range(15): + running.append_generated_token_id(42) + self.assertTrue(scheduler.can_accept_request(candidate, 0)) + + def test_decode_only_step_does_not_scan_reservations(self): + scheduler = self.make_scheduler(num_blocks=16) + running = add_running(scheduler, "running") + with patch.object(scheduler, "_get_running_required_blocks") as scan: + output = scheduler.schedule() + scan.assert_not_called() + self.assertFalse(output.is_prefill) + self.assertEqual(output.scheduled_requests, [running]) + + def test_canceled_waiting_request_and_token_budget(self): + scheduler = self.make_scheduler(num_blocks=16, max_num_batched_tokens=5) + canceled = make_request("canceled") + scheduler.add_request(canceled) + canceled.status = RequestStatus.CANCELED + scheduler.add_request(make_request("first", prompt_length=3)) + scheduler.add_request(make_request("second", prompt_length=3)) + with patch.object( + scheduler, "can_accept_request", wraps=scheduler.can_accept_request + ) as admission: + output = scheduler.schedule() + self.assertEqual(admission.call_count, 1) + self.assertEqual( + [req.request_id for req in output.scheduled_requests], ["first"] + ) + self.assertEqual(queue_ids(scheduler.waiting_queue), ["second"]) + + def test_mamba_rows_still_limit_admission(self): + scheduler = self.make_scheduler( + num_blocks=64, has_mamba_cache=True, num_mamba_cache_blocks=2 + ) + scheduler.add_request(make_request("first")) + scheduler.add_request(make_request("second")) + with self.assertLogs(modules.scheduler.logger, level="WARNING"): + output = scheduler.schedule() + self.assertEqual( + [req.request_id for req in output.scheduled_requests], ["first"] + ) + self.assertEqual(queue_ids(scheduler.waiting_queue), ["second"]) + self.assertEqual(scheduler.mamba_cache_manager.get_num_free_blocks(), 0) + + def test_full_mamba_pool_rejects_without_scanning_running_reservations(self): + scheduler = self.make_scheduler( + num_blocks=64, has_mamba_cache=True, num_mamba_cache_blocks=2 + ) + running = add_running(scheduler, "running") + running.mamba_cache_index = scheduler.mamba_cache_manager.allocate() + scheduler.add_request(make_request("waiting")) + with ( + patch.object(scheduler, "_get_running_required_blocks") as scan, + self.assertLogs(modules.scheduler.logger, level="WARNING"), + ): + output = scheduler.schedule() + scan.assert_not_called() + self.assertFalse(output.is_prefill) + self.assertEqual(output.scheduled_requests, [running]) + self.assertEqual(queue_ids(scheduler.waiting_queue), ["waiting"]) + + +if __name__ == "__main__": + unittest.main() From a717083037796f33bd119c24b2f01def78337741 Mon Sep 17 00:00:00 2001 From: T4t4KAU Date: Thu, 17 Sep 2026 20:51:59 +0800 Subject: [PATCH 2/2] chore: remove added unit tests --- test/llm/cpu_modules.py | 36 ----- test/llm/test_scheduler_admission.py | 197 --------------------------- 2 files changed, 233 deletions(-) delete mode 100644 test/llm/cpu_modules.py delete mode 100644 test/llm/test_scheduler_admission.py diff --git a/test/llm/cpu_modules.py b/test/llm/cpu_modules.py deleted file mode 100644 index 3e65f55ad..000000000 --- a/test/llm/cpu_modules.py +++ /dev/null @@ -1,36 +0,0 @@ -"""Load cache and scheduler modules without initializing the GPU engine.""" - -import importlib.util -import sys -from pathlib import Path -from types import SimpleNamespace - - -def load_cpu_modules(): - directory = Path(__file__).resolve().parents[2] / "python/infinilm/llm" - missing = object() - originals = {} - modules = {} - try: - for name in ( - "prefix_cache", - "sampling_params", - "request", - "cache_manager", - "scheduler", - ): - key = f"infinilm.llm.{name}" - spec = importlib.util.spec_from_file_location(key, directory / f"{name}.py") - module = importlib.util.module_from_spec(spec) - originals[key] = sys.modules.get(key, missing) - sys.modules[key] = module - spec.loader.exec_module(module) - modules[name] = module - finally: - # Retain imported dependencies, including non-reloadable native extensions. - for key, original in originals.items(): - if original is missing: - sys.modules.pop(key, None) - else: - sys.modules[key] = original - return SimpleNamespace(**modules) diff --git a/test/llm/test_scheduler_admission.py b/test/llm/test_scheduler_admission.py deleted file mode 100644 index bd2c64d1f..000000000 --- a/test/llm/test_scheduler_admission.py +++ /dev/null @@ -1,197 +0,0 @@ -"""CPU checks for admission reservations within and across scheduler steps.""" - -import unittest -from unittest.mock import patch - -from cpu_modules import load_cpu_modules - -modules = load_cpu_modules() -Scheduler = modules.scheduler.Scheduler -InferenceRequest = modules.request.InferenceRequest -RequestStatus = modules.request.RequestStatus -SamplingParams = modules.sampling_params.SamplingParams - - -def make_request(request_id, prompt_length=1, max_tokens=16): - return InferenceRequest( - request_id, - prompt_token_ids=list(range(prompt_length)), - sampling_params=SamplingParams(max_tokens=max_tokens), - ) - - -def add_running(scheduler, request_id, max_tokens=32): - request = make_request(request_id, max_tokens=max_tokens) - request.initialize_block_hashes(scheduler.block_size, False) - request.block_table, request.slot_mapping = scheduler.cache_manager.allocate_slots( - request.get_prompt_length() - ) - request.num_blocks = len(request.block_table) - request.append_generated_token_id(42) - request.status = RequestStatus.RUNNING - scheduler.complete_requests([request]) - return request - - -def queue_ids(janus_queue): - items = [janus_queue.sync_q.get_nowait() for _ in range(janus_queue.sync_q.qsize())] - for item in items: - janus_queue.sync_q.put(item) - return [item.request_id for item in items] - - -class RemoteConnector: - def get_num_new_matched_tokens(self, request, num_computed_tokens): - if request.request_id.startswith("remote"): - return request.get_prompt_length() - num_computed_tokens, True - return 0, False - - def update_state_after_alloc(self, *args): - pass - - def build_connector_meta(self): - return None - - def request_finished(self, *args): - return False, None - - -class SchedulerAdmissionTest(unittest.TestCase): - def make_scheduler(self, **kwargs): - scheduler = Scheduler(block_size=16, enable_prefix_caching=False, **kwargs) - self.addCleanup(scheduler.waiting_queue.close) - self.addCleanup(scheduler.running_queue.close) - return scheduler - - def test_running_reservations_scanned_once_and_queue_order_preserved(self): - scheduler = self.make_scheduler(max_batch_size=8, num_blocks=128) - for index in range(3): - add_running(scheduler, f"running-{index}") - for index in range(8): - scheduler.add_request(make_request(f"waiting-{index}")) - before = queue_ids(scheduler.running_queue) - with patch.object( - scheduler, - "_get_running_required_blocks", - wraps=scheduler._get_running_required_blocks, - ) as scan: - output = scheduler.schedule() - self.assertTrue(output.is_prefill) - self.assertEqual(output.num_requests, 8) - self.assertEqual(scan.call_count, 1) - self.assertEqual(queue_ids(scheduler.running_queue), before) - - def test_prefill_reservations_and_remaining_capacity_stay_live(self): - scheduler = self.make_scheduler(num_blocks=5) - for index in range(3): - scheduler.add_request(make_request(str(index))) - with self.assertLogs(modules.scheduler.logger, level="WARNING"): - output = scheduler.schedule() - self.assertEqual( - [req.request_id for req in output.scheduled_requests], ["0", "1"] - ) - self.assertEqual(queue_ids(scheduler.waiting_queue), ["2"]) - self.assertEqual(scheduler.cache_manager.get_num_free_blocks(), 3) - - def test_remote_reservations_stay_live_in_same_step(self): - scheduler = self.make_scheduler(num_blocks=5, connector=RemoteConnector()) - scheduler.add_request(make_request("remote-0")) - scheduler.add_request(make_request("local-1")) - scheduler.add_request(make_request("local-2")) - with self.assertLogs(modules.scheduler.logger, level="WARNING"): - output = scheduler.schedule() - self.assertEqual( - [req.request_id for req in output.scheduled_requests], ["local-1"] - ) - self.assertEqual(scheduler.pending_kv_decode_blocks, 1) - self.assertEqual(set(scheduler.remote_kv_requests), {"remote-0"}) - self.assertEqual(queue_ids(scheduler.waiting_queue), ["local-2"]) - - def test_snapshot_is_recomputed_on_next_step(self): - scheduler = self.make_scheduler(num_blocks=64) - running = add_running(scheduler, "running") - scheduler.add_request(make_request("first")) - with patch.object( - scheduler, "can_accept_request", wraps=scheduler.can_accept_request - ) as admission: - output = scheduler.schedule() - self.assertEqual(admission.call_args.kwargs["running_required_blocks"], 2) - output.scheduled_requests[0].status = RequestStatus.FINISHED - scheduler.complete_requests(output.scheduled_requests) - scheduler.cache_manager.append_slots(running.block_table, 2, 15) - for _ in range(15): - running.append_generated_token_id(42) - scheduler.add_request(make_request("second")) - scheduler.schedule() - self.assertEqual(admission.call_args.kwargs["running_required_blocks"], 1) - - def test_standalone_admission_recomputes_running_reservations(self): - scheduler = self.make_scheduler(num_blocks=5) - running = add_running(scheduler, "running") - candidate = make_request("candidate") - scheduler.pending_kv_decode_blocks = 1 - self.assertFalse(scheduler.can_accept_request(candidate, 0)) - for _ in range(15): - running.append_generated_token_id(42) - self.assertTrue(scheduler.can_accept_request(candidate, 0)) - - def test_decode_only_step_does_not_scan_reservations(self): - scheduler = self.make_scheduler(num_blocks=16) - running = add_running(scheduler, "running") - with patch.object(scheduler, "_get_running_required_blocks") as scan: - output = scheduler.schedule() - scan.assert_not_called() - self.assertFalse(output.is_prefill) - self.assertEqual(output.scheduled_requests, [running]) - - def test_canceled_waiting_request_and_token_budget(self): - scheduler = self.make_scheduler(num_blocks=16, max_num_batched_tokens=5) - canceled = make_request("canceled") - scheduler.add_request(canceled) - canceled.status = RequestStatus.CANCELED - scheduler.add_request(make_request("first", prompt_length=3)) - scheduler.add_request(make_request("second", prompt_length=3)) - with patch.object( - scheduler, "can_accept_request", wraps=scheduler.can_accept_request - ) as admission: - output = scheduler.schedule() - self.assertEqual(admission.call_count, 1) - self.assertEqual( - [req.request_id for req in output.scheduled_requests], ["first"] - ) - self.assertEqual(queue_ids(scheduler.waiting_queue), ["second"]) - - def test_mamba_rows_still_limit_admission(self): - scheduler = self.make_scheduler( - num_blocks=64, has_mamba_cache=True, num_mamba_cache_blocks=2 - ) - scheduler.add_request(make_request("first")) - scheduler.add_request(make_request("second")) - with self.assertLogs(modules.scheduler.logger, level="WARNING"): - output = scheduler.schedule() - self.assertEqual( - [req.request_id for req in output.scheduled_requests], ["first"] - ) - self.assertEqual(queue_ids(scheduler.waiting_queue), ["second"]) - self.assertEqual(scheduler.mamba_cache_manager.get_num_free_blocks(), 0) - - def test_full_mamba_pool_rejects_without_scanning_running_reservations(self): - scheduler = self.make_scheduler( - num_blocks=64, has_mamba_cache=True, num_mamba_cache_blocks=2 - ) - running = add_running(scheduler, "running") - running.mamba_cache_index = scheduler.mamba_cache_manager.allocate() - scheduler.add_request(make_request("waiting")) - with ( - patch.object(scheduler, "_get_running_required_blocks") as scan, - self.assertLogs(modules.scheduler.logger, level="WARNING"), - ): - output = scheduler.schedule() - scan.assert_not_called() - self.assertFalse(output.is_prefill) - self.assertEqual(output.scheduled_requests, [running]) - self.assertEqual(queue_ids(scheduler.waiting_queue), ["waiting"]) - - -if __name__ == "__main__": - unittest.main()