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:
@@ -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(
|
||||
|
||||
@@ -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),
|
||||
|
||||
306
test/srt/test_prefill_adder.py
Normal file
306
test/srt/test_prefill_adder.py
Normal 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()
|
||||
Reference in New Issue
Block a user