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

View File

@@ -903,3 +903,35 @@ PYTHONPATH=python python -m pytest -q \
test/registered/unit/mem_cache/test_cp_hicache_load_back_owner_lanes.py
=> 10 passed, 3 warnings
```
## 18. 2026-06-04 scheduler admission gate 首版打开
目标:先允许 NSA in-seq CP shared-KV 在 scheduler 层组成真实 multi-request prefill batch同时用显式参数把首版风险收窄。
### 参数
1. `--enable-cp-shared-kv-prefill-bs-gt1`
- 默认关闭;关闭时 CP prefill 仍保持旧的 `batch_size == 1` admission gate。
- 打开后,且 `--enable-nsa-prefill-cp-shared-kv` 同时开启时scheduler 可以把多个 waiting requests 放进同一个 prefill batch普通 CP 不因该参数打开 bs>1。
2. `--cp-shared-kv-prefill-max-batch-requests`
- 可选正整数;限制同一个 CP shared-KV prefill batch 的 request 数。
- 不设置时只受通用 `--prefill-max-requests` 约束。
3. `--cp-shared-kv-prefill-max-total-extend-tokens`
- 可选正整数;限制同一个 CP shared-KV prefill batch 的 **page-aligned extend tokens** 总量。
- 单个 request 自身超过该值时仍允许独占 batch避免 scheduler deadlock真正是否能运行仍由 KV allocator 容量判断。
### 容量语义
scheduler admission 不做 bf16/fp8 字节级估算。原因是实际 KV pool 容量已经在初始化时按 dtype、page size、model shape、mem fraction 转换为 token/page capacity。admission 只消费现有 `PrefillAdder` 的 token/page budget
- `rem_input_tokens` / `rem_total_tokens` / `cur_rem_tokens` 继续作为通用容量预算;
- `alloc_for_extend()` 仍是最终容量仲裁点;
- CP owner-lane 不足仍通过 `KVCapacityWaitError` deferred不在 scheduler 新增 collective 或重复推导 dtype bytes。
这样 bf16 与 fp8 都走同一个 admission 逻辑,差异由底层 allocator 的实际可用 token capacity 体现。
### 风险与后续
1. `cp_shared_kv_prefill_max_total_extend_tokens` 首版使用 page-aligned extend 累计,和 allocator 的最小 page 单位一致;日志/指标里如果要展示用户侧 token需要另加 valid-token 统计。
2. 该 gate 只控制 scheduler 组 batch不解决所有 bs>1 runtime correctness打开 ETE 时仍应保留现有 fail-fast特别是 draft/EAGLE、logprob、hidden capture、compute padding 路径。
3. 如果 ETE 中 `KVCapacityWaitError` 频繁出现,下一步应把 owner-lane free pages/evictable pages 的 batch precheck 前移到 scheduler而不是引入 all-reduce。

View File

@@ -390,6 +390,9 @@ class PrefillAdder:
max_prefill_bs: int = 0,
max_running_requests: Optional[int] = None,
prefill_max_requests: Optional[int] = None,
enable_cp_shared_kv_prefill_bs_gt1: bool = False,
cp_shared_kv_prefill_max_batch_requests: Optional[int] = None,
cp_shared_kv_prefill_max_total_extend_tokens: Optional[int] = None,
prefill_delayer_single_pass: Optional[PrefillDelayerSinglePassExecutor] = None,
dllm_config: Optional[DllmConfig] = None,
):
@@ -440,6 +443,14 @@ class PrefillAdder:
self.prefill_max_requests = prefill_max_requests
self.prefill_delayer_single_pass = prefill_delayer_single_pass
self.max_prefill_bs = max_prefill_bs
self.enable_cp_shared_kv_prefill_bs_gt1 = enable_cp_shared_kv_prefill_bs_gt1
self.cp_shared_kv_prefill_max_batch_requests = (
cp_shared_kv_prefill_max_batch_requests
)
self.cp_shared_kv_prefill_max_total_extend_tokens = (
cp_shared_kv_prefill_max_total_extend_tokens
)
self.cp_shared_kv_prefill_total_extend_tokens = 0
def _init_dllm_meta(self, dllm_config: DllmConfig):
self.dllm_block_size = dllm_config.block_size
@@ -502,6 +513,45 @@ class PrefillAdder:
def ceil_paged_tokens(self, tokens: int) -> int:
return -(-tokens // self.page_size) * self.page_size
def _is_cp_prefill_context(self) -> bool:
return self.nsa_prefill_cp_in_seq_split or self.prefill_context_parallel_enabled
def _cp_prefill_multi_request_disabled(self) -> bool:
return (
self._is_cp_prefill_context()
and not self.enable_cp_shared_kv_prefill_bs_gt1
and len(self.can_run_list) >= 1
)
def _cp_prefill_request_limit_reached(self) -> bool:
if not (
self._is_cp_prefill_context()
and self.enable_cp_shared_kv_prefill_bs_gt1
):
return False
limit = self.cp_shared_kv_prefill_max_batch_requests
return limit is not None and len(self.can_run_list) >= limit
def _cp_prefill_extend_limit_exceeded(self, extend_input_len: int) -> bool:
if not (
self._is_cp_prefill_context()
and self.enable_cp_shared_kv_prefill_bs_gt1
):
return False
limit = self.cp_shared_kv_prefill_max_total_extend_tokens
if limit is None:
return False
projected = self.cp_shared_kv_prefill_total_extend_tokens + self.ceil_paged_tokens(
extend_input_len
)
# Do not deadlock a single large request. The limit bounds grouping, not
# the maximum legal request size; actual KV capacity remains allocator-owned.
return projected > limit and len(self.can_run_list) > 0
def _get_available_device_tokens_for_load_back(self) -> int:
if self.is_hybrid_swa:
return (
@@ -549,6 +599,8 @@ class PrefillAdder:
self.rem_total_token_offset += extend_input_len + max_new_tokens + page_overhead
self.cur_rem_token_offset += extend_input_len + page_overhead
self.rem_input_tokens -= extend_input_len
if self._is_cp_prefill_context():
self.cp_shared_kv_prefill_total_extend_tokens += extend_input_len
if self.dllm_config is not None:
self.rem_dllm_tokens -= extend_input_len
@@ -726,6 +778,8 @@ class PrefillAdder:
self.rem_chunk_tokens is None # chunked prefill is disabled
or req.extend_input_len <= self.rem_chunk_tokens # it is the last chunk
):
if self._cp_prefill_extend_limit_exceeded(req.extend_input_len):
return AddReqResult.OTHER
# Non-chunked prefill
self.can_run_list.append(req)
self._update_prefill_budget(
@@ -739,6 +793,8 @@ class PrefillAdder:
# Chunked prefill
trunc_len = self.rem_chunk_tokens
if self._cp_prefill_extend_limit_exceeded(trunc_len):
return AddReqResult.OTHER
req.set_extend_input_len(trunc_len)
req.fill_ids = req.fill_ids[:trunc_len]
@@ -763,12 +819,9 @@ class PrefillAdder:
)
):
return AddReqResult.OTHER
# TODO support cp with multiple requests
# Enabling context parallelism currently presents precision issues;
# therefore, the prefill-batch setting is temporarily set to 1.
if (
self.nsa_prefill_cp_in_seq_split or self.prefill_context_parallel_enabled
) and len(self.can_run_list) >= 1:
if self._cp_prefill_multi_request_disabled():
return AddReqResult.OTHER
if self._cp_prefill_request_limit_reached():
return AddReqResult.OTHER
if (x := self.prefill_max_requests) is not None and len(self.can_run_list) >= x:
@@ -815,6 +868,8 @@ class PrefillAdder:
if input_tokens >= self.rem_input_tokens and len(self.can_run_list) != 0:
return AddReqResult.OTHER
if self._cp_prefill_extend_limit_exceeded(input_tokens):
return AddReqResult.OTHER
if self.dllm_config is not None:
if self.rem_dllm_tokens <= 0:
@@ -857,6 +912,9 @@ class PrefillAdder:
trunc_len // truncation_align_size
)
if self._cp_prefill_extend_limit_exceeded(trunc_len):
return AddReqResult.OTHER
# Chunked prefill
req.set_extend_input_len(trunc_len)
req.fill_ids = req.fill_ids[: len(req.prefix_indices) + trunc_len]

View File

@@ -2337,6 +2337,16 @@ class Scheduler(
max_prefill_bs=self.max_prefill_bs,
max_running_requests=self.max_running_requests,
prefill_max_requests=self.server_args.prefill_max_requests,
enable_cp_shared_kv_prefill_bs_gt1=(
self.server_args.enable_cp_shared_kv_prefill_bs_gt1
and self.server_args.enable_nsa_prefill_cp_shared_kv
),
cp_shared_kv_prefill_max_batch_requests=(
self.server_args.cp_shared_kv_prefill_max_batch_requests
),
cp_shared_kv_prefill_max_total_extend_tokens=(
self.server_args.cp_shared_kv_prefill_max_total_extend_tokens
),
prefill_delayer_single_pass=prefill_delayer_single_pass,
dllm_config=self.dllm_config,
)

View File

@@ -672,6 +672,9 @@ class ServerArgs:
enable_nsa_prefill_context_parallel: bool = False
nsa_prefill_cp_mode: str = "round-robin-split"
enable_nsa_prefill_cp_shared_kv: bool = False
enable_cp_shared_kv_prefill_bs_gt1: bool = False
cp_shared_kv_prefill_max_batch_requests: Optional[int] = None
cp_shared_kv_prefill_max_total_extend_tokens: Optional[int] = None
enable_fused_qk_norm_rope: bool = False
enable_precise_embedding_interpolation: bool = False
enable_fused_moe_sum_all_reduce: bool = False
@@ -931,6 +934,22 @@ class ServerArgs:
"Other backends do not implement CP shared-KV logical-page-position "
"transfer mapping yet."
)
if (
self.cp_shared_kv_prefill_max_batch_requests is not None
and self.cp_shared_kv_prefill_max_batch_requests <= 0
):
raise ValueError(
"cp_shared_kv_prefill_max_batch_requests must be a positive "
"integer when specified."
)
if (
self.cp_shared_kv_prefill_max_total_extend_tokens is not None
and self.cp_shared_kv_prefill_max_total_extend_tokens <= 0
):
raise ValueError(
"cp_shared_kv_prefill_max_total_extend_tokens must be a positive "
"integer when specified."
)
def _handle_cp_hicache_layout_validation(self):
if not (
@@ -5710,6 +5729,40 @@ class ServerArgs:
"Only prefill CP with NSA+MLA is supported; decode CP remains disabled."
),
)
parser.add_argument(
"--enable-cp-shared-kv-prefill-bs-gt1",
action="store_true",
default=ServerArgs.enable_cp_shared_kv_prefill_bs_gt1,
help=(
"Enable experimental multi-request prefill batching for NSA "
"in-seq CP shared KV. Capacity is still enforced by the KV "
"pool allocator; use the CP shared-KV batch limits to bound "
"scheduler admission."
),
)
parser.add_argument(
"--cp-shared-kv-prefill-max-batch-requests",
type=int,
default=ServerArgs.cp_shared_kv_prefill_max_batch_requests,
help=(
"Maximum number of requests admitted into one NSA in-seq CP "
"shared-KV prefill batch when --enable-cp-shared-kv-prefill-bs-gt1 "
"is set. If unset, only the generic --prefill-max-requests limit "
"applies."
),
)
parser.add_argument(
"--cp-shared-kv-prefill-max-total-extend-tokens",
type=human_readable_int,
default=ServerArgs.cp_shared_kv_prefill_max_total_extend_tokens,
help=(
"Maximum page-aligned extend tokens admitted into one NSA in-seq "
"CP shared-KV prefill batch when --enable-cp-shared-kv-prefill-bs-gt1 "
"is set. A single request larger than this limit is still allowed "
"to run alone to avoid scheduler deadlock."
)
+ f"\n\n{human_readable_int.__doc__}",
)
parser.add_argument(
"--nsa-prefill-cp-mode",
type=str,

View File

@@ -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()

View File

@@ -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 = [
{