Gate CP shared-KV prefill batching behind explicit limits

The scheduler can now admit multi-request NSA in-seq CP shared-KV prefill batches only when the shared-KV bs>1 flag is explicitly enabled. The gate is still disabled by default and is scoped to CP shared-KV so ordinary CP is not widened accidentally.

Batch admission is bounded by optional request-count and page-aligned extend-token limits while real memory capacity remains allocator-owned. This keeps bf16 and fp8 on the same scheduler path because dtype differences are already reflected in KV pool token/page capacity.

Constraint: bs>1 runtime paths remain guarded by existing CP shared-KV fail-fast checks.

Constraint: Scheduler must not duplicate bf16/fp8 byte-level capacity estimation.

Rejected: Open the old CP gate unconditionally | ordinary CP would inherit an unverified shared-KV-specific batching path.

Rejected: Treat the extend-token cap as a hard per-request limit | a single large request could deadlock the scheduler.

Confidence: medium

Scope-risk: moderate

Directive: Keep CP shared-KV batching gated until ETE validates EAGLE accept length, output length, and HiCache load/backup behavior under real traffic.

Tested: local py_compile for server_args, schedule_policy, scheduler, prefill_adder tests, and server_args tests.

Tested: remote g0034 py_compile for the same files.

Tested: remote g0034 pytest target set: 5 passed for parser, parameter validation, default single-request CP gate, enabled bs>1 gate, and page-aligned extend cap.

Tested: remote g0034 pytest test_prefill_adder.py => 13 passed.

Not-tested: full server_args test file has an unrelated HuggingFace DNS/config-download failure in TestPrepareServerArgs.test_prepare_server_args.

Not-tested: ETE production traffic with --enable-cp-shared-kv-prefill-bs-gt1.
This commit is contained in:
laoyao0822
2026-06-04 04:17:32 +08:00
parent d7723aca07
commit 108e7d866d
6 changed files with 319 additions and 6 deletions
@@ -140,6 +140,16 @@ class TestPrefillAdder(CustomTestCase):
req.finished.return_value = False
return req
def create_prefill_req(self, rid, extend_input_len, max_new_tokens=1):
req = self.create_mock_req(rid, priority=0, max_new_tokens=max_new_tokens)
req.extend_input_len = extend_input_len
req.host_hit_length = 0
req.prefix_indices = torch.empty((0,), dtype=torch.int64)
req.fill_ids = list(range(extend_input_len))
req.last_node = object()
req.sampling_params.ignore_eos = False
return req
def create_adder(self, running_batch, **kwargs):
defaults = dict(
page_size=1,
@@ -550,6 +560,117 @@ class TestPrefillAdder(CustomTestCase):
self.assertEqual(quota, 90000 + 1024 - 65536 - 64)
def test_cp_prefill_gate_keeps_single_request_by_default(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,
)
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.OTHER,
)
self.assertEqual([req.rid for req in adder.can_run_list], ["first"])
def test_cp_prefill_gate_allows_batched_requests_when_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,
enable_cp_shared_kv_prefill_bs_gt1=True,
cp_shared_kv_prefill_max_batch_requests=2,
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)
third = self.create_prefill_req("third", extend_input_len=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.CONTINUE,
)
self.assertEqual(
adder.add_one_req(third, has_chunked_req=False, truncation_align_size=None),
AddReqResult.OTHER,
)
self.assertEqual([req.rid for req in adder.can_run_list], ["first", "second"])
def test_cp_prefill_total_extend_limit_is_page_aligned_and_allows_first_req(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=128,
)
first = self.create_prefill_req("first", extend_input_len=65)
second = self.create_prefill_req("second", extend_input_len=1)
oversized_first = 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=64,
)
large = self.create_prefill_req("large", extend_input_len=128)
self.assertEqual(
adder.add_one_req(first, has_chunked_req=False, truncation_align_size=None),
AddReqResult.CONTINUE,
)
self.assertEqual(adder.cp_shared_kv_prefill_total_extend_tokens, 128)
self.assertEqual(
adder.add_one_req(second, has_chunked_req=False, truncation_align_size=None),
AddReqResult.OTHER,
)
self.assertEqual(
oversized_first.add_one_req(
large, has_chunked_req=False, truncation_align_size=None
),
AddReqResult.CONTINUE,
)
self.assertEqual([req.rid for req in oversized_first.can_run_list], ["large"])
if __name__ == "__main__":
unittest.main()
@@ -53,6 +53,28 @@ def test_enable_nsa_prefill_cp_shared_kv_parser_flag():
assert args.enable_nsa_prefill_cp_shared_kv is True
def test_cp_shared_kv_prefill_bs_gt1_parser_limits():
import argparse
parser = argparse.ArgumentParser()
ServerArgs.add_cli_args(parser)
raw_args = parser.parse_args(
[
"--model-path",
"dummy",
"--enable-cp-shared-kv-prefill-bs-gt1",
"--cp-shared-kv-prefill-max-batch-requests",
"4",
"--cp-shared-kv-prefill-max-total-extend-tokens",
"8192",
]
)
args = ServerArgs.from_cli_args(raw_args)
assert args.enable_cp_shared_kv_prefill_bs_gt1 is True
assert args.cp_shared_kv_prefill_max_batch_requests == 4
assert args.cp_shared_kv_prefill_max_total_extend_tokens == 8192
class TestLoadBalanceMethod(unittest.TestCase):
def test_non_pd_defaults_to_round_robin(self):
server_args = ServerArgs(model_path="dummy", disaggregation_mode="null")
@@ -438,6 +460,23 @@ class TestHiCacheArgs(CustomTestCase):
enable_hisparse=False,
)
def test_cp_shared_kv_prefill_batch_limits_must_be_positive(self):
with self.assertRaisesRegex(
ValueError, "cp_shared_kv_prefill_max_batch_requests.*positive"
):
ServerArgs(
model_path="dummy",
cp_shared_kv_prefill_max_batch_requests=0,
)
with self.assertRaisesRegex(
ValueError, "cp_shared_kv_prefill_max_total_extend_tokens.*positive"
):
ServerArgs(
model_path="dummy",
cp_shared_kv_prefill_max_total_extend_tokens=0,
)
def test_hicache_io_backend_and_mem_layout_compatibility(self):
cases = [
{