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