Bound CP prefill batches by cached-token pressure
High cache-hit CP shared-KV batches can have small extend tokens while still carrying substantial prefix/load-back work. Add a CP-specific cached-token admission limit so operators can bound that pressure independently from extend-token batching. Constraint: The limit must not deadlock a single high-cache-hit request; it only stops adding additional requests to a non-empty batch. Constraint: Cached tokens are counted after L2 load-back planning via prefix_len, so L1 hits and successful L2 hits share one scheduler budget. Rejected: Reuse max_prefill_tokens | it limits generic input budget and does not represent cached-token work. Rejected: Count only L1 prefix before load-back | would miss L2 hit pressure, which is one of the target cases. Confidence: high Scope-risk: moderate Directive: Keep cached-token and extend-token limits separate; they bound different scheduler costs. Tested: Remote pytest targeted cached-token PrefillAdder cases: 2 passed. Tested: Remote pytest test/registered/unit/managers/test_prefill_adder.py: 18 passed. Tested: Remote ServerArgs CP validation smoke: SERVER_ARGS_CACHED_LIMIT_OK. Not-tested: Full ETE replay with a production cached-token limit value.
This commit is contained in:
@@ -774,6 +774,75 @@ class TestPrefillAdder(CustomTestCase):
|
||||
)
|
||||
self.assertEqual([req.rid for req in adder.can_run_list], ["chunked"])
|
||||
|
||||
def test_cp_prefill_total_cached_limit_stops_second_cached_request(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,
|
||||
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,
|
||||
cp_shared_kv_prefill_max_total_cached_tokens=4096,
|
||||
)
|
||||
|
||||
first = self.create_prefill_req("first", extend_input_len=64)
|
||||
first.prefix_indices = torch.arange(4096, dtype=torch.int64)
|
||||
first.fill_ids = list(range(4096 + 64))
|
||||
second = self.create_prefill_req("second", extend_input_len=64)
|
||||
second.prefix_indices = torch.arange(4096, dtype=torch.int64)
|
||||
second.fill_ids = list(range(4096 + 64))
|
||||
|
||||
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.OTHER,
|
||||
)
|
||||
self.assertEqual([req.rid for req in adder.can_run_list], ["first"])
|
||||
self.assertEqual(adder.cp_shared_kv_prefill_total_cached_tokens, 4096)
|
||||
|
||||
def test_cp_prefill_total_cached_limit_allows_single_oversized_cached_request(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 = 20000
|
||||
adder = self.create_adder(
|
||||
self.create_running_batch(),
|
||||
page_size=64,
|
||||
rem_input_tokens=4096,
|
||||
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,
|
||||
cp_shared_kv_prefill_max_total_cached_tokens=4096,
|
||||
)
|
||||
|
||||
oversized = self.create_prefill_req("oversized", extend_input_len=64)
|
||||
oversized.prefix_indices = torch.arange(8192, dtype=torch.int64)
|
||||
oversized.fill_ids = list(range(8192 + 64))
|
||||
|
||||
self.assertEqual(
|
||||
adder.add_one_req(
|
||||
oversized, has_chunked_req=False, truncation_align_size=None
|
||||
),
|
||||
AddReqResult.CONTINUE,
|
||||
)
|
||||
self.assertEqual([req.rid for req in adder.can_run_list], ["oversized"])
|
||||
self.assertEqual(adder.cp_shared_kv_prefill_total_cached_tokens, 8192)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user