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
This commit is contained in:
laoyao0822
2026-06-11 01:33:28 +08:00
parent 9a9893e571
commit 3a43727216
10 changed files with 1806 additions and 3 deletions
@@ -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")
@@ -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()
@@ -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()
@@ -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