bugfix[schedule]: Refactor sort method and add related UT (#13576)

Co-authored-by: Yuxuan Wei <w1300012920@pku.edu.cn>
Co-authored-by: Yuxuan Wei <w1300012920@gmail.com>
Co-authored-by: Yuxuan Wei🚚 <yuxwei@microsoft.com>
This commit is contained in:
SeanWeiSean
2025-12-23 03:13:57 +08:00
committed by GitHub
parent 3c882db3ad
commit 7759716786
3 changed files with 335 additions and 34 deletions

View File

@@ -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(

View File

@@ -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),

View File

@@ -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()