diff --git a/python/sglang/srt/managers/schedule_policy.py b/python/sglang/srt/managers/schedule_policy.py index ad2abae2f..579a74149 100644 --- a/python/sglang/srt/managers/schedule_policy.py +++ b/python/sglang/srt/managers/schedule_policy.py @@ -417,6 +417,7 @@ class PrefillAdder: self.can_run_list = [] self.preempt_list = [] self.new_chunked_req = None + self.has_chunked_req_in_batch = False self.log_hit_tokens = 0 # TODO(lsyin): report the real input tokens excluding page alignment self.log_input_tokens = 0 @@ -450,6 +451,28 @@ class PrefillAdder: self.cp_shared_kv_prefill_max_total_extend_tokens = ( cp_shared_kv_prefill_max_total_extend_tokens ) + if ( + self._is_cp_prefill_context() + and self.enable_cp_shared_kv_prefill_bs_gt1 + and self.cp_shared_kv_prefill_max_total_extend_tokens is not None + ): + # CP shared-KV bs>1 owns its batch-size/token admission via + # cp_shared_kv_prefill_max_total_extend_tokens. Do not let the + # generic max_prefill_tokens default silently cap this path below + # the CP-specific limit. Keep a larger generic budget intact when + # it is configured, because the CP-specific checker below is the + # actual CP grouping limit. The extra page keeps exact-fill CP + # batches from tripping the legacy generic >= budget check. + # Allocator capacity is still enforced by rem_total_tokens, + # cur_rem_tokens, and prepare_for_extend(). + self.rem_input_tokens = ( + max( + rem_input_tokens, + self.cp_shared_kv_prefill_max_total_extend_tokens + + self.page_size, + ) + - mixed_with_decode_tokens + ) self.cp_shared_kv_prefill_total_extend_tokens = 0 def _init_dllm_meta(self, dllm_config: DllmConfig): @@ -685,6 +708,7 @@ class PrefillAdder: req.set_extend_input_len(min(req.extend_input_len, _rem_tokens)) req.fill_ids = req.fill_ids[: len(req.prefix_indices) + req.extend_input_len] self.can_run_list.append(req) + self.has_chunked_req_in_batch = True self._update_prefill_budget( 0, req.extend_input_len, @@ -819,6 +843,12 @@ class PrefillAdder: ) ): return AddReqResult.OTHER + if ( + self._is_cp_prefill_context() + and len(self.can_run_list) > 0 + and (has_chunked_req or self.has_chunked_req_in_batch) + ): + return AddReqResult.OTHER if self._cp_prefill_multi_request_disabled(): return AddReqResult.OTHER if self._cp_prefill_request_limit_reached(): @@ -921,6 +951,7 @@ class PrefillAdder: self.can_run_list.append(req) self.new_chunked_req = req + self.has_chunked_req_in_batch = True self._req_inc_lock_ref(req) self._update_prefill_budget(prefix_len, trunc_len, 0) diff --git a/test/registered/unit/managers/test_prefill_adder.py b/test/registered/unit/managers/test_prefill_adder.py index b7f3a9ad7..7880d90c7 100644 --- a/test/registered/unit/managers/test_prefill_adder.py +++ b/test/registered/unit/managers/test_prefill_adder.py @@ -148,6 +148,9 @@ class TestPrefillAdder(CustomTestCase): req.fill_ids = list(range(extend_input_len)) req.last_node = object() req.sampling_params.ignore_eos = False + req.set_extend_input_len.side_effect = lambda value: setattr( + req, "extend_input_len", value + ) return req def create_adder(self, running_batch, **kwargs): @@ -671,6 +674,106 @@ class TestPrefillAdder(CustomTestCase): ) self.assertEqual([req.rid for req in oversized_first.can_run_list], ["large"]) + def test_cp_prefill_total_extend_limit_replaces_generic_input_budget(self): + set_global_server_args_for_scheduler( + ServerArgs( + model_path="dummy", + enable_nsa_prefill_context_parallel=True, + nsa_prefill_cp_mode="in-seq-split", + ) + ) + self.mock_token_allocator.available_size.return_value = 10000 + adder = self.create_adder( + self.create_running_batch(), + page_size=64, + # Simulates the generic max_prefill_tokens budget being smaller + # than the CP shared-KV bs>1 budget. + rem_input_tokens=192, + enable_cp_shared_kv_prefill_bs_gt1=True, + cp_shared_kv_prefill_max_batch_requests=8, + cp_shared_kv_prefill_max_total_extend_tokens=256, + ) + + first = self.create_prefill_req("first", extend_input_len=128) + second = self.create_prefill_req("second", extend_input_len=128) + + self.assertEqual( + adder.add_one_req(first, has_chunked_req=False, truncation_align_size=None), + AddReqResult.CONTINUE, + ) + second_result = adder.add_one_req( + second, has_chunked_req=False, truncation_align_size=None + ) + self.assertNotEqual(second_result, AddReqResult.NO_TOKEN) + self.assertEqual([req.rid for req in adder.can_run_list], ["first", "second"]) + self.assertEqual(adder.cp_shared_kv_prefill_total_extend_tokens, 256) + + def test_cp_prefill_total_extend_limit_does_not_bypass_allocator_capacity(self): + set_global_server_args_for_scheduler( + ServerArgs( + model_path="dummy", + enable_nsa_prefill_context_parallel=True, + nsa_prefill_cp_mode="in-seq-split", + ) + ) + self.mock_token_allocator.available_size.return_value = 200 + adder = self.create_adder( + self.create_running_batch(), + page_size=64, + rem_input_tokens=64, + enable_cp_shared_kv_prefill_bs_gt1=True, + cp_shared_kv_prefill_max_batch_requests=8, + cp_shared_kv_prefill_max_total_extend_tokens=4096, + ) + + first = self.create_prefill_req("first", extend_input_len=128) + second = self.create_prefill_req("second", extend_input_len=128) + + self.assertEqual( + adder.add_one_req(first, has_chunked_req=False, truncation_align_size=None), + AddReqResult.CONTINUE, + ) + self.assertEqual( + adder.add_one_req(second, has_chunked_req=False, truncation_align_size=None), + AddReqResult.NO_TOKEN, + ) + self.assertEqual([req.rid for req in adder.can_run_list], ["first"]) + + def test_cp_prefill_chunked_req_excludes_new_requests_even_when_bs_gt1_enabled(self): + set_global_server_args_for_scheduler( + ServerArgs( + model_path="dummy", + enable_nsa_prefill_context_parallel=True, + nsa_prefill_cp_mode="in-seq-split", + ) + ) + self.mock_token_allocator.available_size.return_value = 10000 + adder = self.create_adder( + self.create_running_batch(), + page_size=64, + rem_input_tokens=4096, + rem_chunk_tokens=256, + enable_cp_shared_kv_prefill_bs_gt1=True, + cp_shared_kv_prefill_max_batch_requests=8, + cp_shared_kv_prefill_max_total_extend_tokens=4096, + ) + + chunked = self.create_prefill_req("chunked", extend_input_len=128) + normal = self.create_prefill_req("normal", extend_input_len=128) + + adder.new_chunked_req = adder.add_chunked_req(chunked) + self.assertEqual([req.rid for req in adder.can_run_list], ["chunked"]) + self.assertIsNone(adder.new_chunked_req) + self.assertEqual(chunked.extend_input_len, 128) + + self.assertEqual( + adder.add_one_req( + normal, has_chunked_req=True, truncation_align_size=None + ), + AddReqResult.OTHER, + ) + self.assertEqual([req.rid for req in adder.can_run_list], ["chunked"]) + if __name__ == "__main__": unittest.main()