Decouple CP prefill batching from generic token cap

CP shared-KV bs>1 uses cp_shared_kv_prefill_max_total_extend_tokens as its grouping admission limit, but the generic max_prefill_tokens budget could still stop batching earlier. Raise only the legacy input-token admission budget for this CP path while keeping allocator-owned capacity checks unchanged.

Constraint: CP shared-KV bs>1 needs large cache-hit batches without relying on the generic max_prefill_tokens default.

Constraint: Allocator capacity must remain enforced by rem_total_tokens, cur_rem_tokens, and prepare_for_extend().

Rejected: Increase server-wide max_prefill_tokens | would change generic scheduler behavior and non-CP paths.

Confidence: high

Scope-risk: narrow

Directive: Do not use max_prefill_tokens as the CP shared-KV bs>1 grouping limit; use the CP-specific total extend token knob.

Tested: Local py_compile for schedule_policy.py and test_prefill_adder.py.

Tested: Remote pytest targeted CP PrefillAdder cases: 5 passed.

Tested: Remote pytest test/registered/unit/managers/test_prefill_adder.py: 16 passed.

Not-tested: Full ETE replay after this scheduler-only change.
This commit is contained in:
laoyao0822
2026-06-10 21:36:22 +08:00
parent afdd1d0992
commit 342c552ab3
2 changed files with 134 additions and 0 deletions
@@ -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()