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:
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user