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

View File

@@ -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)

View File

@@ -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()