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:
laoyao0822
2026-06-10 22:21:46 +08:00
parent 342c552ab3
commit 4f65d7a176
4 changed files with 137 additions and 1 deletions
@@ -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()