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:
@@ -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 = [
|
||||
{
|
||||
|
||||
Reference in New Issue
Block a user