From 3a4372721609c185f0a5eeea3447cdedebac8718 Mon Sep 17 00:00:00 2001 From: laoyao0822 Date: Thu, 11 Jun 2026 01:33:28 +0800 Subject: [PATCH] Bound CP prefill batching by estimated temp memory CP shared-KV bs>1 batching was only bounded by request count, extend tokens, and cached tokens. That left temporary GPU buffers such as MLA/index materialization, remap metadata, logits windows, and transfer descriptors implicit, and raw extend-token limits could exceed the active chunked-prefill budget.\n\nThis adds an explicit max-buffer-size admission gate with a CPU-only stream-aware estimator, wires it through PrefillAdder/Scheduler, performs a startup CUDA smoke allocation when configured, and reports the estimate in the scheduler admission benchmark. When chunked prefill is active, the effective CP extend-token limit is capped by the current chunk budget so the CP path does not advertise unreachable batch capacity or lift max-prefill-tokens too far.\n\nConstraint: Admission estimation must stay CPU-only on the scheduler hot path; CUDA allocation is limited to startup smoke checking.\nConstraint: Single oversized requests must still be allowed to run alone to avoid scheduler deadlock.\nRejected: Rely only on --max-prefill-tokens | it does not reliably bound the first oversized request and does not model cache-hit/load-back pressure.\nRejected: Let CP extend limit exceed chunked-prefill size | it creates an unreachable effective capacity and misleading budget lift.\nConfidence: medium\nScope-risk: moderate\nDirective: If bs>1 L1 prefetch is enabled later, update CPSharedKVPrefillBufferEstimatorContext.bs_gt1_l1_prefetch_enabled and include the live prefetch dense buffers in overlap windows.\nTested: local py_compile for touched files\nTested: local PYTHONPATH=python pytest -q test/registered/unit/managers/test_cp_shared_kv_prefill_buffer_estimator.py (4 passed)\nTested: remote cjy-glm5-new targeted pytest for new server_args, PrefillAdder, estimator, and benchmark cases (10 passed)\nTested: remote cjy-glm5-new PYTHONPATH=python pytest -q test/registered/unit/managers/test_cp_shared_kv_prefill_buffer_estimator.py test/registered/unit/managers/test_prefill_adder.py test/registered/unit/managers/test_prefill_scheduler_admission_bench.py (29 passed before chunk cap, then test_prefill_adder.py 21 passed after chunk cap)\nNot-tested: full server_args suite because existing TestPrepareServerArgs tries to reach HuggingFace and fails under container DNS/network\nNot-tested: GLM5 ETE smoke with --cp-shared-kv-prefill-max-buffer-size --- .../bench_prefill_scheduler_admission.py | 85 +- ...cp_bs_gt1_buffer_size_admission_plan_zh.md | 869 ++++++++++++++++++ .../cp_shared_kv_prefill_buffer_estimator.py | 249 +++++ python/sglang/srt/managers/schedule_policy.py | 139 +++ python/sglang/srt/managers/scheduler.py | 70 ++ python/sglang/srt/server_args.py | 42 + ...t_cp_shared_kv_prefill_buffer_estimator.py | 151 +++ .../unit/managers/test_prefill_adder.py | 115 +++ .../test_prefill_scheduler_admission_bench.py | 33 + .../unit/server_args/test_server_args.py | 56 ++ 10 files changed, 1806 insertions(+), 3 deletions(-) create mode 100644 docs/advanced_features/nsa_prefill_cp_bs_gt1_buffer_size_admission_plan_zh.md create mode 100644 python/sglang/srt/managers/cp_shared_kv_prefill_buffer_estimator.py create mode 100644 test/registered/unit/managers/test_cp_shared_kv_prefill_buffer_estimator.py diff --git a/benchmark/hicache/bench_prefill_scheduler_admission.py b/benchmark/hicache/bench_prefill_scheduler_admission.py index 1f9c545a3..2fc0e810e 100755 --- a/benchmark/hicache/bench_prefill_scheduler_admission.py +++ b/benchmark/hicache/bench_prefill_scheduler_admission.py @@ -143,6 +143,13 @@ class SchedulerBenchConfig: cp_shared_kv_prefill_max_batch_requests: Optional[int] = None cp_shared_kv_prefill_max_total_extend_tokens: Optional[int] = None cp_shared_kv_prefill_max_total_cached_tokens: Optional[int] = None + cp_shared_kv_prefill_max_buffer_size: Optional[int] = None + kv_cache_dim: int = 656 + kv_dtype_bytes: int = 1 + index_head_dim: int = 128 + vocab_size: int = 128_000 + tp_size: int = 8 + enable_bs_gt1_prefetch_estimate: bool = False max_ticks: int = 1 consume_l2_load_back_capacity: bool = True @@ -198,6 +205,8 @@ class TickResult: log_input_tokens: int allocator_available_after_tick: int load_back_events: list[LoadBackEvent] + cp_estimated_peak_buffer_bytes: int + cp_buffer_breakdown: dict[str, int] duration_us: float @@ -211,8 +220,15 @@ class TraceResult: class FakeTokenAllocator: - def __init__(self, available_tokens: int): + def __init__(self, available_tokens: int, cfg: SchedulerBenchConfig): self.available_tokens = int(available_tokens) + self.kvcache = SimpleNamespace( + kv_cache_dim=cfg.kv_cache_dim, + store_dtype=SimpleNamespace(itemsize=cfg.kv_dtype_bytes), + index_head_dim=cfg.index_head_dim, + quant_block_size=128, + index_k_with_scale_buffer_dtype=SimpleNamespace(itemsize=1), + ) def available_size(self) -> int: return self.available_tokens @@ -231,6 +247,9 @@ class FakeTokenAllocator: self.available_tokens -= tokens return True + def get_kvcache(self): + return self.kvcache + class FakeTreeCache: def __init__( @@ -382,6 +401,9 @@ def _configure_scheduler_globals(enable_cp_context: bool) -> None: def _make_prefill_adder(cfg: SchedulerBenchConfig, tree_cache: FakeTreeCache, allocator: FakeTokenAllocator): _install_sgl_kernel_stubs() + from sglang.srt.managers.cp_shared_kv_prefill_buffer_estimator import ( + CPSharedKVPrefillBufferEstimatorContext, + ) from sglang.srt.managers.schedule_policy import PrefillAdder return PrefillAdder( @@ -398,6 +420,18 @@ def _make_prefill_adder(cfg: SchedulerBenchConfig, tree_cache: FakeTreeCache, al cp_shared_kv_prefill_max_batch_requests=cfg.cp_shared_kv_prefill_max_batch_requests, cp_shared_kv_prefill_max_total_extend_tokens=cfg.cp_shared_kv_prefill_max_total_extend_tokens, cp_shared_kv_prefill_max_total_cached_tokens=cfg.cp_shared_kv_prefill_max_total_cached_tokens, + cp_shared_kv_prefill_max_buffer_size=cfg.cp_shared_kv_prefill_max_buffer_size, + cp_shared_kv_prefill_buffer_estimator_context=( + CPSharedKVPrefillBufferEstimatorContext( + kvcache=allocator.get_kvcache(), + model_config=SimpleNamespace(vocab_size=cfg.vocab_size), + tp_size=cfg.tp_size, + page_size=cfg.page_size, + logprob_chunk_enabled=False, + logprob_chunk_size=2048, + bs_gt1_l1_prefetch_enabled=cfg.enable_bs_gt1_prefetch_estimate, + ) + ), ) @@ -426,7 +460,7 @@ def run_scheduler_admission_trace( pending = list(requests) ticks: list[TickResult] = [] - allocator = FakeTokenAllocator(cfg.available_tokens) + allocator = FakeTokenAllocator(cfg.available_tokens, cfg) tree_cache = FakeTreeCache( allocator=allocator, page_size=cfg.page_size, @@ -484,6 +518,22 @@ def run_scheduler_admission_trace( accepted_rids = {req.rid for req in accepted} pending = [spec for spec in pending if spec.rid not in accepted_rids] + peak_estimate = adder.cp_shared_kv_prefill_last_buffer_estimate + peak_breakdown = ( + { + "layer_forward_peak_bytes": peak_estimate.layer_forward_peak_bytes, + "logits_window_peak_bytes": peak_estimate.logits_window_peak_bytes, + "load_back_window_peak_bytes": peak_estimate.load_back_window_peak_bytes, + "materialize_peak_bytes": peak_estimate.materialize_peak_bytes, + "prefetch_peak_bytes": peak_estimate.prefetch_peak_bytes, + "logits_peak_bytes": peak_estimate.logits_peak_bytes, + "remap_peak_bytes": peak_estimate.remap_peak_bytes, + "transfer_descriptor_peak_bytes": peak_estimate.transfer_descriptor_peak_bytes, + "backup_descriptor_peak_bytes": peak_estimate.backup_descriptor_peak_bytes, + } + if peak_estimate is not None + else {} + ) ticks.append( TickResult( tick=tick_idx, @@ -499,6 +549,10 @@ def run_scheduler_admission_trace( log_input_tokens=int(adder.log_input_tokens), allocator_available_after_tick=allocator.available_size(), load_back_events=list(tree_cache.load_back_events[load_event_start:]), + cp_estimated_peak_buffer_bytes=int( + adder.cp_shared_kv_prefill_estimated_peak_buffer_bytes + ), + cp_buffer_breakdown=peak_breakdown, duration_us=duration_us, ) ) @@ -566,7 +620,8 @@ def _print_text(trace: TraceResult) -> None: f"page_size={trace.config.page_size} available={trace.config.available_tokens} " f"evictable={trace.config.evictable_tokens} max_prefill={trace.config.max_prefill_tokens} " f"cp_extend_limit={trace.config.cp_shared_kv_prefill_max_total_extend_tokens} " - f"cp_cached_limit={trace.config.cp_shared_kv_prefill_max_total_cached_tokens}" + f"cp_cached_limit={trace.config.cp_shared_kv_prefill_max_total_cached_tokens} " + f"cp_buffer_limit={trace.config.cp_shared_kv_prefill_max_buffer_size}" ) for tick in trace.ticks: accepted = ",".join( @@ -578,6 +633,7 @@ def _print_text(trace: TraceResult) -> None: f"tick={tick.tick} bs={len(tick.accepted)} accepted=[{accepted}] " f"stop={tick.stopped_on_rid}:{tick.stopped_result} " f"cp_extend={tick.cp_total_extend_tokens} cp_cached={tick.cp_total_cached_tokens} " + f"cp_peak_buffer={tick.cp_estimated_peak_buffer_bytes} " f"log_hit={tick.log_hit_tokens} " f"log_input={tick.log_input_tokens} rem_input={tick.rem_input_tokens_after_tick} " f"rem_total={tick.rem_total_tokens_after_tick:.1f} " @@ -618,6 +674,18 @@ def build_arg_parser() -> argparse.ArgumentParser: parser.add_argument("--cp-max-batch-requests", type=int, default=8) parser.add_argument("--cp-max-total-extend-tokens", type=int, default=65_536) parser.add_argument("--cp-max-total-cached-tokens", type=int, default=None) + parser.add_argument( + "--cp-max-buffer-size", + type=float, + default=None, + help="CP shared-KV estimated temp-buffer gate in decimal GB.", + ) + parser.add_argument("--kv-cache-dim", type=int, default=656) + parser.add_argument("--kv-dtype-bytes", type=int, default=1) + parser.add_argument("--index-head-dim", type=int, default=128) + parser.add_argument("--vocab-size", type=int, default=128_000) + parser.add_argument("--tp-size", type=int, default=8) + parser.add_argument("--enable-bs-gt1-prefetch-estimate", action="store_true") parser.add_argument("--no-consume-l2-load-back-capacity", action="store_true") parser.add_argument("--output", choices=("text", "json"), default="text") return parser @@ -642,6 +710,17 @@ def main(argv: Optional[list[str]] = None) -> int: cp_shared_kv_prefill_max_batch_requests=args.cp_max_batch_requests, cp_shared_kv_prefill_max_total_extend_tokens=args.cp_max_total_extend_tokens, cp_shared_kv_prefill_max_total_cached_tokens=args.cp_max_total_cached_tokens, + cp_shared_kv_prefill_max_buffer_size=( + None + if args.cp_max_buffer_size is None + else int(args.cp_max_buffer_size * 1e9) + ), + kv_cache_dim=args.kv_cache_dim, + kv_dtype_bytes=args.kv_dtype_bytes, + index_head_dim=args.index_head_dim, + vocab_size=args.vocab_size, + tp_size=args.tp_size, + enable_bs_gt1_prefetch_estimate=args.enable_bs_gt1_prefetch_estimate, max_ticks=args.max_ticks, consume_l2_load_back_capacity=not args.no_consume_l2_load_back_capacity, ) diff --git a/docs/advanced_features/nsa_prefill_cp_bs_gt1_buffer_size_admission_plan_zh.md b/docs/advanced_features/nsa_prefill_cp_bs_gt1_buffer_size_admission_plan_zh.md new file mode 100644 index 000000000..38061db01 --- /dev/null +++ b/docs/advanced_features/nsa_prefill_cp_bs_gt1_buffer_size_admission_plan_zh.md @@ -0,0 +1,869 @@ +# NSA Prefill CP bs>1 max buffer size admission 计划 + +> 日期:2026-06-11 +> 分支:`cjy-cp-refactor` +> 范围:在 `--enable-cp-shared-kv-prefill-bs-gt1` 下新增 batch 级峰值 buffer size gate,同时保留 `--cp-shared-kv-prefill-max-total-extend-tokens` 和 `--cp-shared-kv-prefill-max-total-cached-tokens` 的现有语义。 + +## 0. 目标 + +新增一个以 GB 为默认单位的 CP shared-KV prefill batch admission 参数: + +```bash +--cp-shared-kv-prefill-max-buffer-size +``` + +建议字段名: + +```python +cp_shared_kv_prefill_max_buffer_size: Optional[int] +``` + +语义: + +1. 仅在 `enable_cp_shared_kv_prefill_bs_gt1 && enable_nsa_prefill_cp_shared_kv` 时生效。 +2. 默认单位是 GB:裸数字 `8` 表示 `8G`,显式后缀仍支持 `8G` / `8Gi` / `8192Mi`。内部字段统一保存为 bytes。 +3. 它是 batch 级 **峰值临时 buffer** admission gate,不替代 KV allocator capacity,也不替代现有 token gates。 +4. 现有三个限制同时生效: + - `cp_shared_kv_prefill_max_batch_requests` + - `cp_shared_kv_prefill_max_total_extend_tokens` + - `cp_shared_kv_prefill_max_total_cached_tokens` + - 新增 `cp_shared_kv_prefill_max_buffer_size` +5. 与现有 token gates 一样,若单个 request 自身超过该 limit,允许它单独运行,避免 scheduler deadlock;如果 batch 已非空,则停止继续加入。 +6. scheduler 启动完成、主要静态显存分配完成后,要按该 limit 做一次 CUDA buffer smoke allocation,提前暴露配置过大导致的 OOM。 + +## 1. 当前代码事实 + +### C1. admission 入口在 `PrefillAdder` + +相关文件: + +- `python/sglang/srt/managers/schedule_policy.py:378-481` +- `python/sglang/srt/managers/schedule_policy.py:554-600` +- `python/sglang/srt/managers/schedule_policy.py:638-662` +- `python/sglang/srt/managers/schedule_policy.py:916-936` + +当前 CP bs>1 admission 已有: + +```python +cp_shared_kv_prefill_max_batch_requests +cp_shared_kv_prefill_max_total_extend_tokens +cp_shared_kv_prefill_max_total_cached_tokens +``` + +其中: + +- extend limit 按 `ceil_paged_tokens(extend_input_len)` 累计。 +- cached limit 按 `ceil_paged_tokens(prefix_len)` 累计。 +- L2/HiCache hit 会先在 `add_one_req()` 里调用 `tree_cache.init_load_back()`,更新 `prefix_indices` 和 `extend_input_len`,再进入 CP gates。 + +结论:新增 buffer size gate 应该放在同一个位置:`init_load_back()` 后、`_update_prefill_budget()` 前。 + +### C2. scheduler 已经把 CP 参数传给 `PrefillAdder` + +相关文件: + +- `python/sglang/srt/managers/scheduler.py:2399-2424` +- `python/sglang/srt/server_args.py:677-680` +- `python/sglang/srt/server_args.py:5885-5921` + +结论:新增参数需要贯穿: + +```text +ServerArgs dataclass + -> ServerArgs._post_init validation + -> ServerArgs.add_cli_args + -> Scheduler.get_new_batch_prefill PrefillAdder(...) + -> PrefillAdder constructor / gate +``` + +### C3. exact L1 KV capacity 仍由 allocator 管 + +相关文件: + +- `python/sglang/srt/mem_cache/common.py:470-620` +- `python/sglang/srt/mem_cache/allocator.py:862-909` +- `python/sglang/srt/mem_cache/allocator.py:982-1045` + +`alloc_paged_token_slots_extend()` 会根据 CP owner-lane page owners 分配 L1 KV pages。这个路径已经包含: + +- page owner 规划; +- L1 free room eviction; +- exact capacity wait / fail-fast。 + +结论:`max_buffer_size` 不应该重复充当 L1 KV allocator 的 correctness gate。它主要限制 CUDA temporary / materialize / logits / descriptor 这类 batch 临时峰值。 + +### C4. KV pool 和 HiCache 的 per-token bytes 可从现有结构推导 + +相关文件: + +- `python/sglang/srt/mem_cache/memory_pool.py:1492-1580` +- `python/sglang/srt/mem_cache/memory_pool.py:1853-1938` +- `python/sglang/srt/mem_cache/memory_pool.py:2080-2084` +- `python/sglang/srt/mem_cache/hiradix_cache.py:69-94` + +NSA KV pool 中: + +- MLA KV buffer 形状约为 `(size + page_size, 1, kv_cache_dim)` per layer。 +- index buffer 形状约为 `(num_pages, page_size * (index_head_dim + index_head_dim / quant_block_size * 4))`。 +- host HiCache 已有 `_estimate_hicache_size_per_token()`,可作为 bytes 计算参考。 + +结论:scheduler 可以用 `token_to_kv_pool_allocator.get_kvcache()` 推导 dtype、layer_num、kv_cache_dim、index_head_dim、active index layer count。 + +### C5. L1 CP shared-KV prefetch 当前仍 gate 掉 bs>1 + +相关文件: + +- `python/sglang/srt/layers/attention/nsa/cp_shared_kv_prefetch.py:380-397` +- `python/sglang/srt/layers/attention/nsa/cp_shared_kv_prefetch.py:915-917` +- `python/sglang/srt/layers/attention/nsa/cp_shared_kv_prefetch.py:1191-1202` +- `python/sglang/srt/layers/attention/nsa/cp_shared_kv_prefetch.py:1717-1719` + +当前: + +- MLA prefetch `batch_size != 1` 时 skip。 +- index prefetch `batch_size != 1` 时 skip。 +- MLA prefetch 若启用,会持有完整 `dense_kv_cache` handle。 +- index prefetch 若启用,会持有完整 `dense_page_buffer` handle。 + +结论:当前 bs>1 运行时,`max_buffer_size` 不需要计入 L1 shared-KV prefetch dense buffer;但 estimator 必须预留字段,后续打开 bs>1 prefetch 时直接生效。 + +### C6. logits / logprob 峰值与 request 选项强相关 + +相关文件: + +- `python/sglang/srt/layers/logits_processor.py:315-405` +- `python/sglang/srt/layers/logits_processor.py:413-545` +- `python/sglang/srt/layers/logits_processor.py:675-745` +- `python/sglang/srt/layers/logits_processor.py:849-879` + +普通 prefill 不返回 input logprob 时,logits 通常只对 batch last-token rows 计算,峰值近似: + +```text +batch_size * vocab_shard * logits_dtype_size +``` + +返回 input logprob 时,`pruned_states` 可能包含更多 extend rows;如果启用 logits chunk,峰值按 `min(pruned_rows, logprobs_chunk_size)` 估算;如果未启用 chunk,峰值按 `pruned_rows` 估算。 + +结论:buffer gate 至少要 conservative 地覆盖 logits/logprob 峰值,否则 `extend_tokens` 很小但 logprob 请求可能放大 CUDA temp。 + +## 2. 设计原则 + +### 2.1 估算同一时间 live 的峰值,并按 stream 并发叠加 + +不能按 layer 数累加 temp buffer;但也不能把所有类别简单取 `max()`。真正需要防 OOM 的是同一时间 live 的最大值。只要不同 stream 上可能同时存在,就应该在同一个并发窗口里相加。 + +需要区分至少三类 stream / lifetime: + +```text +current/default forward stream: + 当前 layer 的 MLA/index materialize、remap、attention 使用的 dense/slot buffer + +prefetch stream: + next/next+1 layer 的 L2->L1 load、MLA/index prefix prefetch handle + +backup / transfer stream: + async backup descriptor、D2H staging、pending write ack 相关小 buffer + +post-forward/logits window: + logits/logprob/lm-head output buffer,可能与仍未完成的 prefetch/backup buffer 重叠 +``` + +第一版估算不要写成: + +```python +# 错误:会低估不同 stream 同时 live 的 buffer。 +estimated_peak_buffer_bytes = max(A, B, C, D, E) +``` + +应该写成 stream-aware overlap windows: + +```python +layer_forward_peak = ( + materialize_peak_bytes + + remap_peak_bytes + + prefetch_peak_bytes + + backup_descriptor_peak_bytes +) + +logits_window_peak = ( + logits_peak_bytes + + prefetch_peak_bytes + + backup_descriptor_peak_bytes +) + +load_back_window_peak = ( + transfer_descriptor_peak_bytes + + prefetch_peak_bytes + + backup_descriptor_peak_bytes +) + +estimated_peak_buffer_bytes = max( + layer_forward_peak, + logits_window_peak, + load_back_window_peak, +) +``` + +当前 bs>1 L1 shared-KV prefetch 仍 gate 关闭时,`prefetch_peak_bytes=0`。后续打开 bs>1 prefetch 后,这个字段必须反映真实 live prefetch buffer,而不是继续当作独立窗口取 max。 + +### 2.2 保持 resource type 分离 + +内部 estimator 不应该只有一个 token scalar。至少保留这些字段,便于日志和 benchmark 判断到底被什么卡住: + +```python +@dataclass(frozen=True) +class CPSharedKVPrefillBufferEstimate: + total_peak_bytes: int + layer_forward_peak_bytes: int + logits_window_peak_bytes: int + load_back_window_peak_bytes: int + materialize_peak_bytes: int + prefetch_peak_bytes: int + logits_peak_bytes: int + remap_peak_bytes: int + transfer_descriptor_peak_bytes: int + backup_descriptor_peak_bytes: int + l1_load_back_bytes: int + l1_extend_bytes: int + host_backup_bytes: int +``` + +其中 gate 第一阶段只使用 `total_peak_bytes`。其它字段用于 debug、benchmark 和后续校准。 + +### 2.3 第一阶段宁可保守,也不能引入正式路径额外 CUDA 操作 + +estimator 必须是 CPU-only、无 CUDA allocation、无 collective、无 tensor dump。正式推理 hot path 不能为了 admission gate 做额外 GPU 计算。 + +### 2.4 single request 不 deadlock + +与现有 token gates 一致: + +```python +return projected > limit and len(self.can_run_list) > 0 +``` + +单个大请求允许独占 batch。真正 capacity 仍由 allocator / CUDA OOM / fail-fast 路径处理。 + +## 3. 峰值估算公式 + +### 3.1 输入量 + +每个 request 在 `init_load_back()` 后已有: + +```text +prefix_len = len(req.prefix_indices) +extend_len = req.extend_input_len +seq_len = len(req.fill_ids) +``` + +batch 级: + +```text +batch_size = len(can_run_list) + 1 +paged_prefix_tokens = sum(ceil_page(prefix_len_i)) +paged_extend_tokens = sum(ceil_page(extend_len_i)) +prefix_pages = paged_prefix_tokens / page_size +extend_pages = paged_extend_tokens / page_size +``` + +### 3.2 MLA materialize peak + +对 CP shared-KV cache-hit/partial-current 路径,dense MLA buffer 近似: + +```text +dense_mla_pages = prefix_pages + extend_pages +mla_page_bytes = page_size * kv_cache_dim * kv_dtype_size +materialize_mla_bytes = dense_mla_pages * mla_page_bytes +``` + +如果 current-only 且不 materialize prefix,可按 `extend_pages` 估算;第一阶段可以统一用 `prefix_pages + extend_pages` 保守估算。 + +### 3.3 index materialize peak + +NSA index page buffer: + +```text +index_page_bytes = page_size * (index_head_dim + index_head_dim / quant_block_size * 4) * uint8_size +materialize_index_bytes = dense_index_pages * index_page_bytes +``` + +`dense_index_pages` 第一阶段同样可用 `prefix_pages + extend_pages` 保守估算。后续 index skip 生效后,可按 active index layer 或当前 layer 是否需要 index 细化;scheduler admission 第一阶段不按 layer 区分,只估最坏 active layer。 + +### 3.4 remap / page_inverse peak + +当前 slot remap 结构包含: + +- `slot_logical_pages` +- `page_inverse` +- `dense_locs` / `dense_pages` +- sorted slot metadata + +保守估算: + +```text +logical_page_capacity = max_page_id_capacity_or_pages_per_request_capacity +page_inverse_bytes = batch_size * logical_page_capacity * int64_size +slot_metadata_bytes = dense_pages * int64_size * 4 +remap_peak_bytes = page_inverse_bytes + slot_metadata_bytes +``` + +注意:这里必须避免重新引入早期 bs>1 的巨型 per-request dense inverse 设计。实现时如果无法可靠得到 `logical_page_capacity` 的紧上界,先用实际 page-table width 而不是 allocator logical capacity。 + +### 3.5 logits / logprob peak + +普通 generation: + +```text +logits_rows = batch_size +logits_peak_bytes = logits_rows * vocab_shard_size * logits_dtype_size +``` + +返回 input logprob: + +```text +pruned_rows = sum(max(1, extend_len_i - extend_logprob_start_len_i)) +logits_rows = min(pruned_rows, logprobs_chunk_size) if logprob_chunk_enabled else pruned_rows +logits_peak_bytes = logits_rows * vocab_shard_size * logits_dtype_size +``` + +第一阶段如 scheduler 难以拿到 vocab shard / logits dtype,可用 conservative 默认: + +```text +vocab_size from model_config.vocab_size / tp_size +logits_dtype_size = 2 unless enable_fp32_lm_head then 4 +``` + +### 3.6 prefetch peak + +当前 bs>1 L1 shared-KV prefetch disabled: + +```text +prefetch_peak_bytes = 0 +``` + +后续打开 bs>1 prefetch 后: + +```text +prefetch_peak_bytes = prefetch_mla_dense_bytes + prefetch_index_dense_bytes +``` + +因为 prefetch handle 在 next-layer consume 前持有完整 dense buffer,必须按完整 dense pages 估算,而不是 prefix rows。 + +L2->L1 transfer prefetch 不一定有大 GPU staging,但必须计入: + +```text +transfer_descriptor_peak_bytes = descriptor_count * descriptor_entry_size +``` + +如果 direct backend 只使用 CPU descriptor,GPU peak 可为 0,但 benchmark 仍应报告 CPU descriptor count/bytes。 + +### 3.7 backup peak + +per-layer async backup 当前主要是 D2H transfer,GPU memory 大头通常不是 staging,而是 descriptor。第一阶段估算: + +```text +backup_descriptor_peak_bytes = backup_page_count * descriptor_entry_size +host_backup_bytes = page_count * host_page_bytes +``` + +host reservation/evict 仍由 HiCache host allocator 管,`max_buffer_size` 不负责 host correctness。 + +## 4. 实现计划 + +### P0. 增加参数、GB 默认解析与校验 + +文件: + +- `python/sglang/srt/server_args.py` +- `test/registered/unit/server_args/test_server_args.py` + +步骤: + +1. 在 `ServerArgs` dataclass 增加内部 bytes 字段: + +```python +cp_shared_kv_prefill_max_buffer_size: Optional[int] = None +``` + +2. 新增 parser helper,裸数字按 decimal GB 解析,显式后缀复用现有 `human_readable_int` 语义: + +```python +def human_readable_gb_size(value: str) -> int: + value = value.strip() + if re.fullmatch(r"\d+(?:\.\d+)?", value): + return int(Decimal(value) * Decimal(10**9)) + return human_readable_int(value) +``` + +这样启动参数含义为: + +```text +--cp-shared-kv-prefill-max-buffer-size 8 -> 8_000_000_000 bytes +--cp-shared-kv-prefill-max-buffer-size 8G -> 8_000_000_000 bytes +--cp-shared-kv-prefill-max-buffer-size 8Gi -> 8_589_934_592 bytes +``` + +3. 在 `_post_init` 校验: + +```python +if self.cp_shared_kv_prefill_max_buffer_size is not None and self.cp_shared_kv_prefill_max_buffer_size <= 0: + raise ValueError("cp_shared_kv_prefill_max_buffer_size must be a positive integer when specified.") +``` + +4. 在 CLI 增加: + +```bash +--cp-shared-kv-prefill-max-buffer-size +``` + +help 必须明确:裸数字单位为 GB,显式后缀支持 SI/IEC。 + +5. 单测: + +- CLI 能解析裸数字 `8` 为 `8_000_000_000`。 +- CLI 能解析 `8Gi` 为 `8 * 2**30`。 +- `0` 抛正数校验错误。 +- 未设置时为 `None`。 + +### P1. 定义 CPU-only buffer estimate 数据结构 + +文件: + +- 推荐新增:`python/sglang/srt/managers/cp_shared_kv_prefill_buffer_estimator.py` +- 测试:`test/registered/unit/managers/test_cp_shared_kv_prefill_buffer_estimator.py` + +新增: + +```python +@dataclass(frozen=True) +class CPSharedKVPrefillBufferEstimate: + total_peak_bytes: int + layer_forward_peak_bytes: int + logits_window_peak_bytes: int + load_back_window_peak_bytes: int + materialize_peak_bytes: int + prefetch_peak_bytes: int + logits_peak_bytes: int + remap_peak_bytes: int + transfer_descriptor_peak_bytes: int + backup_descriptor_peak_bytes: int + l1_load_back_bytes: int + l1_extend_bytes: int + host_backup_bytes: int +``` + +新增 helper: + +```python +def estimate_cp_shared_kv_prefill_buffer_bytes( + *, + page_size: int, + batch_size: int, + prefix_lens: Sequence[int], + extend_lens: Sequence[int], + kvcache: object | None, + model_config: object | None, + tp_size: int, + return_logprob_rows: int = 0, + logprob_chunk_enabled: bool = False, + logprob_chunk_size: int = 2048, + bs_gt1_l1_prefetch_enabled: bool = False, +) -> CPSharedKVPrefillBufferEstimate: +``` + +第一阶段要求: + +- 不 import CUDA-only package。 +- 不分配 CUDA tensor。 +- `kvcache is None` 时使用保守 fallback 并返回非零 remap/logits 估算。 +- fp8/bf16 根据 `kvcache.store_dtype.itemsize` 计算 MLA KV bytes。 +- NSA index bytes 根据 `index_head_dim` / `quant_block_size` / `index_k_with_scale_buffer_dtype.itemsize` 计算。 +- `total_peak_bytes` 必须按 2.1 的 stream-aware overlap windows 计算,而不是对各字段简单取 max。 + +### P2. 在 `PrefillAdder` 中接入 projected peak gate + +文件: + +- `python/sglang/srt/managers/schedule_policy.py` +- `python/sglang/srt/managers/scheduler.py` +- `test/registered/unit/managers/test_prefill_adder.py` + +新增 constructor 参数: + +```python +cp_shared_kv_prefill_max_buffer_size: Optional[int] = None +cp_shared_kv_prefill_buffer_estimator_context: Optional[CPSharedKVPrefillBufferEstimatorContext] = None +``` + +`PrefillAdder` 内维护: + +```python +self.cp_shared_kv_prefill_estimate_prefix_lens: list[int] = [] +self.cp_shared_kv_prefill_estimate_extend_lens: list[int] = [] +self.cp_shared_kv_prefill_estimated_peak_buffer_bytes: int = 0 +``` + +在 `add_one_req()` 中,`init_load_back()` 后计算 candidate: + +```text +candidate_prefix_lens = accepted_prefix_lens + [prefix_len] +candidate_extend_lens = accepted_extend_lens + [input_tokens] +estimate = estimate_cp_shared_kv_prefill_buffer_bytes(...) +if estimate.total_peak_bytes > max_buffer_size and can_run_list 非空: + return AddReqResult.OTHER +``` + +接受 request 后,在 `_update_prefill_budget()` 附近同步更新 estimator 状态。 + +单测: + +1. 两个小 extend request 在 token gates 未超限时,因为 projected buffer 超限而第二个返回 `OTHER`。 +2. 单个 request 超过 buffer limit 仍 `CONTINUE`。 +3. buffer limit 与 extend/cached token limit 同时存在时,任一超限都会停止 batch。 +4. `enable_cp_shared_kv_prefill_bs_gt1=False` 时 buffer gate 不生效。 + +### P3. 从 scheduler 传入 estimator context + +文件: + +- `python/sglang/srt/managers/scheduler.py` + +context 应包含: + +```python +@dataclass(frozen=True) +class CPSharedKVPrefillBufferEstimatorContext: + kvcache: object | None + model_config: object | None + tp_size: int + page_size: int + logprob_chunk_enabled: bool + logprob_chunk_size: int + bs_gt1_l1_prefetch_enabled: bool +``` + +第一阶段 `bs_gt1_l1_prefetch_enabled=False`,因为当前 prefetcher 明确拒绝 bs>1。 + +`kvcache` 来源: + +```python +self.token_to_kv_pool_allocator.get_kvcache() +``` + +`model_config` 来源: + +```python +self.model_config +``` + +`tp_size` 来源: + +```python +self.tp_size +``` + +logprob chunk env 可从 `sglang.srt.environ.envs` 读取,或在 context 中使用与 `LogitsProcessor` 一致的默认值。 + +### P3.5. 启动后 CUDA buffer smoke allocation + +文件: + +- `python/sglang/srt/managers/scheduler.py` +- 可选新增 helper:`python/sglang/srt/managers/cp_shared_kv_prefill_buffer_estimator.py` + +目的:启动完成、主要静态显存对象已经分配后,按 `cp_shared_kv_prefill_max_buffer_size` 尝试临时分配一次,提前发现配置值超过剩余显存。 + +推荐接入点:`Scheduler.__init__()` 末尾、`self.is_initializing = False` 前。此时模型、KV pool、attention backend、CUDA graph、disaggregation、overlap 等主要对象已经初始化,且 scheduler 尚未 ready。 + +触发条件: + +```python +should_smoke_check = ( + self.server_args.enable_cp_shared_kv_prefill_bs_gt1 + and self.server_args.enable_nsa_prefill_cp_shared_kv + and self.server_args.cp_shared_kv_prefill_max_buffer_size is not None + and self.server_args.disaggregation_mode in (None, "prefill") +) +``` + +实现要求: + +```python +def smoke_check_cp_shared_kv_prefill_buffer_size(device: torch.device | str, size_bytes: int) -> None: + try: + torch.cuda.synchronize() + probe = torch.empty(size_bytes, dtype=torch.uint8, device=device) + # 只触碰首尾,避免 fill 整个大 buffer 带来长启动耗时;cudaMalloc 本身已能暴露大部分 OOM。 + if size_bytes > 0: + probe[0] = 0 + probe[-1] = 0 + torch.cuda.synchronize() + except torch.cuda.OutOfMemoryError as exc: + raise RuntimeError( + "[CP_SHARED_KV_FAIL_FAST][prefill_buffer_smoke] " + f"cannot allocate cp_shared_kv_prefill_max_buffer_size={size_bytes} bytes " + "after scheduler startup. Lower --cp-shared-kv-prefill-max-buffer-size " + "or reduce other static GPU memory consumers." + ) from exc + finally: + try: + del probe + except UnboundLocalError: + pass + torch.cuda.empty_cache() +``` + +注意点: + +1. 这是启动期 fail-fast,不在 scheduler hot path 中运行。 +2. 每个 CP/TP rank 都执行一次,验证 per-rank 剩余显存。 +3. 不长期持有 probe buffer,避免实际推理容量下降。 +4. 日志需要打印 rank、size GB、smoke 通过/失败。 +5. 该 smoke check 只能证明启动后能分配一个同等大小的连续 CUDA block,不能替代 runtime estimator;runtime 中仍可能因为多个 stream live buffer 同时存在而 OOM,所以 admission gate 仍必须按 2.1 叠加并发窗口。 + +### P4. 更新 benchmark,能离线观察 buffer gate + +文件: + +- `benchmark/hicache/bench_prefill_scheduler_admission.py` +- `test/registered/unit/managers/test_prefill_scheduler_admission_bench.py` +- `docs/advanced_features/nsa_prefill_cp_scheduler_admission_benchmark.md` + +新增 benchmark 参数: + +```bash +--cp-max-buffer-size +--kv-cache-dim +--kv-dtype-bytes +--layer-num +--index-head-dim +--vocab-size +--tp-size +--logprob-rows-per-extend-token +--enable-bs-gt1-prefetch-estimate +``` + +输出新增: + +```json +"cp_estimated_peak_buffer_bytes": ..., +"cp_buffer_breakdown": { + "layer_forward_peak_bytes": ..., + "logits_window_peak_bytes": ..., + "load_back_window_peak_bytes": ..., + "materialize_peak_bytes": ..., + "prefetch_peak_bytes": ..., + "logits_peak_bytes": ..., + "remap_peak_bytes": ... +} +``` + +单测: + +1. `--cp-max-buffer-size 1` 按 1GB 解析,并能让 benchmark 第二个 request 被挡住。 +2. text/json 输出包含 buffer bytes 与 breakdown。 +3. token gates 仍保持原测试行为。 + +### P5. 加入限频 debug 日志 + +文件: + +- `python/sglang/srt/managers/schedule_policy.py` + +复用现有 CP bs>1 timing/debug 环境变量,不新增新 env。仅在以下情况日志: + +1. buffer gate 阻止加入 request; +2. estimate 超过 slow 阈值或 debug 开启; +3. estimator fallback 到保守模型。 + +日志格式建议: + +```text +[CP_SHARED_KV_BS_GT1_ADMISSION] stop reason=max_buffer_size projected=... limit=... batch_size=... rid=... breakdown=... +``` + +避免每个 request 刷屏;默认限频。 + +### P6. 远端验证顺序 + +1. 本地 CPU 单测: + +```bash +PYTHONPATH=python python -m pytest -q \ + test/registered/unit/server_args/test_server_args.py \ + test/registered/unit/managers/test_prefill_adder.py \ + test/registered/unit/managers/test_prefill_scheduler_admission_bench.py \ + test/registered/unit/managers/test_cp_shared_kv_prefill_buffer_estimator.py +``` + +2. 远端容器 CPU 单测: + +```bash +cd /sgl-workspace/sglang +PYTHONPATH=python python -m pytest -q \ + test/registered/unit/server_args/test_server_args.py \ + test/registered/unit/managers/test_prefill_adder.py \ + test/registered/unit/managers/test_prefill_scheduler_admission_bench.py \ + test/registered/unit/managers/test_cp_shared_kv_prefill_buffer_estimator.py +``` + +3. 远端短 ETE: + +- 使用原启动参数; +- 追加 `--cp-shared-kv-prefill-max-buffer-size 8` 起步;裸数字表示 8GB。若要二进制 GiB,显式写 `8Gi`。 +- 保留: + - `--cp-shared-kv-prefill-max-total-extend-tokens 65536` + - `--cp-shared-kv-prefill-max-total-cached-tokens` 如已配置继续保留; +- 用 50/200 条 GSM8K 或 replay 小样本确认 correctness 不掉点; +- 再跑 replay 观察 batch size、吞吐、OOM/卡死情况。 + +## 5. 风险与边界 + +### R1. 估算不应替代 allocator correctness + +L1 device KV pages、owner-lane free room、HiCache host reservation 仍由现有 allocator / radix / HiCache 路径保证。buffer gate 只减少 batch 过大导致的临时 CUDA OOM 和 CPU descriptor pressure。 + +### R2. bs>1 prefetch 未来打开时需要更新 context 和并发窗口 + +当前代码明确 gate 掉 bs>1 prefetch,因此第一版 `prefetch_peak_bytes=0` 是符合当前运行时的。后续打开 bs>1 prefetch 必须把 `bs_gt1_l1_prefetch_enabled` 改成真实状态,并将 prefetch MLA/index dense handle 加到可能重叠的 layer/logits/load-back windows 中,不能把它作为独立 max window。 + +### R3. logits 估算只能 conservative + +scheduler admission 阶段拿不到所有 logits processor runtime 状态。第一版按 model_config、request logprob flags、env chunk size 做 conservative estimate。若发现过度保守,再用实际 runtime telemetry 校准。 + +### R4. 不要引入 CUDA 操作到 admission + +所有 estimator 逻辑必须 CPU-only。不能在 scheduler scan waiting queue 时创建 CUDA tensor、调用 collective 或触发 TAI kernel。 + +### R5. 参数默认不改变现有行为 + +`cp_shared_kv_prefill_max_buffer_size=None` 时完全保持当前行为。只有用户显式设置时才新增 gate。 + +## 6. 推荐初始参数 + +对当前 GLM5 FP8 + CP shared-KV + bs>1,建议从保守值开始: + +```bash +--cp-shared-kv-prefill-max-buffer-size 8 +``` + +裸数字单位是 GB。若需要二进制 GiB,可显式写: + +```bash +--cp-shared-kv-prefill-max-buffer-size 8Gi +``` + +如果 replay 中 batch size 被频繁 buffer gate 卡住且无 OOM,可逐步调到: + +```bash +--cp-shared-kv-prefill-max-buffer-size 12 +--cp-shared-kv-prefill-max-buffer-size 16 +``` + +不要直接去掉 extend/cached token gates。三个 gate 分别约束不同风险: + +```text +max_total_extend_tokens -> compute / logits / fresh KV growth +max_total_cached_tokens -> cache-hit materialize / load-back / descriptor pressure +max_buffer_size -> dtype/model/config-aware CUDA temp 峰值 +``` + +## 7. 当前实现状态(2026-06-11) + +已完成: + +1. `--cp-shared-kv-prefill-max-buffer-size` 已接入 `ServerArgs`,内部保存 bytes;裸数字按 decimal GB 解析,显式 `G/Gi/Mi` 后缀复用现有 size parser。 +2. 新增 CPU-only estimator:`python/sglang/srt/managers/cp_shared_kv_prefill_buffer_estimator.py`。 + - 按 page 粒度估算 prefix/extend materialize。 + - 同时估算 MLA KV、NSA index、remap、logits、load-back/backup descriptor。 + - `total_peak_bytes` 使用 stream-aware window:layer-forward / logits / load-back 三个并发窗口取最大;窗口内会叠加可重叠的 prefetch/backup descriptor。 +3. `PrefillAdder` 已接入 projected buffer gate。 + - 仍保留 request-count / total-extend / total-cached 三个 gate。 + - `max_buffer_size=None` 时不维护 prefix/extend estimator list,避免默认 hot path 额外 CPU 开销。 + - 单个超限 request 仍允许独占 batch,避免 scheduler deadlock。 + - buffer gate 命中时复用 `SGLANG_CP_SHARED_KV_BS_GT1_DEBUG` 做限频 debug 日志,不新增 env。 +4. `Scheduler` 已接入 estimator context。 + - 只有 `cp_shared_kv_prefill_max_buffer_size` 设置时才构建 context。 + - 当前 bs>1 L1 prefetch 仍关闭,因此 context 中 `bs_gt1_l1_prefetch_enabled=False`。 +5. 启动期 CUDA smoke allocation 已接入。 + - 触发条件:CP shared-KV + bs>1 + max-buffer-size 设置 + prefill/null disagg mode。 + - 每个 rank 在 scheduler ready 前临时分配一次同等大小 `uint8` CUDA tensor,触碰首尾后释放,用于提前暴露配置过大。 +6. `benchmark/hicache/bench_prefill_scheduler_admission.py` 已增加 buffer limit、KV/index/logits 参数和 breakdown 输出。 + +验证: + +```bash +# 本地 py_compile +python -m py_compile \ + python/sglang/srt/server_args.py \ + python/sglang/srt/managers/cp_shared_kv_prefill_buffer_estimator.py \ + python/sglang/srt/managers/schedule_policy.py \ + python/sglang/srt/managers/scheduler.py \ + benchmark/hicache/bench_prefill_scheduler_admission.py \ + test/registered/unit/server_args/test_server_args.py \ + test/registered/unit/managers/test_cp_shared_kv_prefill_buffer_estimator.py \ + test/registered/unit/managers/test_prefill_adder.py \ + test/registered/unit/managers/test_prefill_scheduler_admission_bench.py + +# 本地可运行的 estimator 单测 +PYTHONPATH=python python -m pytest -q \ + test/registered/unit/managers/test_cp_shared_kv_prefill_buffer_estimator.py +# 结果:4 passed + +# 远端 cjy-glm5-new /sgl-workspace/sglang-tai targeted 单测 +PYTHONPATH=python python -m pytest -q \ + test/registered/unit/managers/test_cp_shared_kv_prefill_buffer_estimator.py \ + test/registered/unit/server_args/test_server_args.py::test_cp_shared_kv_prefill_max_buffer_size_defaults_to_gb \ + test/registered/unit/server_args/test_server_args.py::test_cp_shared_kv_prefill_max_buffer_size_accepts_iec_suffix \ + test/registered/unit/server_args/test_server_args.py::test_cp_shared_kv_prefill_max_buffer_size_rejects_non_positive_value \ + test/registered/unit/managers/test_prefill_adder.py::TestPrefillAdder::test_cp_prefill_buffer_limit_stops_second_request_without_token_gate \ + test/registered/unit/managers/test_prefill_adder.py::TestPrefillAdder::test_cp_prefill_buffer_limit_allows_single_oversized_request \ + test/registered/unit/managers/test_prefill_scheduler_admission_bench.py::TestPrefillSchedulerAdmissionBench::test_max_buffer_size_is_observable_in_scheduler_trace +# 结果:10 passed + +# 远端 managers 相关全文件 +PYTHONPATH=python python -m pytest -q \ + test/registered/unit/managers/test_cp_shared_kv_prefill_buffer_estimator.py \ + test/registered/unit/managers/test_prefill_adder.py \ + test/registered/unit/managers/test_prefill_scheduler_admission_bench.py +# 结果:29 passed +``` + +未完全验证: + +- `test/registered/unit/server_args/test_server_args.py` 全文件在远端有 1 个既有网络依赖失败:`TestPrepareServerArgs.test_prepare_server_args` 需要从 HuggingFace 拉 `Qwen/Qwen2.5-1.5B-Instruct/config.json`,当前容器 DNS/HF 网络失败;本次新增的 3 个 server_args targeted tests 已通过。 +- 尚未跑真实 ETE 启动验证 `--cp-shared-kv-prefill-max-buffer-size` smoke allocation 对当前 GLM5 配置的实际可用值。 + +## 8. Chunked prefill 下的 effective extend limit(2026-06-11) + +新增修正:当 CP shared-KV bs>1 与 chunked prefill 同时开启时,`PrefillAdder` 中实际使用的 extend grouping limit 取: + +```python +effective_extend_limit = min( + cp_shared_kv_prefill_max_total_extend_tokens, + current_rem_chunk_tokens, +) +``` + +原因:chunked prefill 已经限制一个 batch 当前最多消耗的 extend token。如果 CP 专用 extend limit 大于当前 chunk budget,较大的值不可达;继续用 raw CP limit 还会把 generic `max_prefill_tokens` lift 到过大的数值,造成 admission debug/benchmark 表达的 batch 容量与真实 chunk 容量不一致。 + +实现位置: + +- `python/sglang/srt/managers/schedule_policy.py` 的 `PrefillAdder.__init__()`。 + +验证: + +```bash +# 先确认 red test:旧实现下 256 != 128,失败。 +PYTHONPATH=python python -m pytest -q \ + test/registered/unit/managers/test_prefill_adder.py::TestPrefillAdder::test_cp_prefill_total_extend_limit_is_capped_by_chunked_prefill_size + +# 修复后远端 cjy-glm5-new: +PYTHONPATH=python python -m pytest -q \ + test/registered/unit/managers/test_prefill_adder.py::TestPrefillAdder::test_cp_prefill_total_extend_limit_is_capped_by_chunked_prefill_size +# 结果:1 passed + +PYTHONPATH=python python -m pytest -q \ + test/registered/unit/managers/test_prefill_adder.py +# 结果:21 passed +``` diff --git a/python/sglang/srt/managers/cp_shared_kv_prefill_buffer_estimator.py b/python/sglang/srt/managers/cp_shared_kv_prefill_buffer_estimator.py new file mode 100644 index 000000000..5a3f92427 --- /dev/null +++ b/python/sglang/srt/managers/cp_shared_kv_prefill_buffer_estimator.py @@ -0,0 +1,249 @@ +"""CPU-only admission estimates for CP shared-KV prefill batching. + +The scheduler uses this module before it builds a real batch. Keep this file +free of CUDA allocations and collectives except for the explicit startup smoke +check helper. +""" + +from __future__ import annotations + +import logging +from dataclasses import dataclass +from typing import Optional, Sequence + +import torch + +logger = logging.getLogger(__name__) + + +_DEFAULT_DESCRIPTOR_BYTES = 64 +_DEFAULT_LOGITS_DTYPE_BYTES = 2 +_DEFAULT_FALLBACK_KV_CACHE_DIM = 1 +_DEFAULT_FALLBACK_INDEX_HEAD_DIM = 128 +_DEFAULT_FALLBACK_QUANT_BLOCK_SIZE = 128 + + +@dataclass(frozen=True) +class CPSharedKVPrefillBufferEstimatorContext: + kvcache: object | None + model_config: object | None + tp_size: int + page_size: int + logprob_chunk_enabled: bool + logprob_chunk_size: int + bs_gt1_l1_prefetch_enabled: bool = False + + +@dataclass(frozen=True) +class CPSharedKVPrefillBufferEstimate: + total_peak_bytes: int + layer_forward_peak_bytes: int + logits_window_peak_bytes: int + load_back_window_peak_bytes: int + materialize_peak_bytes: int + prefetch_peak_bytes: int + logits_peak_bytes: int + remap_peak_bytes: int + transfer_descriptor_peak_bytes: int + backup_descriptor_peak_bytes: int + l1_load_back_bytes: int + l1_extend_bytes: int + host_backup_bytes: int + + +def ceil_paged_tokens(tokens: int, page_size: int) -> int: + if tokens <= 0: + return 0 + return -(-int(tokens) // page_size) * page_size + + +def ceil_pages(tokens: int, page_size: int) -> int: + if tokens <= 0: + return 0 + return -(-int(tokens) // page_size) + + +def _dtype_size(dtype: object | None, default: int = 2) -> int: + if dtype is None: + return default + itemsize = getattr(dtype, "itemsize", None) + if itemsize is not None: + return int(itemsize) + try: + return int(torch.empty((), dtype=dtype, device="cpu").element_size()) + except Exception: + return default + + +def _kv_cache_dim(kvcache: object | None) -> int: + return int( + getattr(kvcache, "kv_cache_dim", None) + or getattr(kvcache, "element_dim", None) + or _DEFAULT_FALLBACK_KV_CACHE_DIM + ) + + +def _kv_dtype_bytes(kvcache: object | None) -> int: + return _dtype_size(getattr(kvcache, "store_dtype", None), default=2) + + +def _index_page_bytes(kvcache: object | None, page_size: int) -> int: + index_head_dim = int( + getattr(kvcache, "index_head_dim", None) or _DEFAULT_FALLBACK_INDEX_HEAD_DIM + ) + quant_block_size = int( + getattr(kvcache, "quant_block_size", None) + or _DEFAULT_FALLBACK_QUANT_BLOCK_SIZE + ) + index_dtype_bytes = _dtype_size( + getattr(kvcache, "index_k_with_scale_buffer_dtype", None), + default=1, + ) + scale_bytes_per_token = (index_head_dim // quant_block_size) * 4 + return page_size * (index_head_dim + scale_bytes_per_token) * index_dtype_bytes + + +def _vocab_shard_size(model_config: object | None, tp_size: int) -> int: + vocab_size = int(getattr(model_config, "vocab_size", 0) or 0) + if vocab_size <= 0: + return 0 + return -(-vocab_size // max(int(tp_size), 1)) + + +def estimate_cp_shared_kv_prefill_buffer_bytes( + *, + page_size: int, + batch_size: int, + prefix_lens: Sequence[int], + extend_lens: Sequence[int], + context: CPSharedKVPrefillBufferEstimatorContext | None, + return_logprob_rows: int = 0, + logits_dtype_bytes: int = _DEFAULT_LOGITS_DTYPE_BYTES, + descriptor_bytes: int = _DEFAULT_DESCRIPTOR_BYTES, +) -> CPSharedKVPrefillBufferEstimate: + """Estimate peak temporary GPU buffer pressure for one candidate batch. + + The estimate is intentionally conservative and stream-aware. It sums + buffers that can be live concurrently on different streams instead of + taking a max over independent resource categories. + """ + + if context is None: + context = CPSharedKVPrefillBufferEstimatorContext( + kvcache=None, + model_config=None, + tp_size=1, + page_size=page_size, + logprob_chunk_enabled=False, + logprob_chunk_size=2048, + bs_gt1_l1_prefetch_enabled=False, + ) + + if page_size <= 0: + raise ValueError(f"page_size must be positive, got {page_size}") + if len(prefix_lens) != len(extend_lens): + raise ValueError( + "prefix_lens and extend_lens must have the same length: " + f"{len(prefix_lens)} != {len(extend_lens)}" + ) + + prefix_pages = [ceil_pages(tokens, page_size) for tokens in prefix_lens] + extend_pages = [ceil_pages(tokens, page_size) for tokens in extend_lens] + total_prefix_pages = sum(prefix_pages) + total_extend_pages = sum(extend_pages) + total_materialize_pages = sum(p + e for p, e in zip(prefix_pages, extend_pages)) + + kv_page_bytes = page_size * _kv_cache_dim(context.kvcache) * _kv_dtype_bytes( + context.kvcache + ) + index_page_bytes = _index_page_bytes(context.kvcache, page_size) + + materialize_peak_bytes = total_materialize_pages * ( + kv_page_bytes + index_page_bytes + ) + remap_peak_bytes = max(total_materialize_pages, int(batch_size)) * 8 + prefetch_peak_bytes = ( + total_prefix_pages * (kv_page_bytes + index_page_bytes) + if context.bs_gt1_l1_prefetch_enabled + else 0 + ) + + logits_rows = max(int(batch_size), 0) + if return_logprob_rows > 0: + logits_rows = int(return_logprob_rows) + if context.logprob_chunk_enabled: + logits_rows = min(logits_rows, int(context.logprob_chunk_size)) + vocab_shard = _vocab_shard_size(context.model_config, context.tp_size) + logits_peak_bytes = logits_rows * vocab_shard * int(logits_dtype_bytes) + + transfer_descriptor_peak_bytes = total_prefix_pages * int(descriptor_bytes) + backup_descriptor_peak_bytes = total_extend_pages * int(descriptor_bytes) + + l1_load_back_bytes = total_prefix_pages * kv_page_bytes + l1_extend_bytes = total_extend_pages * kv_page_bytes + host_backup_bytes = total_extend_pages * kv_page_bytes + + layer_forward_peak_bytes = ( + materialize_peak_bytes + + remap_peak_bytes + + prefetch_peak_bytes + + backup_descriptor_peak_bytes + ) + logits_window_peak_bytes = ( + logits_peak_bytes + prefetch_peak_bytes + backup_descriptor_peak_bytes + ) + load_back_window_peak_bytes = ( + transfer_descriptor_peak_bytes + + prefetch_peak_bytes + + backup_descriptor_peak_bytes + ) + total_peak_bytes = max( + layer_forward_peak_bytes, + logits_window_peak_bytes, + load_back_window_peak_bytes, + ) + + return CPSharedKVPrefillBufferEstimate( + total_peak_bytes=total_peak_bytes, + layer_forward_peak_bytes=layer_forward_peak_bytes, + logits_window_peak_bytes=logits_window_peak_bytes, + load_back_window_peak_bytes=load_back_window_peak_bytes, + materialize_peak_bytes=materialize_peak_bytes, + prefetch_peak_bytes=prefetch_peak_bytes, + logits_peak_bytes=logits_peak_bytes, + remap_peak_bytes=remap_peak_bytes, + transfer_descriptor_peak_bytes=transfer_descriptor_peak_bytes, + backup_descriptor_peak_bytes=backup_descriptor_peak_bytes, + l1_load_back_bytes=l1_load_back_bytes, + l1_extend_bytes=l1_extend_bytes, + host_backup_bytes=host_backup_bytes, + ) + + +def smoke_check_cp_shared_kv_prefill_buffer_size( + *, + device: torch.device | str, + size_bytes: int, + torch_module=torch, +) -> None: + """Fail fast if startup cannot allocate the configured temp buffer size.""" + + probe = None + try: + torch_module.cuda.synchronize() + probe = torch_module.empty(size_bytes, dtype=torch_module.uint8, device=device) + if size_bytes > 0: + probe[0] = 0 + probe[-1] = 0 + torch_module.cuda.synchronize() + except torch_module.cuda.OutOfMemoryError as exc: + raise RuntimeError( + "[CP_SHARED_KV_FAIL_FAST][prefill_buffer_smoke] " + f"cannot allocate cp_shared_kv_prefill_max_buffer_size={size_bytes} " + "bytes after scheduler startup. Lower " + "--cp-shared-kv-prefill-max-buffer-size or reduce other static GPU " + "memory consumers." + ) from exc + finally: + del probe + torch_module.cuda.empty_cache() diff --git a/python/sglang/srt/managers/schedule_policy.py b/python/sglang/srt/managers/schedule_policy.py index 661ef5aba..99a696b92 100644 --- a/python/sglang/srt/managers/schedule_policy.py +++ b/python/sglang/srt/managers/schedule_policy.py @@ -34,8 +34,14 @@ from typing import TYPE_CHECKING, Dict, List, Optional, Set, Union import torch from sglang.srt.dllm.config import DllmConfig +from sglang.srt.environ import envs from sglang.srt.layers.attention.nsa.utils import is_nsa_prefill_cp_in_seq_split from sglang.srt.layers.utils.cp_utils import is_prefill_context_parallel_enabled +from sglang.srt.managers.cp_shared_kv_prefill_buffer_estimator import ( + CPSharedKVPrefillBufferEstimate, + CPSharedKVPrefillBufferEstimatorContext, + estimate_cp_shared_kv_prefill_buffer_bytes, +) from sglang.srt.managers.schedule_batch import Req, ScheduleBatch from sglang.srt.mem_cache.base_prefix_cache import ( BasePrefixCache, @@ -77,6 +83,24 @@ IN_BATCH_PREFIX_CACHING_DEPRIORITIZE_THRESHOLD = int( IGNORE_EOS_RESERVE_TOKENS = 1 +_CP_SHARED_KV_BS_GT1_ADMISSION_DEBUG_COUNTS: Dict[str, int] = {} + + +def _cp_shared_kv_bs_gt1_admission_debug( + key: str, + message: str, + *args, +) -> None: + if not envs.SGLANG_CP_SHARED_KV_BS_GT1_DEBUG.get(): + return + limit = int(envs.SGLANG_CP_SHARED_KV_BS_GT1_DEBUG_LIMIT.get()) + count = _CP_SHARED_KV_BS_GT1_ADMISSION_DEBUG_COUNTS.get(key, 0) + if limit > 0 and count >= limit: + return + _CP_SHARED_KV_BS_GT1_ADMISSION_DEBUG_COUNTS[key] = count + 1 + logger.info("[CP_SHARED_KV_BS_GT1_DEBUG] event=%s " + message, key, *args) + + class CacheAwarePolicy(Enum): """Scheduling policies that are aware of the tree cache.""" @@ -394,6 +418,10 @@ class PrefillAdder: cp_shared_kv_prefill_max_batch_requests: Optional[int] = None, cp_shared_kv_prefill_max_total_extend_tokens: Optional[int] = None, cp_shared_kv_prefill_max_total_cached_tokens: Optional[int] = None, + cp_shared_kv_prefill_max_buffer_size: Optional[int] = None, + cp_shared_kv_prefill_buffer_estimator_context: Optional[ + CPSharedKVPrefillBufferEstimatorContext + ] = None, prefill_delayer_single_pass: Optional[PrefillDelayerSinglePassExecutor] = None, dllm_config: Optional[DllmConfig] = None, ): @@ -452,9 +480,30 @@ class PrefillAdder: self.cp_shared_kv_prefill_max_total_extend_tokens = ( cp_shared_kv_prefill_max_total_extend_tokens ) + if ( + self._is_cp_prefill_context() + and self.enable_cp_shared_kv_prefill_bs_gt1 + and self.cp_shared_kv_prefill_max_total_extend_tokens is not None + and self.rem_chunk_tokens is not None + ): + # When chunked prefill is active, a batch cannot consume more + # extend tokens than the current chunk budget. Use the smaller + # value as the effective CP-specific grouping limit so the generic + # max_prefill_tokens lift below does not advertise unreachable + # capacity. + self.cp_shared_kv_prefill_max_total_extend_tokens = min( + self.cp_shared_kv_prefill_max_total_extend_tokens, + max(self.rem_chunk_tokens, 0), + ) self.cp_shared_kv_prefill_max_total_cached_tokens = ( cp_shared_kv_prefill_max_total_cached_tokens ) + self.cp_shared_kv_prefill_max_buffer_size = ( + cp_shared_kv_prefill_max_buffer_size + ) + self.cp_shared_kv_prefill_buffer_estimator_context = ( + cp_shared_kv_prefill_buffer_estimator_context + ) if ( self._is_cp_prefill_context() and self.enable_cp_shared_kv_prefill_bs_gt1 @@ -479,6 +528,12 @@ class PrefillAdder: ) self.cp_shared_kv_prefill_total_extend_tokens = 0 self.cp_shared_kv_prefill_total_cached_tokens = 0 + self.cp_shared_kv_prefill_estimate_prefix_lens: list[int] = [] + self.cp_shared_kv_prefill_estimate_extend_lens: list[int] = [] + self.cp_shared_kv_prefill_estimated_peak_buffer_bytes = 0 + self.cp_shared_kv_prefill_last_buffer_estimate: Optional[ + CPSharedKVPrefillBufferEstimate + ] = None def _init_dllm_meta(self, dllm_config: DllmConfig): self.dllm_block_size = dllm_config.block_size @@ -599,6 +654,68 @@ class PrefillAdder: # grouping pressure from prefix/load-back work, not legal request size. return projected > limit and len(self.can_run_list) > 0 + def _estimate_cp_prefill_projected_buffer( + self, prefix_len: int, extend_input_len: int + ) -> CPSharedKVPrefillBufferEstimate: + return estimate_cp_shared_kv_prefill_buffer_bytes( + page_size=self.page_size, + batch_size=len(self.cp_shared_kv_prefill_estimate_prefix_lens) + 1, + prefix_lens=[ + *self.cp_shared_kv_prefill_estimate_prefix_lens, + prefix_len, + ], + extend_lens=[ + *self.cp_shared_kv_prefill_estimate_extend_lens, + extend_input_len, + ], + context=self.cp_shared_kv_prefill_buffer_estimator_context, + ) + + def _cp_prefill_buffer_limit_exceeded( + self, prefix_len: int, extend_input_len: int, rid: Optional[str] = None + ) -> 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_buffer_size + if limit is None: + return False + + estimate = self._estimate_cp_prefill_projected_buffer( + prefix_len, extend_input_len + ) + self.cp_shared_kv_prefill_last_buffer_estimate = estimate + + # Do not deadlock a single large request. The limit bounds grouping, not + # the maximum legal request size; actual runtime capacity remains + # allocator/CUDA-owned. + should_stop = estimate.total_peak_bytes > limit and len(self.can_run_list) > 0 + if should_stop: + _cp_shared_kv_bs_gt1_admission_debug( + "admission_max_buffer_size", + "stop rid=%s projected=%s limit=%s batch_size=%s " + "layer_forward=%s logits_window=%s load_back_window=%s " + "materialize=%s prefetch=%s logits=%s remap=%s " + "transfer_desc=%s backup_desc=%s", + rid, + estimate.total_peak_bytes, + limit, + len(self.can_run_list) + 1, + estimate.layer_forward_peak_bytes, + estimate.logits_window_peak_bytes, + estimate.load_back_window_peak_bytes, + estimate.materialize_peak_bytes, + estimate.prefetch_peak_bytes, + estimate.logits_peak_bytes, + estimate.remap_peak_bytes, + estimate.transfer_descriptor_peak_bytes, + estimate.backup_descriptor_peak_bytes, + ) + return should_stop + def _get_available_device_tokens_for_load_back(self) -> int: if self.is_hybrid_swa: return ( @@ -648,6 +765,20 @@ class PrefillAdder: 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.cp_shared_kv_prefill_max_buffer_size is not None: + self.cp_shared_kv_prefill_estimate_prefix_lens.append(prefix_len) + self.cp_shared_kv_prefill_estimate_extend_lens.append(extend_input_len) + estimate = estimate_cp_shared_kv_prefill_buffer_bytes( + page_size=self.page_size, + batch_size=len(self.cp_shared_kv_prefill_estimate_prefix_lens), + prefix_lens=self.cp_shared_kv_prefill_estimate_prefix_lens, + extend_lens=self.cp_shared_kv_prefill_estimate_extend_lens, + context=self.cp_shared_kv_prefill_buffer_estimator_context, + ) + self.cp_shared_kv_prefill_last_buffer_estimate = estimate + self.cp_shared_kv_prefill_estimated_peak_buffer_bytes = ( + estimate.total_peak_bytes + ) if self.dllm_config is not None: self.rem_dllm_tokens -= extend_input_len @@ -934,6 +1065,10 @@ class PrefillAdder: return AddReqResult.OTHER if self._cp_prefill_cached_limit_exceeded(prefix_len): return AddReqResult.OTHER + if self._cp_prefill_buffer_limit_exceeded( + prefix_len, input_tokens, req.rid + ): + return AddReqResult.OTHER if self.dllm_config is not None: if self.rem_dllm_tokens <= 0: @@ -980,6 +1115,10 @@ class PrefillAdder: return AddReqResult.OTHER if self._cp_prefill_cached_limit_exceeded(prefix_len): return AddReqResult.OTHER + if self._cp_prefill_buffer_limit_exceeded( + prefix_len, trunc_len, req.rid + ): + return AddReqResult.OTHER # Chunked prefill req.set_extend_input_len(trunc_len) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index f43709358..29fa1dc4d 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -147,6 +147,10 @@ from sglang.srt.managers.prefill_delayer import ( PrefillDelayer, PrefillDelayerSinglePassExecutor, ) +from sglang.srt.managers.cp_shared_kv_prefill_buffer_estimator import ( + CPSharedKVPrefillBufferEstimatorContext, + smoke_check_cp_shared_kv_prefill_buffer_size, +) from sglang.srt.managers.schedule_batch import ( FINISH_ABORT, ModelWorkerBatch, @@ -492,6 +496,8 @@ class Scheduler( # Init the grammar backend for constrained generation self.grammar_manager = GrammarManager(self) + self.maybe_smoke_check_cp_shared_kv_prefill_buffer() + self.is_initializing = False def init_model_config(self): @@ -515,6 +521,62 @@ class Scheduler( ) self.page_size = self.dllm_config.block_size + def make_cp_shared_kv_prefill_buffer_estimator_context( + self, + ) -> CPSharedKVPrefillBufferEstimatorContext: + return CPSharedKVPrefillBufferEstimatorContext( + kvcache=self.token_to_kv_pool_allocator.get_kvcache(), + model_config=self.model_config, + tp_size=self.tp_size, + page_size=self.page_size, + logprob_chunk_enabled=envs.SGLANG_ENABLE_LOGITS_PROCESSER_CHUNK.get(), + logprob_chunk_size=envs.SGLANG_LOGITS_PROCESSER_CHUNK_SIZE.get(), + # As of this scheduler path, CP shared-KV L1 prefetch explicitly + # gates off bs>1. Keep the estimate aligned with runtime until that + # path is enabled. + bs_gt1_l1_prefetch_enabled=False, + ) + + def maybe_smoke_check_cp_shared_kv_prefill_buffer(self) -> None: + size_bytes = self.server_args.cp_shared_kv_prefill_max_buffer_size + if not ( + self.server_args.enable_cp_shared_kv_prefill_bs_gt1 + and self.server_args.enable_nsa_prefill_cp_shared_kv + and size_bytes is not None + and self.server_args.disaggregation_mode in (None, "null", "prefill") + ): + return + + if not str(self.device).startswith("cuda"): + logger.warning( + "[CP_SHARED_KV_FALLBACK][prefill_buffer_smoke] " + "skip CUDA smoke allocation on non-CUDA device=%s size_bytes=%s", + self.device, + size_bytes, + ) + return + + logger.info( + "[CP_SHARED_KV_BS_GT1] prefill buffer smoke allocation begin: " + "rank=%s cp_rank=%s size_bytes=%s size_gb=%.3f", + self.tp_rank, + self.attn_cp_rank, + size_bytes, + size_bytes / 1e9, + ) + smoke_check_cp_shared_kv_prefill_buffer_size( + device=self.device, + size_bytes=size_bytes, + ) + logger.info( + "[CP_SHARED_KV_BS_GT1] prefill buffer smoke allocation passed: " + "rank=%s cp_rank=%s size_bytes=%s size_gb=%.3f", + self.tp_rank, + self.attn_cp_rank, + size_bytes, + size_bytes / 1e9, + ) + def init_ipc_channels(self, port_args: PortArgs): context = zmq.Context(2) self.idle_sleeper = None @@ -2422,6 +2484,14 @@ class Scheduler( cp_shared_kv_prefill_max_total_cached_tokens=( self.server_args.cp_shared_kv_prefill_max_total_cached_tokens ), + cp_shared_kv_prefill_max_buffer_size=( + self.server_args.cp_shared_kv_prefill_max_buffer_size + ), + cp_shared_kv_prefill_buffer_estimator_context=( + self.make_cp_shared_kv_prefill_buffer_estimator_context() + if self.server_args.cp_shared_kv_prefill_max_buffer_size is not None + else None + ), prefill_delayer_single_pass=prefill_delayer_single_pass, dllm_config=self.dllm_config, ) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 2a4189890..cf2a5a2dc 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -23,7 +23,9 @@ import json import logging import os import random +import re import tempfile +from decimal import Decimal from typing import Any, Callable, Dict, List, Literal, Optional, Union from sglang.srt.connector import ConnectorType @@ -72,6 +74,23 @@ from sglang.utils import is_in_ci logger = logging.getLogger(__name__) + +def human_readable_gb_size(value: str) -> int: + """Parse sizes where a bare number means decimal GB. + + Examples: + '8' -> 8000000000 + '8G' -> 8000000000 + '8Gi' -> 8589934592 + + Explicit suffixes reuse human_readable_int semantics. + """ + + value = value.strip() + if re.fullmatch(r"\d+(?:\.\d+)?", value): + return int(Decimal(value) * Decimal(10**9)) + return human_readable_int(value) + # Define constants DEFAULT_UVICORN_ACCESS_LOG_EXCLUDE_PREFIXES = () SAMPLING_BACKEND_CHOICES = {"flashinfer", "pytorch", "ascend"} @@ -678,6 +697,7 @@ class ServerArgs: cp_shared_kv_prefill_max_batch_requests: Optional[int] = None cp_shared_kv_prefill_max_total_extend_tokens: Optional[int] = None cp_shared_kv_prefill_max_total_cached_tokens: Optional[int] = None + cp_shared_kv_prefill_max_buffer_size: Optional[int] = None enable_fused_qk_norm_rope: bool = False enable_precise_embedding_interpolation: bool = False enable_fused_moe_sum_all_reduce: bool = False @@ -1002,6 +1022,14 @@ class ServerArgs: "cp_shared_kv_prefill_max_total_cached_tokens must be a positive " "integer when specified." ) + if ( + self.cp_shared_kv_prefill_max_buffer_size is not None + and self.cp_shared_kv_prefill_max_buffer_size <= 0 + ): + raise ValueError( + "cp_shared_kv_prefill_max_buffer_size must be a positive " + "integer when specified." + ) def _handle_cp_hicache_layout_validation(self): if not ( @@ -5919,6 +5947,20 @@ class ServerArgs: ) + f"\n\n{human_readable_int.__doc__}", ) + parser.add_argument( + "--cp-shared-kv-prefill-max-buffer-size", + type=human_readable_gb_size, + default=ServerArgs.cp_shared_kv_prefill_max_buffer_size, + help=( + "Maximum estimated peak temporary GPU buffer size admitted into " + "one NSA in-seq CP shared-KV prefill batch when " + "--enable-cp-shared-kv-prefill-bs-gt1 is set. Bare numbers are " + "interpreted as decimal GB (for example, 8 means 8G). Explicit " + "SI/IEC suffixes such as 8G or 8Gi are also accepted. A single " + "request larger than this limit is still allowed to run alone." + ) + + f"\n\n{human_readable_gb_size.__doc__}", + ) parser.add_argument( "--nsa-prefill-cp-mode", type=str, diff --git a/test/registered/unit/managers/test_cp_shared_kv_prefill_buffer_estimator.py b/test/registered/unit/managers/test_cp_shared_kv_prefill_buffer_estimator.py new file mode 100644 index 000000000..6a09378a7 --- /dev/null +++ b/test/registered/unit/managers/test_cp_shared_kv_prefill_buffer_estimator.py @@ -0,0 +1,151 @@ +from types import SimpleNamespace + +import torch + +from sglang.srt.managers.cp_shared_kv_prefill_buffer_estimator import ( + CPSharedKVPrefillBufferEstimatorContext, + estimate_cp_shared_kv_prefill_buffer_bytes, + smoke_check_cp_shared_kv_prefill_buffer_size, +) + + +def _fake_kvcache( + *, + kv_cache_dim: int = 4, + index_head_dim: int = 8, + store_dtype: torch.dtype = torch.bfloat16, +): + return SimpleNamespace( + kv_cache_dim=kv_cache_dim, + store_dtype=store_dtype, + index_head_dim=index_head_dim, + quant_block_size=4, + index_k_with_scale_buffer_dtype=torch.uint8, + ) + + +def test_estimator_uses_stream_aware_peak_instead_of_independent_max(): + estimate = estimate_cp_shared_kv_prefill_buffer_bytes( + page_size=4, + batch_size=2, + prefix_lens=[5, 0], + extend_lens=[3, 4], + context=CPSharedKVPrefillBufferEstimatorContext( + kvcache=_fake_kvcache(), + model_config=SimpleNamespace(vocab_size=32), + tp_size=1, + page_size=4, + logprob_chunk_enabled=False, + logprob_chunk_size=2048, + bs_gt1_l1_prefetch_enabled=True, + ), + ) + + assert estimate.prefetch_peak_bytes > 0 + assert estimate.layer_forward_peak_bytes == ( + estimate.materialize_peak_bytes + + estimate.remap_peak_bytes + + estimate.prefetch_peak_bytes + + estimate.backup_descriptor_peak_bytes + ) + assert estimate.total_peak_bytes == max( + estimate.layer_forward_peak_bytes, + estimate.logits_window_peak_bytes, + estimate.load_back_window_peak_bytes, + ) + assert estimate.total_peak_bytes > max( + estimate.materialize_peak_bytes, + estimate.prefetch_peak_bytes, + estimate.logits_peak_bytes, + ) + + +def test_estimator_keeps_bs_gt1_prefetch_zero_until_enabled(): + estimate = estimate_cp_shared_kv_prefill_buffer_bytes( + page_size=64, + batch_size=1, + prefix_lens=[128], + extend_lens=[64], + context=CPSharedKVPrefillBufferEstimatorContext( + kvcache=_fake_kvcache(), + model_config=SimpleNamespace(vocab_size=32), + tp_size=1, + page_size=64, + logprob_chunk_enabled=False, + logprob_chunk_size=2048, + bs_gt1_l1_prefetch_enabled=False, + ), + ) + + assert estimate.prefetch_peak_bytes == 0 + assert estimate.total_peak_bytes >= estimate.materialize_peak_bytes + + +def test_smoke_check_allocates_and_releases_probe_with_device_module(monkeypatch): + events = [] + + class FakeCuda: + class OutOfMemoryError(RuntimeError): + pass + + @staticmethod + def synchronize(): + events.append("sync") + + @staticmethod + def empty_cache(): + events.append("empty_cache") + + class FakeProbe: + def __setitem__(self, index, value): + events.append(("set", index, value)) + + fake_torch = SimpleNamespace( + cuda=FakeCuda, + uint8=object(), + empty=lambda size, dtype, device: events.append( + ("empty", size, dtype, device) + ) + or FakeProbe(), + ) + + smoke_check_cp_shared_kv_prefill_buffer_size( + device="cuda:0", size_bytes=16, torch_module=fake_torch + ) + + assert events == [ + "sync", + ("empty", 16, fake_torch.uint8, "cuda:0"), + ("set", 0, 0), + ("set", -1, 0), + "sync", + "empty_cache", + ] + + +def test_smoke_check_raises_fail_fast_on_oom(): + class FakeCuda: + class OutOfMemoryError(RuntimeError): + pass + + @staticmethod + def synchronize(): + return None + + @staticmethod + def empty_cache(): + return None + + def _raise_oom(size, dtype, device): + raise FakeCuda.OutOfMemoryError("oom") + + fake_torch = SimpleNamespace(cuda=FakeCuda, uint8=object(), empty=_raise_oom) + + try: + smoke_check_cp_shared_kv_prefill_buffer_size( + device="cuda:0", size_bytes=16, torch_module=fake_torch + ) + except RuntimeError as exc: + assert "[CP_SHARED_KV_FAIL_FAST][prefill_buffer_smoke]" in str(exc) + else: + raise AssertionError("expected fail-fast RuntimeError") diff --git a/test/registered/unit/managers/test_prefill_adder.py b/test/registered/unit/managers/test_prefill_adder.py index 7778a7eca..b22cb0a4b 100644 --- a/test/registered/unit/managers/test_prefill_adder.py +++ b/test/registered/unit/managers/test_prefill_adder.py @@ -63,6 +63,9 @@ for _schema in ( raise from sglang.srt.managers.schedule_batch import Req +from sglang.srt.managers.cp_shared_kv_prefill_buffer_estimator import ( + CPSharedKVPrefillBufferEstimatorContext, +) from sglang.srt.managers.schedule_policy import AddReqResult, PrefillAdder from sglang.srt.mem_cache.base_prefix_cache import ( DecLockRefResult, @@ -168,6 +171,23 @@ class TestPrefillAdder(CustomTestCase): defaults.update(kwargs) return PrefillAdder(**defaults) + def create_buffer_estimator_context(self, *, kv_cache_dim=1, vocab_size=16): + return CPSharedKVPrefillBufferEstimatorContext( + kvcache=SimpleNamespace( + kv_cache_dim=kv_cache_dim, + store_dtype=torch.bfloat16, + index_head_dim=8, + quant_block_size=4, + index_k_with_scale_buffer_dtype=torch.uint8, + ), + model_config=SimpleNamespace(vocab_size=vocab_size), + tp_size=1, + page_size=64, + logprob_chunk_enabled=False, + logprob_chunk_size=2048, + bs_gt1_l1_prefetch_enabled=False, + ) + def test_preempt_success_high_priority_values_first(self): params = [ ("run1", 0, 50), @@ -708,6 +728,32 @@ class TestPrefillAdder(CustomTestCase): self.assertEqual([req.rid for req in adder.can_run_list], ["first", "second"]) self.assertEqual(adder.cp_shared_kv_prefill_total_extend_tokens, 256) + def test_cp_prefill_total_extend_limit_is_capped_by_chunked_prefill_size(self): + set_global_server_args_for_scheduler( + ServerArgs( + model_path="dummy", + enable_nsa_prefill_context_parallel=True, + nsa_prefill_cp_mode="in-seq-split", + ) + ) + adder = self.create_adder( + self.create_running_batch(), + page_size=64, + # The CP-specific extend limit is larger than the chunked prefill + # budget. Effective admission should use the smaller chunk budget + # to avoid advertising an unreachable per-batch extend capacity. + rem_input_tokens=192, + rem_chunk_tokens=128, + enable_cp_shared_kv_prefill_bs_gt1=True, + cp_shared_kv_prefill_max_batch_requests=8, + cp_shared_kv_prefill_max_total_extend_tokens=256, + ) + + self.assertEqual(adder.cp_shared_kv_prefill_max_total_extend_tokens, 128) + # The generic max_prefill_tokens lift should also use the effective + # limit, not the raw 256-token CP limit. + self.assertEqual(adder.rem_input_tokens, 192) + def test_cp_prefill_total_extend_limit_does_not_bypass_allocator_capacity(self): set_global_server_args_for_scheduler( ServerArgs( @@ -843,6 +889,75 @@ class TestPrefillAdder(CustomTestCase): self.assertEqual([req.rid for req in adder.can_run_list], ["oversized"]) self.assertEqual(adder.cp_shared_kv_prefill_total_cached_tokens, 8192) + def test_cp_prefill_buffer_limit_stops_second_request_without_token_gate(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, + cp_shared_kv_prefill_max_buffer_size=1, + cp_shared_kv_prefill_buffer_estimator_context=( + self.create_buffer_estimator_context() + ), + ) + + first = self.create_prefill_req("first", extend_input_len=64) + second = self.create_prefill_req("second", 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.OTHER, + ) + self.assertEqual([req.rid for req in adder.can_run_list], ["first"]) + self.assertGreater(adder.cp_shared_kv_prefill_estimated_peak_buffer_bytes, 1) + + def test_cp_prefill_buffer_limit_allows_single_oversized_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_buffer_size=1, + cp_shared_kv_prefill_buffer_estimator_context=( + self.create_buffer_estimator_context() + ), + ) + + oversized = self.create_prefill_req("oversized", extend_input_len=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"]) + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/managers/test_prefill_scheduler_admission_bench.py b/test/registered/unit/managers/test_prefill_scheduler_admission_bench.py index fb4a462fd..4ef2cd856 100644 --- a/test/registered/unit/managers/test_prefill_scheduler_admission_bench.py +++ b/test/registered/unit/managers/test_prefill_scheduler_admission_bench.py @@ -136,6 +136,39 @@ class TestPrefillSchedulerAdmissionBench(unittest.TestCase): self.assertEqual(trace.ticks[0].stopped_result, "OTHER") self.assertEqual(trace.ticks[0].cp_total_cached_tokens, 4096) + def test_max_buffer_size_is_observable_in_scheduler_trace(self): + bench = _load_bench_module() + cfg = bench.SchedulerBenchConfig( + page_size=64, + available_tokens=10000, + evictable_tokens=0, + max_prefill_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, + cp_shared_kv_prefill_max_buffer_size=1, + kv_cache_dim=1, + kv_dtype_bytes=2, + index_head_dim=8, + vocab_size=16, + tp_size=1, + ) + trace = bench.run_scheduler_admission_trace( + [ + bench.RequestSpec("a", 0, 0, 64, max_new_tokens=1), + bench.RequestSpec("b", 0, 0, 64, max_new_tokens=1), + ], + cfg, + ) + + tick = trace.ticks[0] + self.assertEqual([req.rid for req in tick.accepted], ["a"]) + self.assertEqual(tick.stopped_on_rid, "b") + self.assertEqual(tick.stopped_result, "OTHER") + self.assertGreater(tick.cp_estimated_peak_buffer_bytes, 1) + self.assertIn("layer_forward_peak_bytes", tick.cp_buffer_breakdown) + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index 1c9f98c4d..c214369ee 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -3,6 +3,8 @@ import tempfile import unittest from unittest.mock import MagicMock, patch +import pytest + from sglang.srt.server_args import PortArgs, ServerArgs, prepare_server_args from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import ( @@ -75,6 +77,60 @@ def test_cp_shared_kv_prefill_bs_gt1_parser_limits(): assert args.cp_shared_kv_prefill_max_total_extend_tokens == 8192 +def test_cp_shared_kv_prefill_max_buffer_size_defaults_to_gb(): + 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-buffer-size", + "8", + ] + ) + args = ServerArgs.from_cli_args(raw_args) + assert args.cp_shared_kv_prefill_max_buffer_size == 8_000_000_000 + + +def test_cp_shared_kv_prefill_max_buffer_size_accepts_iec_suffix(): + 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-buffer-size", + "8Gi", + ] + ) + args = ServerArgs.from_cli_args(raw_args) + assert args.cp_shared_kv_prefill_max_buffer_size == 8 * 2**30 + + +def test_cp_shared_kv_prefill_max_buffer_size_rejects_non_positive_value(): + 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-buffer-size", + "0", + ] + ) + with pytest.raises(ValueError, match="cp_shared_kv_prefill_max_buffer_size"): + ServerArgs.from_cli_args(raw_args) + + def test_hicache_mem_layout_parser_accepts_layer_page_first(): import argparse