add mixed chunk unit test and make small refactors (#18776)

This commit is contained in:
Glen Liu
2026-03-08 07:56:09 -04:00
committed by GitHub
parent ab9de886c5
commit 1eed219f47

View File

@@ -3,7 +3,7 @@ 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.srt.managers.schedule_policy import AddReqResult, PrefillAdder
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.test_utils import CustomTestCase
@@ -79,8 +79,8 @@ class TestPrefillAdder(CustomTestCase):
req.finished.return_value = False
return req
def create_adder(self, running_batch):
return PrefillAdder(
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,
@@ -91,6 +91,8 @@ class TestPrefillAdder(CustomTestCase):
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 = [
@@ -350,6 +352,96 @@ class TestPrefillAdder(CustomTestCase):
) # 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()