diff --git a/python/sglang/srt/managers/schedule_policy.py b/python/sglang/srt/managers/schedule_policy.py index 10f0fe47b..aa72acfe5 100644 --- a/python/sglang/srt/managers/schedule_policy.py +++ b/python/sglang/srt/managers/schedule_policy.py @@ -382,7 +382,7 @@ class PrefillAdder: new_token_ratio: float, rem_input_tokens: int, rem_chunk_tokens: Optional[int], - mixed_with_decode_tokens: int = 0, + num_mixed_decode_tokens: int = 0, priority_scheduling_preemption_threshold: int = 0, max_prefill_bs: int = 0, max_running_requests: Optional[int] = None, @@ -395,7 +395,7 @@ class PrefillAdder: self.token_to_kv_pool_allocator = token_to_kv_pool_allocator self.running_batch = running_batch self.new_token_ratio = new_token_ratio - self.rem_input_tokens = rem_input_tokens - mixed_with_decode_tokens + self.rem_input_tokens = rem_input_tokens - num_mixed_decode_tokens self.rem_chunk_tokens = rem_chunk_tokens self.dllm_config = dllm_config @@ -403,9 +403,9 @@ class PrefillAdder: self._init_dllm_meta(dllm_config) 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 + self.rem_chunk_tokens -= num_mixed_decode_tokens + self.rem_total_token_offset = num_mixed_decode_tokens + self.cur_rem_token_offset = num_mixed_decode_tokens self.req_states = None self.can_run_list = [] @@ -416,6 +416,7 @@ class PrefillAdder: self.log_input_tokens = 0 if running_batch is not None: + # Estimate the offset in the remaining token space self.rem_total_token_offset += sum( [ self._get_running_request_total_token_offset(r) diff --git a/test/registered/unit/managers/test_prefill_adder.py b/test/registered/unit/managers/test_prefill_adder.py index b0b611e8e..85a4acf39 100644 --- a/test/registered/unit/managers/test_prefill_adder.py +++ b/test/registered/unit/managers/test_prefill_adder.py @@ -88,7 +88,7 @@ class TestPrefillAdder(CustomTestCase): new_token_ratio=1.0, rem_input_tokens=10000, rem_chunk_tokens=None, - mixed_with_decode_tokens=0, + num_mixed_decode_tokens=0, priority_scheduling_preemption_threshold=0, ) defaults.update(kwargs) @@ -365,7 +365,7 @@ class TestPrefillAdder(CustomTestCase): running_batch, rem_input_tokens=200, rem_chunk_tokens=64, - mixed_with_decode_tokens=len(decode_reqs), + num_mixed_decode_tokens=len(decode_reqs), ) self.assertEqual(adder.rem_input_tokens, 192) # 200 - 8 @@ -400,7 +400,7 @@ class TestPrefillAdder(CustomTestCase): running_batch2, rem_input_tokens=200, rem_chunk_tokens=64, - mixed_with_decode_tokens=len(remaining_decode_reqs), + num_mixed_decode_tokens=len(remaining_decode_reqs), ) self.assertEqual(adder2.rem_input_tokens, 195) # 200 - 5