diff --git a/python/sglang/srt/managers/schedule_policy.py b/python/sglang/srt/managers/schedule_policy.py index 138ea55ee..32c0215c2 100644 --- a/python/sglang/srt/managers/schedule_policy.py +++ b/python/sglang/srt/managers/schedule_policy.py @@ -93,6 +93,7 @@ class SchedulePolicy: self.enable_hierarchical_cache = enable_hierarchical_cache self.enable_priority_scheduling = enable_priority_scheduling self.schedule_low_priority_values_first = schedule_low_priority_values_first + self.priority_sign = 1 if schedule_low_priority_values_first else -1 # It is used to find the matching prefix for in-batch prefix caching. self.waiting_queue_radix_tree = RadixCache.create_simulated() @@ -101,7 +102,7 @@ class SchedulePolicy: if self.policy == CacheAgnosticPolicy.FCFS: if self.enable_priority_scheduling: SchedulePolicy._sort_by_priority_and_fcfs( - waiting_queue, self.schedule_low_priority_values_first + waiting_queue, self.priority_sign ) return False @@ -128,7 +129,7 @@ class SchedulePolicy: SchedulePolicy._sort_by_longest_output( waiting_queue, self.enable_priority_scheduling, - self.schedule_low_priority_values_first, + self.priority_sign, ) elif policy == CacheAgnosticPolicy.RANDOM: SchedulePolicy._sort_randomly(waiting_queue) @@ -173,7 +174,6 @@ class SchedulePolicy: for r in waiting_queue: prefix_ids = r.origin_input_ids + r.output_ids extra_key = r.extra_key - # NOTE: the prefix_indices must always be aligned with last_node match_result = self.tree_cache.match_prefix( rid=r.rid, key=RadixKey(token_ids=prefix_ids, extra_key=extra_key) @@ -255,18 +255,16 @@ class SchedulePolicy: def _sort_by_longest_output( waiting_queue: List[Req], enable_priority_scheduling: bool, - schedule_low_priority_values_first: bool, + priority_sign: int, ) -> None: """Sorts the waiting queue based on the longest output (max_new_tokens). If using priority scheduling, sort by priority first.""" if enable_priority_scheduling: - if schedule_low_priority_values_first: - waiting_queue.sort( - key=lambda x: (x.priority, -x.sampling_params.max_new_tokens) - ) - else: - waiting_queue.sort( - key=lambda x: (-x.priority, -x.sampling_params.max_new_tokens) + waiting_queue.sort( + key=lambda x: ( + x.priority * priority_sign, + -x.sampling_params.max_new_tokens, ) + ) else: waiting_queue.sort(key=lambda x: -x.sampling_params.max_new_tokens) @@ -277,17 +275,15 @@ class SchedulePolicy: @staticmethod def _sort_by_priority_and_fcfs( - waiting_queue: List[Req], schedule_low_priority_values_first: bool + waiting_queue: List[Req], priority_sign: int ) -> None: """Sorts the waiting queue based on the request priority then received titmestamp.""" - if schedule_low_priority_values_first: - waiting_queue.sort( - key=lambda x: (x.priority, x.time_stats.wait_queue_entry_time) - ) - else: - waiting_queue.sort( - key=lambda x: (-x.priority, x.time_stats.wait_queue_entry_time) + waiting_queue.sort( + key=lambda x: ( + x.priority * priority_sign, + x.time_stats.wait_queue_entry_time, ) + ) @staticmethod def _calc_weight(cur_node: TreeNode, node_to_weight: Dict[TreeNode, int]) -> None: @@ -340,7 +336,6 @@ class PrefillAdder: self.rem_chunk_tokens = rem_chunk_tokens if self.rem_chunk_tokens is not None: self.rem_chunk_tokens -= mixed_with_decode_tokens - self.rem_total_token_offset = mixed_with_decode_tokens self.cur_rem_token_offset = mixed_with_decode_tokens @@ -399,7 +394,6 @@ class PrefillAdder: self.token_to_kv_pool_allocator.available_size() + self.tree_cache.evictable_size() ) - return available_and_evictable - self.rem_total_token_offset @property @@ -677,19 +671,20 @@ class PrefillAdder: Returns True if preemption was committed, and the new request can be scheduled. """ # Iterate running requests to find preemptible requests + priority_sign = 1 if server_args.schedule_low_priority_values_first else -1 + valid_running_reqs = ( r for r in self.running_batch.reqs if r not in self.preempt_list ) - if server_args.schedule_low_priority_values_first: - sorted_valid_running_reqs = sorted( - valid_running_reqs, - key=lambda x: (-x.priority, -x.time_stats.wait_queue_entry_time), - ) - else: - sorted_valid_running_reqs = sorted( - valid_running_reqs, - key=lambda x: (x.priority, -x.time_stats.wait_queue_entry_time), - ) + + sorted_valid_running_reqs = sorted( + valid_running_reqs, + key=lambda x: ( + x.priority * (-priority_sign), + -x.time_stats.wait_queue_entry_time, + ), + ) + preemptible_reqs = [] min_tokens_to_remove = ( req.extend_input_len @@ -698,9 +693,8 @@ class PrefillAdder: ) for running_req in sorted_valid_running_reqs: # Priority difference needs to meet the threshold to be preemptible. - priority_diff = req.priority - running_req.priority - if server_args.schedule_low_priority_values_first: - priority_diff *= -1 + priority_diff = (req.priority - running_req.priority) * (-priority_sign) + if priority_diff > self.priority_scheduling_preemption_threshold: preemptible_reqs.append(running_req) min_tokens_to_remove -= self._get_running_request_total_token_offset( diff --git a/test/srt/run_suite.py b/test/srt/run_suite.py index 5291f9fea..6574e2183 100644 --- a/test/srt/run_suite.py +++ b/test/srt/run_suite.py @@ -88,6 +88,7 @@ suites = { TestFile("test_original_logprobs.py", 41), TestFile("test_page_size.py", 60), TestFile("test_penalty.py", 82), + TestFile("test_prefill_adder.py", 1), TestFile("test_priority_scheduling.py", 130), TestFile("test_pytorch_sampling_backend.py", 66), TestFile("test_radix_attention.py", 105), diff --git a/test/srt/test_prefill_adder.py b/test/srt/test_prefill_adder.py new file mode 100644 index 000000000..a7308a401 --- /dev/null +++ b/test/srt/test_prefill_adder.py @@ -0,0 +1,306 @@ +import unittest +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +from sglang.srt.managers.schedule_batch import Req +from sglang.srt.managers.schedule_policy import PrefillAdder +from sglang.test.test_utils import CustomTestCase + + +class TestPrefillAdder(CustomTestCase): + def setUp(self): + self.mock_tree_cache = self.create_tree_cache() + self.mock_token_allocator = self.create_token_allocator() + patcher = patch( + "sglang.srt.managers.schedule_policy.is_nsa_enable_prefill_cp", + return_value=False, + ) + self.mock_is_nsa = patcher.start() + self.addCleanup(patcher.stop) + + def create_tree_cache( + self, + *, + full_evictable_size: int = 0, + swa_evictable_size: int = 0, + evictable_size: int = 0, + ) -> MagicMock: + tree_cache = MagicMock() + tree_cache.full_evictable_size.return_value = full_evictable_size + tree_cache.swa_evictable_size.return_value = swa_evictable_size + tree_cache.evictable_size.return_value = evictable_size + tree_cache.disable = False + tree_cache.inc_lock_ref.return_value = None + tree_cache.dec_lock_ref.return_value = None + return tree_cache + + def create_token_allocator( + self, + *, + full_available_size: int = 0, + swa_available_size: int = 0, + available_size: int = 0, + ) -> MagicMock: + allocator = MagicMock() + allocator.full_available_size.return_value = full_available_size + allocator.swa_available_size.return_value = swa_available_size + allocator.available_size.return_value = available_size + return allocator + + def create_running_batch(self, reqs=None) -> MagicMock: + batch = MagicMock() + batch.reqs = list(reqs or []) + batch.release_req.return_value = None + batch.filter_batch.return_value = None + return batch + + def create_server_args( + self, *, schedule_low_priority_values_first: bool + ) -> MagicMock: + server_args = MagicMock() + server_args.schedule_low_priority_values_first = ( + schedule_low_priority_values_first + ) + return server_args + + def create_mock_req(self, rid, priority, max_new_tokens, output_len=0, wait_time=0): + req = MagicMock(spec=Req) + req.rid = str(rid) + req.priority = priority + req.extend_input_len = 0 + req.output_ids = [0] * output_len + req.sampling_params = SimpleNamespace(max_new_tokens=max_new_tokens) + req.time_stats = SimpleNamespace(wait_queue_entry_time=wait_time) + return req + + def create_adder(self, running_batch): + return PrefillAdder( + page_size=1, + tree_cache=self.mock_tree_cache, + token_to_kv_pool_allocator=self.mock_token_allocator, + running_batch=running_batch, + new_token_ratio=1.0, + rem_input_tokens=10000, + rem_chunk_tokens=None, + mixed_with_decode_tokens=0, + priority_scheduling_preemption_threshold=0, + ) + + def test_preempt_success_high_priority_values_first(self): + params = [ + ("run1", 0, 50), + ("run2", 1, 75), + ("run3", 2, 100), + ] + running_reqs = [ + self.create_mock_req(rid, priority, max_new_tokens) + for rid, priority, max_new_tokens in params + ] + mock_server_args = self.create_server_args( + schedule_low_priority_values_first=False + ) + running_batch = self.create_running_batch(running_reqs) + adder = self.create_adder(running_batch) + + self.assertEqual(adder.rem_total_token_offset, 225) + + self.mock_token_allocator.full_available_size.return_value = ( + 225 # full occupation of GRam + ) + self.mock_token_allocator.available_size.return_value = 225 + + new_req = self.create_mock_req("new1", priority=1, max_new_tokens=49) + + success = adder.preempt_to_schedule(new_req, mock_server_args) + + self.assertTrue(success) + self.assertIn(running_reqs[0], adder.preempt_list) + self.assertEqual(adder.rem_total_token_offset, 175) # 50 + 75 + 100 - 50 = 175 + running_batch.release_req.assert_called_once() + + def test_preempt_success_low_priority_values_first(self): + params = [ + ("run1", 0, 50), + ("run2", 1, 75), + ("run3", 2, 100), + ] + running_reqs = [ + self.create_mock_req(rid, priority, max_new_tokens) + for rid, priority, max_new_tokens in params + ] + mock_server_args = self.create_server_args( + schedule_low_priority_values_first=True + ) + running_batch = self.create_running_batch(running_reqs) + adder = self.create_adder(running_batch) + + self.assertEqual(adder.rem_total_token_offset, 225) + + self.mock_token_allocator.full_available_size.return_value = ( + 225 # full occupation of GRam + ) + self.mock_token_allocator.available_size.return_value = 225 + + new_req = self.create_mock_req("new1", priority=1, max_new_tokens=49) + + success = adder.preempt_to_schedule(new_req, mock_server_args) + + self.assertTrue(success) + self.assertIn(running_reqs[2], adder.preempt_list) + self.assertEqual(adder.rem_total_token_offset, 125) # 50 + 75 + 100 - 100 = 125 + running_batch.release_req.assert_called_once() + + def test_preempt_fail_low_priority_values_first(self): + params = [ + ("run1", 0, 50), + ("run2", 1, 75), + ("run3", 2, 100), + ] + running_reqs = [ + self.create_mock_req(rid, priority, max_new_tokens) + for rid, priority, max_new_tokens in params + ] + mock_server_args = self.create_server_args( + schedule_low_priority_values_first=True + ) + running_batch = self.create_running_batch(running_reqs) + adder = self.create_adder(running_batch) + + self.assertEqual(adder.rem_total_token_offset, 225) + + self.mock_token_allocator.full_available_size.return_value = ( + 225 # full occupation of GRam + ) + self.mock_token_allocator.available_size.return_value = 225 + + new_req_fail_by_priority_check = self.create_mock_req( + "new1", priority=2, max_new_tokens=49 + ) + + success_by_priority_check = adder.preempt_to_schedule( + new_req_fail_by_priority_check, mock_server_args + ) + self.assertFalse(success_by_priority_check) + + new_req_fail_by_priority_check = self.create_mock_req( + "new2", priority=1, max_new_tokens=110 + ) + success_by_capacity_check = adder.preempt_to_schedule( + new_req_fail_by_priority_check, mock_server_args + ) + self.assertFalse(success_by_capacity_check) + + def test_preempt_fail_high_priority_values_first(self): + params = [ + ("run1", 0, 50), + ("run2", 1, 75), + ("run3", 2, 100), + ] + running_reqs = [ + self.create_mock_req(rid, priority, max_new_tokens) + for rid, priority, max_new_tokens in params + ] + mock_server_args = self.create_server_args( + schedule_low_priority_values_first=False + ) + running_batch = self.create_running_batch(running_reqs) + adder = self.create_adder(running_batch) + + self.assertEqual(adder.rem_total_token_offset, 225) + + self.mock_token_allocator.full_available_size.return_value = ( + 225 # full occupation of GRam + ) + self.mock_token_allocator.available_size.return_value = 225 + + new_req_fail_by_priority_check = self.create_mock_req( + "new1", priority=0, max_new_tokens=49 + ) + + success_by_priority_check = adder.preempt_to_schedule( + new_req_fail_by_priority_check, mock_server_args + ) + self.assertFalse(success_by_priority_check) + + new_req_fail_by_priority_check = self.create_mock_req( + "new2", priority=-1, max_new_tokens=110 + ) + success_by_capacity_check = adder.preempt_to_schedule( + new_req_fail_by_priority_check, mock_server_args + ) + self.assertFalse(success_by_capacity_check) + + def test_preempt_success_low_priority_values_first_exact_once(self): + params = [ + ("run1", 0, 50), + ("run2", 1, 75), + ("run3", 2, 100), + ("run4", 2, 125), + ("run4", 2, 125), + ] + running_reqs = [ + self.create_mock_req(rid, priority, max_new_tokens) + for rid, priority, max_new_tokens in params + ] + mock_server_args = self.create_server_args( + schedule_low_priority_values_first=True + ) + running_batch = self.create_running_batch(running_reqs) + adder = self.create_adder(running_batch) + + self.assertEqual(adder.rem_total_token_offset, 475) + + self.mock_token_allocator.full_available_size.return_value = ( + 475 # full occupation of GRam + ) + self.mock_token_allocator.available_size.return_value = 475 + + new_req = self.create_mock_req("new1", priority=1, max_new_tokens=75) + + success = adder.preempt_to_schedule(new_req, mock_server_args) + self.assertTrue(success) + self.assertIn(running_reqs[2], adder.preempt_list) + self.assertEqual( + adder.rem_total_token_offset, 375 + ) # 50 + 75 + 100 + 125 + 125 - 100 = 375 + running_batch.release_req.assert_called_once() + + def test_preempt_success_low_priority_values_first_exact_twice(self): + params = [ + ("run1", 0, 50), + ("run2", 1, 75), + ("run3", 2, 100), + ("run4", 2, 125), + ("run4", 2, 125), + ] + running_reqs = [ + self.create_mock_req(rid, priority, max_new_tokens) + for rid, priority, max_new_tokens in params + ] + mock_server_args = self.create_server_args( + schedule_low_priority_values_first=True + ) + running_batch = self.create_running_batch(running_reqs) + adder = self.create_adder(running_batch) + + self.assertEqual(adder.rem_total_token_offset, 475) + + self.mock_token_allocator.full_available_size.return_value = ( + 475 # full occupation of GRam + ) + self.mock_token_allocator.available_size.return_value = 475 + + new_req = self.create_mock_req("new1", priority=1, max_new_tokens=200) + + success = adder.preempt_to_schedule(new_req, mock_server_args) + self.assertTrue(success) + self.assertIn(running_reqs[2], adder.preempt_list) + self.assertIn(running_reqs[3], adder.preempt_list) + self.assertEqual( + adder.rem_total_token_offset, 250 + ) # 50 + 75 + 100 + 125 + 125 - 100 - 125 = 250 + self.assertEqual(running_batch.release_req.call_count, 2) + + +if __name__ == "__main__": + unittest.main()