Files
sglang/test/registered/unit/managers/test_prefill_adder.py
2026-03-23 00:18:45 -07:00

448 lines
16 KiB
Python

import unittest
from types import SimpleNamespace
from unittest.mock import MagicMock
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.managers.schedule_policy import AddReqResult, PrefillAdder
from sglang.srt.mem_cache.base_prefix_cache import (
DecLockRefResult,
IncLockRefResult,
)
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=1, suite="stage-b-test-1-gpu-small")
register_amd_ci(est_time=2, suite="stage-b-test-1-gpu-small-amd")
class TestPrefillAdder(CustomTestCase):
def setUp(self):
set_global_server_args_for_scheduler(ServerArgs(model_path="dummy"))
self.mock_tree_cache = self.create_tree_cache()
self.mock_token_allocator = self.create_token_allocator()
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 = IncLockRefResult()
tree_cache.dec_lock_ref.return_value = DecLockRefResult()
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.extend_logprob_start_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)
req.finished.return_value = False
return req
def create_adder(self, running_batch, **kwargs):
defaults = dict(
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,
)
defaults.update(kwargs)
return PrefillAdder(**defaults)
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_skip_already_preempted_request(self):
params = [
("req_prio_0", 0, 50),
("req_prio_1", 1, 75),
("req_prio_2", 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
self.mock_token_allocator.available_size.return_value = 225
# New request preempts req_prio_0
first_req = self.create_mock_req(
"new_req_prio_1", priority=1, max_new_tokens=49
)
first_success = adder.preempt_to_schedule(first_req, mock_server_args)
self.assertTrue(first_success)
self.assertIn(running_reqs[0], adder.preempt_list)
self.assertEqual(adder.rem_total_token_offset, 175)
running_batch.release_req.assert_called_once()
# Second call needs more tokens than currently free, so it would need to
# preempt req_prio_0 again if already-preempted requests were not filtered out.
second_req = self.create_mock_req(
"second_new_req_prio_1", priority=1, max_new_tokens=76
)
second_success = adder.preempt_to_schedule(second_req, mock_server_args)
self.assertFalse(second_success)
self.assertEqual(adder.rem_total_token_offset, 175)
self.assertEqual(adder.preempt_list.count(running_reqs[0]), 1)
running_batch.release_req.assert_called_once()
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)
def test_mixed_chunk_prefill_budgets(self):
self.mock_token_allocator.available_size.return_value = 1000
decode_reqs = [
self.create_mock_req(f"decode_{i}", priority=0, max_new_tokens=50)
for i in range(8)
]
running_batch = self.create_running_batch(decode_reqs)
adder = self.create_adder(
running_batch,
rem_input_tokens=200,
rem_chunk_tokens=64,
mixed_with_decode_tokens=len(decode_reqs),
)
self.assertEqual(adder.rem_input_tokens, 192) # 200 - 8
self.assertEqual(adder.rem_chunk_tokens, 56) # 64 - 8
self.assertEqual(adder.rem_total_token_offset, 408) # 8 + 8 * 50
self.assertEqual(adder.cur_rem_token_offset, 8)
self.assertEqual(adder.budget_state(), AddReqResult.CONTINUE)
# Add a prefill that exactly consumes the chunk budget
req1 = self.create_mock_req("req1", priority=0, max_new_tokens=64)
req1.extend_input_len = 56
req1.host_hit_length = 0
req1.prefix_indices = []
req1.fill_ids = list(range(56))
req1.last_node = MagicMock()
req1.sampling_params.ignore_eos = False
result1 = adder.add_one_req(
req1, has_chunked_req=False, truncation_align_size=None
)
self.assertEqual(len(adder.can_run_list), 1)
self.assertEqual(adder.rem_chunk_tokens, 0) # 56 - 56
self.assertEqual(adder.rem_input_tokens, 136) # 192 - 56
self.assertEqual(result1, AddReqResult.OTHER)
# 3 decode requests finished
remaining_decode_reqs = decode_reqs[3:]
running_batch2 = self.create_running_batch(remaining_decode_reqs)
adder2 = self.create_adder(
running_batch2,
rem_input_tokens=200,
rem_chunk_tokens=64,
mixed_with_decode_tokens=len(remaining_decode_reqs),
)
self.assertEqual(adder2.rem_input_tokens, 195) # 200 - 5
self.assertEqual(adder2.rem_chunk_tokens, 59) # 64 - 5
self.assertEqual(adder2.rem_total_token_offset, 255) # 5 + 5 * 50
self.assertEqual(adder2.budget_state(), AddReqResult.CONTINUE)
# Same prefill no longer exhausts the chunk budget
req2 = self.create_mock_req("req2", priority=0, max_new_tokens=64)
req2.extend_input_len = 56
req2.host_hit_length = 0
req2.prefix_indices = []
req2.fill_ids = list(range(56))
req2.last_node = MagicMock()
req2.sampling_params.ignore_eos = False
result2 = adder2.add_one_req(
req2, has_chunked_req=False, truncation_align_size=None
)
self.assertEqual(len(adder2.can_run_list), 1)
self.assertEqual(adder2.rem_chunk_tokens, 3) # 59 - 56 = 3 remaining
self.assertEqual(result2, AddReqResult.CONTINUE)
# Fit last small prefill request
req3 = self.create_mock_req("req3", priority=0, max_new_tokens=16)
req3.extend_input_len = 3
req3.host_hit_length = 0
req3.prefix_indices = []
req3.fill_ids = list(range(3))
req3.last_node = MagicMock()
req3.sampling_params.ignore_eos = False
result3 = adder2.add_one_req(
req3, has_chunked_req=False, truncation_align_size=None
)
self.assertEqual(len(adder2.can_run_list), 2)
self.assertEqual(adder2.rem_chunk_tokens, 0) # 3 - 3 = 0
self.assertEqual(result3, AddReqResult.OTHER)
if __name__ == "__main__":
unittest.main()