Files
sglang/test/registered/unit/managers/test_prefill_adder.py
leavelet 6a1e862f48 Group cache-hit prefills into dense batches (SGLANG_CP_PREFILL_AFFINITY_GROUP)
专题 S4 (design: docs_internal/perf/prefill-compute-intensity-plan.md S4,
amended).  Under FCFS a cold request joining a warm-led batch turns a
1-2s cache-hit forward into a 5-10s one, splitting the warm work into
the 新-cache-新 pattern.  The policy prevents exactly that one thing:

- WARM candidates always admit (into a cold-led batch they are free
  density — the cold extend dominates the forward anyway).
- COLD admits into an empty or cold-led batch (small colds co-batch
  today; the FCFS head always starts a batch so the queue keeps moving).
- COLD into a WARM-led batch is skipped, bounded by a per-pass window
  (W=16 skips), a head defer count (K=3 passes) and an age bound
  (T=5s).  On any bound the scan STOPS instead of force-admitting: the
  cold waits for the same forward either way, but leads its own clean
  batch next pass instead of polluting this one.

The skip is strictly post-match / pre-admit (after init_next_round_input,
before add_one_req): no lock, no allocation, no budget mutation to
unwind, and re-matching a skipped candidate next pass is exactly what
the scan already does after a cap rejection.  Classification is the
in-scan match result (device prefix + host hit vs a 64-token floor) —
under FCFS+L2 no pre-scan signal exists, so this adds zero matching
work for inspected candidates.  Disabled wholesale under priority
scheduling (the skip must not reorder across priority classes).

Three amendments vs the design draft, reasoned in the decision-table
docstring: cold+cold-led admits (STOP would regress today's small-cold
co-batching); starved heads STOP rather than force-admit (clean batch
boundaries at identical latency); priority interaction handled by
disabling rather than per-request comparison.

Decision logic is a pure function with table + bounds unit tests
(28/28 adder suite green).  Default OFF.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-06-12 07:44:39 +00:00

1206 lines
47 KiB
Python

import sys
import types
import unittest
from types import SimpleNamespace
from unittest.mock import MagicMock
import torch
from sglang.srt.environ import envs
if "sgl_kernel" not in sys.modules:
sys.modules["sgl_kernel"] = types.ModuleType("sgl_kernel")
sys.modules["sgl_kernel"].__file__ = "sgl_kernel_stub.py"
sys.modules["sgl_kernel"].__path__ = []
if not hasattr(sys.modules["sgl_kernel"], "__getattr__"):
def _sgl_kernel_getattr(name):
if name.startswith("__"):
raise AttributeError(name)
fn = lambda *args, **kwargs: None
setattr(sys.modules["sgl_kernel"], name, fn)
return fn
sys.modules["sgl_kernel"].__getattr__ = _sgl_kernel_getattr
if "sgl_kernel.kvcacheio" not in sys.modules:
sys.modules["sgl_kernel.kvcacheio"] = types.ModuleType("sgl_kernel.kvcacheio")
for _name in (
"sgl_per_token_group_quant_8bit",
"sgl_per_token_group_quant_fp8",
"sgl_per_token_quant_fp8",
"fp8_blockwise_scaled_mm",
"fp8_scaled_mm",
"silu_and_mul",
):
if not hasattr(sys.modules["sgl_kernel"], _name):
setattr(sys.modules["sgl_kernel"], _name, lambda *args, **kwargs: None)
if "sgl_kernel.quantization" not in sys.modules:
quantization_stub = types.ModuleType("sgl_kernel.quantization")
for _name in (
"ggml_dequantize",
"ggml_moe_a8",
"ggml_moe_a8_vec",
"ggml_moe_get_block_size",
"ggml_mul_mat_a8",
"ggml_mul_mat_vec_a8",
):
setattr(quantization_stub, _name, lambda *args, **kwargs: None)
sys.modules["sgl_kernel.quantization"] = quantization_stub
_sgl_kernel_lib = torch.library.Library("sgl_kernel", "FRAGMENT")
for _schema in (
"sgl_per_token_group_quant_8bit(Tensor input, Tensor(a!) output_q, Tensor(b!) output_s, int group_size, float eps, float fp8_min, float fp8_max, bool scale_ue8m0) -> ()",
"sgl_per_token_group_quant_fp8(Tensor input, Tensor(a!) output_q, Tensor(b!) output_s, int group_size, float eps, float fp8_min, float fp8_max, bool scale_ue8m0) -> ()",
"sgl_per_token_quant_fp8(Tensor input, Tensor(a!) output_q, Tensor(b!) output_s) -> ()",
"fp8_scaled_mm(Tensor mat_a, Tensor mat_b, Tensor scales_a, Tensor scales_b, ScalarType out_dtype, Tensor? bias=None) -> Tensor",
"fp8_blockwise_scaled_mm(Tensor mat_a, Tensor mat_b, Tensor scales_a, Tensor scales_b, ScalarType out_dtype) -> Tensor",
):
try:
_sgl_kernel_lib.define(_schema)
except RuntimeError as exc:
if "already" not in str(exc).lower() and "duplicate" not in str(exc).lower():
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,
IncLockRefResult,
)
from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=1, suite="stage-b-test-1-gpu-small")
register_amd_ci(est_time=2, suite="stage-b-test-1-gpu-small-amd")
class TestPrefillAdder(CustomTestCase):
def setUp(self):
set_global_server_args_for_scheduler(ServerArgs(model_path="dummy"))
self.mock_tree_cache = self.create_tree_cache()
self.mock_token_allocator = self.create_token_allocator()
def create_tree_cache(
self,
*,
full_evictable_size: int = 0,
swa_evictable_size: int = 0,
evictable_size: int = 0,
) -> MagicMock:
tree_cache = MagicMock()
tree_cache.full_evictable_size.return_value = full_evictable_size
tree_cache.swa_evictable_size.return_value = swa_evictable_size
tree_cache.evictable_size.return_value = evictable_size
tree_cache.supports_mamba.return_value = False
tree_cache.disable = False
tree_cache.inc_lock_ref.return_value = IncLockRefResult()
tree_cache.dec_lock_ref.return_value = DecLockRefResult()
return tree_cache
def create_token_allocator(
self,
*,
full_available_size: int = 0,
swa_available_size: int = 0,
available_size: int = 0,
) -> MagicMock:
allocator = MagicMock()
allocator.full_available_size.return_value = full_available_size
allocator.swa_available_size.return_value = swa_available_size
allocator.available_size.return_value = available_size
return allocator
def create_running_batch(self, reqs=None) -> MagicMock:
batch = MagicMock()
batch.reqs = list(reqs or [])
batch.release_req.return_value = None
batch.filter_batch.return_value = None
return batch
def create_server_args(
self, *, schedule_low_priority_values_first: bool
) -> MagicMock:
server_args = MagicMock()
server_args.schedule_low_priority_values_first = (
schedule_low_priority_values_first
)
return server_args
def create_mock_req(self, rid, priority, max_new_tokens, output_len=0, wait_time=0):
req = MagicMock(spec=Req)
req.rid = str(rid)
req.priority = priority
req.extend_input_len = 0
req.extend_logprob_start_len = 0
req.output_ids = [0] * output_len
req.sampling_params = SimpleNamespace(max_new_tokens=max_new_tokens)
req.time_stats = SimpleNamespace(wait_queue_entry_time=wait_time)
req.finished.return_value = False
return req
def create_prefill_req(self, rid, extend_input_len, max_new_tokens=1):
req = self.create_mock_req(rid, priority=0, max_new_tokens=max_new_tokens)
req.extend_input_len = extend_input_len
req.host_hit_length = 0
req.prefix_indices = torch.empty((0,), dtype=torch.int64)
req.fill_ids = list(range(extend_input_len))
req.last_node = object()
req.sampling_params.ignore_eos = False
req.set_extend_input_len.side_effect = lambda value: setattr(
req, "extend_input_len", value
)
return req
def create_adder(self, running_batch, **kwargs):
defaults = dict(
page_size=1,
tree_cache=self.mock_tree_cache,
token_to_kv_pool_allocator=self.mock_token_allocator,
running_batch=running_batch,
new_token_ratio=1.0,
rem_input_tokens=10000,
rem_chunk_tokens=None,
mixed_with_decode_tokens=0,
priority_scheduling_preemption_threshold=0,
)
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),
("run2", 1, 75),
("run3", 2, 100),
]
running_reqs = [
self.create_mock_req(rid, priority, max_new_tokens)
for rid, priority, max_new_tokens in params
]
mock_server_args = self.create_server_args(
schedule_low_priority_values_first=False
)
running_batch = self.create_running_batch(running_reqs)
adder = self.create_adder(running_batch)
self.assertEqual(adder.rem_total_token_offset, 225)
self.mock_token_allocator.full_available_size.return_value = (
225 # full occupation of GRam
)
self.mock_token_allocator.available_size.return_value = 225
new_req = self.create_mock_req("new1", priority=1, max_new_tokens=49)
success = adder.preempt_to_schedule(new_req, mock_server_args)
self.assertTrue(success)
self.assertIn(running_reqs[0], adder.preempt_list)
self.assertEqual(adder.rem_total_token_offset, 175) # 50 + 75 + 100 - 50 = 175
running_batch.release_req.assert_called_once()
def test_preempt_success_low_priority_values_first(self):
params = [
("run1", 0, 50),
("run2", 1, 75),
("run3", 2, 100),
]
running_reqs = [
self.create_mock_req(rid, priority, max_new_tokens)
for rid, priority, max_new_tokens in params
]
mock_server_args = self.create_server_args(
schedule_low_priority_values_first=True
)
running_batch = self.create_running_batch(running_reqs)
adder = self.create_adder(running_batch)
self.assertEqual(adder.rem_total_token_offset, 225)
self.mock_token_allocator.full_available_size.return_value = (
225 # full occupation of GRam
)
self.mock_token_allocator.available_size.return_value = 225
new_req = self.create_mock_req("new1", priority=1, max_new_tokens=49)
success = adder.preempt_to_schedule(new_req, mock_server_args)
self.assertTrue(success)
self.assertIn(running_reqs[2], adder.preempt_list)
self.assertEqual(adder.rem_total_token_offset, 125) # 50 + 75 + 100 - 100 = 125
running_batch.release_req.assert_called_once()
def test_preempt_fail_low_priority_values_first(self):
params = [
("run1", 0, 50),
("run2", 1, 75),
("run3", 2, 100),
]
running_reqs = [
self.create_mock_req(rid, priority, max_new_tokens)
for rid, priority, max_new_tokens in params
]
mock_server_args = self.create_server_args(
schedule_low_priority_values_first=True
)
running_batch = self.create_running_batch(running_reqs)
adder = self.create_adder(running_batch)
self.assertEqual(adder.rem_total_token_offset, 225)
self.mock_token_allocator.full_available_size.return_value = (
225 # full occupation of GRam
)
self.mock_token_allocator.available_size.return_value = 225
new_req_fail_by_priority_check = self.create_mock_req(
"new1", priority=2, max_new_tokens=49
)
success_by_priority_check = adder.preempt_to_schedule(
new_req_fail_by_priority_check, mock_server_args
)
self.assertFalse(success_by_priority_check)
new_req_fail_by_priority_check = self.create_mock_req(
"new2", priority=1, max_new_tokens=110
)
success_by_capacity_check = adder.preempt_to_schedule(
new_req_fail_by_priority_check, mock_server_args
)
self.assertFalse(success_by_capacity_check)
def test_preempt_fail_high_priority_values_first(self):
params = [
("run1", 0, 50),
("run2", 1, 75),
("run3", 2, 100),
]
running_reqs = [
self.create_mock_req(rid, priority, max_new_tokens)
for rid, priority, max_new_tokens in params
]
mock_server_args = self.create_server_args(
schedule_low_priority_values_first=False
)
running_batch = self.create_running_batch(running_reqs)
adder = self.create_adder(running_batch)
self.assertEqual(adder.rem_total_token_offset, 225)
self.mock_token_allocator.full_available_size.return_value = (
225 # full occupation of GRam
)
self.mock_token_allocator.available_size.return_value = 225
new_req_fail_by_priority_check = self.create_mock_req(
"new1", priority=0, max_new_tokens=49
)
success_by_priority_check = adder.preempt_to_schedule(
new_req_fail_by_priority_check, mock_server_args
)
self.assertFalse(success_by_priority_check)
new_req_fail_by_priority_check = self.create_mock_req(
"new2", priority=-1, max_new_tokens=110
)
success_by_capacity_check = adder.preempt_to_schedule(
new_req_fail_by_priority_check, mock_server_args
)
self.assertFalse(success_by_capacity_check)
def test_preempt_skip_already_preempted_request(self):
params = [
("req_prio_0", 0, 50),
("req_prio_1", 1, 75),
("req_prio_2", 2, 100),
]
running_reqs = [
self.create_mock_req(rid, priority, max_new_tokens)
for rid, priority, max_new_tokens in params
]
mock_server_args = self.create_server_args(
schedule_low_priority_values_first=False
)
running_batch = self.create_running_batch(running_reqs)
adder = self.create_adder(running_batch)
self.assertEqual(adder.rem_total_token_offset, 225)
self.mock_token_allocator.full_available_size.return_value = 225
self.mock_token_allocator.available_size.return_value = 225
# New request preempts req_prio_0
first_req = self.create_mock_req(
"new_req_prio_1", priority=1, max_new_tokens=49
)
first_success = adder.preempt_to_schedule(first_req, mock_server_args)
self.assertTrue(first_success)
self.assertIn(running_reqs[0], adder.preempt_list)
self.assertEqual(adder.rem_total_token_offset, 175)
running_batch.release_req.assert_called_once()
# Second call needs more tokens than currently free, so it would need to
# preempt req_prio_0 again if already-preempted requests were not filtered out.
second_req = self.create_mock_req(
"second_new_req_prio_1", priority=1, max_new_tokens=76
)
second_success = adder.preempt_to_schedule(second_req, mock_server_args)
self.assertFalse(second_success)
self.assertEqual(adder.rem_total_token_offset, 175)
self.assertEqual(adder.preempt_list.count(running_reqs[0]), 1)
running_batch.release_req.assert_called_once()
def test_preempt_success_low_priority_values_first_exact_once(self):
params = [
("run1", 0, 50),
("run2", 1, 75),
("run3", 2, 100),
("run4", 2, 125),
("run4", 2, 125),
]
running_reqs = [
self.create_mock_req(rid, priority, max_new_tokens)
for rid, priority, max_new_tokens in params
]
mock_server_args = self.create_server_args(
schedule_low_priority_values_first=True
)
running_batch = self.create_running_batch(running_reqs)
adder = self.create_adder(running_batch)
self.assertEqual(adder.rem_total_token_offset, 475)
self.mock_token_allocator.full_available_size.return_value = (
475 # full occupation of GRam
)
self.mock_token_allocator.available_size.return_value = 475
new_req = self.create_mock_req("new1", priority=1, max_new_tokens=75)
success = adder.preempt_to_schedule(new_req, mock_server_args)
self.assertTrue(success)
self.assertIn(running_reqs[2], adder.preempt_list)
self.assertEqual(
adder.rem_total_token_offset, 375
) # 50 + 75 + 100 + 125 + 125 - 100 = 375
running_batch.release_req.assert_called_once()
def test_preempt_success_low_priority_values_first_exact_twice(self):
params = [
("run1", 0, 50),
("run2", 1, 75),
("run3", 2, 100),
("run4", 2, 125),
("run4", 2, 125),
]
running_reqs = [
self.create_mock_req(rid, priority, max_new_tokens)
for rid, priority, max_new_tokens in params
]
mock_server_args = self.create_server_args(
schedule_low_priority_values_first=True
)
running_batch = self.create_running_batch(running_reqs)
adder = self.create_adder(running_batch)
self.assertEqual(adder.rem_total_token_offset, 475)
self.mock_token_allocator.full_available_size.return_value = (
475 # full occupation of GRam
)
self.mock_token_allocator.available_size.return_value = 475
new_req = self.create_mock_req("new1", priority=1, max_new_tokens=200)
success = adder.preempt_to_schedule(new_req, mock_server_args)
self.assertTrue(success)
self.assertIn(running_reqs[2], adder.preempt_list)
self.assertIn(running_reqs[3], adder.preempt_list)
self.assertEqual(
adder.rem_total_token_offset, 250
) # 50 + 75 + 100 + 125 + 125 - 100 - 125 = 250
self.assertEqual(running_batch.release_req.call_count, 2)
def test_mixed_chunk_prefill_budgets(self):
self.mock_token_allocator.available_size.return_value = 1000
decode_reqs = [
self.create_mock_req(f"decode_{i}", priority=0, max_new_tokens=50)
for i in range(8)
]
running_batch = self.create_running_batch(decode_reqs)
adder = self.create_adder(
running_batch,
rem_input_tokens=200,
rem_chunk_tokens=64,
mixed_with_decode_tokens=len(decode_reqs),
)
self.assertEqual(adder.rem_input_tokens, 192) # 200 - 8
self.assertEqual(adder.rem_chunk_tokens, 56) # 64 - 8
self.assertEqual(adder.rem_total_token_offset, 408) # 8 + 8 * 50
self.assertEqual(adder.cur_rem_token_offset, 8)
self.assertEqual(adder.budget_state(), AddReqResult.CONTINUE)
# Add a prefill that exactly consumes the chunk budget
req1 = self.create_mock_req("req1", priority=0, max_new_tokens=64)
req1.extend_input_len = 56
req1.host_hit_length = 0
req1.prefix_indices = []
req1.fill_ids = list(range(56))
req1.last_node = MagicMock()
req1.sampling_params.ignore_eos = False
result1 = adder.add_one_req(
req1, has_chunked_req=False, truncation_align_size=None
)
self.assertEqual(len(adder.can_run_list), 1)
self.assertEqual(adder.rem_chunk_tokens, 0) # 56 - 56
self.assertEqual(adder.rem_input_tokens, 136) # 192 - 56
self.assertEqual(result1, AddReqResult.OTHER)
# 3 decode requests finished
remaining_decode_reqs = decode_reqs[3:]
running_batch2 = self.create_running_batch(remaining_decode_reqs)
adder2 = self.create_adder(
running_batch2,
rem_input_tokens=200,
rem_chunk_tokens=64,
mixed_with_decode_tokens=len(remaining_decode_reqs),
)
self.assertEqual(adder2.rem_input_tokens, 195) # 200 - 5
self.assertEqual(adder2.rem_chunk_tokens, 59) # 64 - 5
self.assertEqual(adder2.rem_total_token_offset, 255) # 5 + 5 * 50
self.assertEqual(adder2.budget_state(), AddReqResult.CONTINUE)
# Same prefill no longer exhausts the chunk budget
req2 = self.create_mock_req("req2", priority=0, max_new_tokens=64)
req2.extend_input_len = 56
req2.host_hit_length = 0
req2.prefix_indices = []
req2.fill_ids = list(range(56))
req2.last_node = MagicMock()
req2.sampling_params.ignore_eos = False
result2 = adder2.add_one_req(
req2, has_chunked_req=False, truncation_align_size=None
)
self.assertEqual(len(adder2.can_run_list), 1)
self.assertEqual(adder2.rem_chunk_tokens, 3) # 59 - 56 = 3 remaining
self.assertEqual(result2, AddReqResult.CONTINUE)
# Fit last small prefill request
req3 = self.create_mock_req("req3", priority=0, max_new_tokens=16)
req3.extend_input_len = 3
req3.host_hit_length = 0
req3.prefix_indices = []
req3.fill_ids = list(range(3))
req3.last_node = MagicMock()
req3.sampling_params.ignore_eos = False
result3 = adder2.add_one_req(
req3, has_chunked_req=False, truncation_align_size=None
)
self.assertEqual(len(adder2.can_run_list), 2)
self.assertEqual(adder2.rem_chunk_tokens, 0) # 3 - 3 = 0
self.assertEqual(result3, AddReqResult.OTHER)
def test_host_load_back_passes_mem_quota(self):
running_batch = self.create_running_batch()
self.mock_token_allocator.available_size.return_value = 512
self.mock_tree_cache.init_load_back.return_value = (
__import__("torch").tensor([1, 2, 3, 4], dtype=__import__("torch").int64),
"loaded_node",
)
adder = self.create_adder(
running_batch,
page_size=64,
rem_input_tokens=4096,
)
req = self.create_mock_req("req", priority=0, max_new_tokens=16)
req.extend_input_len = 256
req.host_hit_length = 128
req.prefix_indices = __import__("torch").empty(
(0,), dtype=__import__("torch").int64
)
req.last_node = object()
req.last_host_node = object()
req.fill_ids = list(range(256))
req.cache_protected_len = 0
req.set_extend_input_len = lambda value: setattr(req, "extend_input_len", value)
req.sampling_params.ignore_eos = False
result = adder.add_one_req(
req, has_chunked_req=False, truncation_align_size=None
)
self.assertNotEqual(result, AddReqResult.NO_TOKEN)
params = self.mock_tree_cache.init_load_back.call_args.args[0]
self.assertEqual(params.mem_quota, 320)
def test_load_back_mem_quota_counts_evictable_device_tokens(self):
self.mock_tree_cache = self.create_tree_cache(evictable_size=90000)
self.mock_token_allocator = self.create_token_allocator(available_size=1024)
adder = self.create_adder(
self.create_running_batch(),
page_size=64,
tree_cache=self.mock_tree_cache,
token_to_kv_pool_allocator=self.mock_token_allocator,
)
quota = adder._get_load_back_mem_quota(real_input_tokens=65536)
self.assertEqual(quota, 90000 + 1024 - 65536 - 64)
def test_cp_prefill_gate_keeps_single_request_by_default(self):
set_global_server_args_for_scheduler(
ServerArgs(
model_path="dummy",
enable_nsa_prefill_context_parallel=True,
nsa_prefill_cp_mode="in-seq-split",
)
)
self.mock_token_allocator.available_size.return_value = 10000
adder = self.create_adder(
self.create_running_batch(),
page_size=64,
rem_input_tokens=4096,
)
first = self.create_prefill_req("first", extend_input_len=128)
second = self.create_prefill_req("second", extend_input_len=128)
self.assertEqual(
adder.add_one_req(first, has_chunked_req=False, truncation_align_size=None),
AddReqResult.CONTINUE,
)
self.assertEqual(
adder.add_one_req(second, has_chunked_req=False, truncation_align_size=None),
AddReqResult.OTHER,
)
self.assertEqual([req.rid for req in adder.can_run_list], ["first"])
def test_cp_prefill_gate_allows_batched_requests_when_enabled(self):
set_global_server_args_for_scheduler(
ServerArgs(
model_path="dummy",
enable_nsa_prefill_context_parallel=True,
nsa_prefill_cp_mode="in-seq-split",
)
)
self.mock_token_allocator.available_size.return_value = 10000
adder = self.create_adder(
self.create_running_batch(),
page_size=64,
rem_input_tokens=4096,
enable_cp_shared_kv_prefill_bs_gt1=True,
cp_shared_kv_prefill_max_batch_requests=2,
cp_shared_kv_prefill_max_total_extend_tokens=256,
)
first = self.create_prefill_req("first", extend_input_len=128)
second = self.create_prefill_req("second", extend_input_len=128)
third = self.create_prefill_req("third", extend_input_len=64)
self.assertEqual(
adder.add_one_req(first, has_chunked_req=False, truncation_align_size=None),
AddReqResult.CONTINUE,
)
self.assertEqual(
adder.add_one_req(second, has_chunked_req=False, truncation_align_size=None),
AddReqResult.CONTINUE,
)
self.assertEqual(
adder.add_one_req(third, has_chunked_req=False, truncation_align_size=None),
AddReqResult.OTHER,
)
self.assertEqual([req.rid for req in adder.can_run_list], ["first", "second"])
def test_cp_prefill_total_extend_limit_is_page_aligned_and_allows_first_req(self):
set_global_server_args_for_scheduler(
ServerArgs(
model_path="dummy",
enable_nsa_prefill_context_parallel=True,
nsa_prefill_cp_mode="in-seq-split",
)
)
self.mock_token_allocator.available_size.return_value = 10000
adder = self.create_adder(
self.create_running_batch(),
page_size=64,
rem_input_tokens=4096,
enable_cp_shared_kv_prefill_bs_gt1=True,
cp_shared_kv_prefill_max_batch_requests=8,
cp_shared_kv_prefill_max_total_extend_tokens=128,
)
first = self.create_prefill_req("first", extend_input_len=65)
second = self.create_prefill_req("second", extend_input_len=1)
oversized_first = self.create_adder(
self.create_running_batch(),
page_size=64,
rem_input_tokens=4096,
enable_cp_shared_kv_prefill_bs_gt1=True,
cp_shared_kv_prefill_max_batch_requests=8,
cp_shared_kv_prefill_max_total_extend_tokens=64,
)
large = self.create_prefill_req("large", extend_input_len=128)
self.assertEqual(
adder.add_one_req(first, has_chunked_req=False, truncation_align_size=None),
AddReqResult.CONTINUE,
)
self.assertEqual(adder.cp_shared_kv_prefill_total_extend_tokens, 128)
self.assertEqual(
adder.add_one_req(second, has_chunked_req=False, truncation_align_size=None),
AddReqResult.OTHER,
)
self.assertEqual(
oversized_first.add_one_req(
large, has_chunked_req=False, truncation_align_size=None
),
AddReqResult.CONTINUE,
)
self.assertEqual([req.rid for req in oversized_first.can_run_list], ["large"])
def test_cp_prefill_total_extend_limit_replaces_generic_input_budget(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,
# Simulates the generic max_prefill_tokens budget being smaller
# than the CP shared-KV bs>1 budget.
rem_input_tokens=192,
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,
)
first = self.create_prefill_req("first", extend_input_len=128)
second = self.create_prefill_req("second", extend_input_len=128)
self.assertEqual(
adder.add_one_req(first, has_chunked_req=False, truncation_align_size=None),
AddReqResult.CONTINUE,
)
second_result = adder.add_one_req(
second, has_chunked_req=False, truncation_align_size=None
)
self.assertNotEqual(second_result, AddReqResult.NO_TOKEN)
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_defaults_to_chunk_budget(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,
rem_input_tokens=192,
rem_chunk_tokens=128,
enable_cp_shared_kv_prefill_bs_gt1=True,
cp_shared_kv_prefill_max_batch_requests=None,
cp_shared_kv_prefill_max_total_extend_tokens=None,
cp_shared_kv_prefill_max_total_cached_tokens=None,
)
self.assertIsNone(adder.cp_shared_kv_prefill_max_batch_requests)
self.assertEqual(adder.cp_shared_kv_prefill_max_total_extend_tokens, 128)
self.assertIsNone(adder.cp_shared_kv_prefill_max_total_cached_tokens)
# The generic budget lift uses the effective defaulted extend 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(
model_path="dummy",
enable_nsa_prefill_context_parallel=True,
nsa_prefill_cp_mode="in-seq-split",
)
)
self.mock_token_allocator.available_size.return_value = 200
adder = self.create_adder(
self.create_running_batch(),
page_size=64,
rem_input_tokens=64,
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,
)
first = self.create_prefill_req("first", extend_input_len=128)
second = self.create_prefill_req("second", extend_input_len=128)
self.assertEqual(
adder.add_one_req(first, has_chunked_req=False, truncation_align_size=None),
AddReqResult.CONTINUE,
)
self.assertEqual(
adder.add_one_req(second, has_chunked_req=False, truncation_align_size=None),
AddReqResult.NO_TOKEN,
)
self.assertEqual([req.rid for req in adder.can_run_list], ["first"])
def test_cp_prefill_chunked_req_excludes_new_requests_even_when_bs_gt1_enabled(self):
set_global_server_args_for_scheduler(
ServerArgs(
model_path="dummy",
enable_nsa_prefill_context_parallel=True,
nsa_prefill_cp_mode="in-seq-split",
)
)
self.mock_token_allocator.available_size.return_value = 10000
adder = self.create_adder(
self.create_running_batch(),
page_size=64,
rem_input_tokens=4096,
rem_chunk_tokens=256,
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,
)
chunked = self.create_prefill_req("chunked", extend_input_len=128)
normal = self.create_prefill_req("normal", extend_input_len=128)
adder.new_chunked_req = adder.add_chunked_req(chunked)
self.assertEqual([req.rid for req in adder.can_run_list], ["chunked"])
self.assertIsNone(adder.new_chunked_req)
self.assertEqual(chunked.extend_input_len, 128)
self.assertEqual(
adder.add_one_req(
normal, has_chunked_req=True, truncation_align_size=None
),
AddReqResult.OTHER,
)
self.assertEqual([req.rid for req in adder.can_run_list], ["chunked"])
def test_affinity_decision_table(self):
# Plan doc S4 (amended): the policy prevents exactly one thing — a
# COLD candidate joining a WARM-led batch. Everything else admits.
from sglang.srt.managers.schedule_policy import (
AffinityDecision,
decide_cp_prefill_affinity,
)
base = dict(
is_head=False,
head_defer_count=0,
head_age_s=0.0,
window_used=0,
)
# WARM always admits, whatever the batch looks like.
for empty, warm_led in [(True, False), (False, True), (False, False)]:
self.assertIs(
decide_cp_prefill_affinity(
is_warm=True,
batch_empty_for_affinity=empty,
batch_warm_led=warm_led,
**base,
),
AffinityDecision.ADMIT,
)
# COLD into an empty batch admits (the head always starts a batch).
self.assertIs(
decide_cp_prefill_affinity(
is_warm=False,
batch_empty_for_affinity=True,
batch_warm_led=False,
**base,
),
AffinityDecision.ADMIT,
)
# COLD into a COLD-led batch admits (small colds co-batch today).
self.assertIs(
decide_cp_prefill_affinity(
is_warm=False,
batch_empty_for_affinity=False,
batch_warm_led=False,
**base,
),
AffinityDecision.ADMIT,
)
# COLD into a WARM-led batch is skipped (within bounds).
self.assertIs(
decide_cp_prefill_affinity(
is_warm=False,
batch_empty_for_affinity=False,
batch_warm_led=True,
**base,
),
AffinityDecision.SKIP_COLD,
)
def test_affinity_anti_starvation_bounds_stop_the_scan(self):
from sglang.srt.managers.schedule_policy import (
AffinityDecision,
decide_cp_prefill_affinity,
)
cold_in_warm = dict(
is_warm=False,
batch_empty_for_affinity=False,
batch_warm_led=True,
)
# Head deferred K times -> STOP (end the warm batch cleanly; the
# cold leads the next, empty batch — never force-polluted into this
# one).
self.assertIs(
decide_cp_prefill_affinity(
**cold_in_warm,
is_head=True,
head_defer_count=3,
head_age_s=0.0,
window_used=0,
),
AffinityDecision.STOP,
)
# Head older than T -> STOP regardless of defer count.
self.assertIs(
decide_cp_prefill_affinity(
**cold_in_warm,
is_head=True,
head_defer_count=0,
head_age_s=10.0,
window_used=0,
),
AffinityDecision.STOP,
)
# Window exhausted -> STOP, head or not.
self.assertIs(
decide_cp_prefill_affinity(
**cold_in_warm,
is_head=False,
head_defer_count=0,
head_age_s=0.0,
window_used=16,
),
AffinityDecision.STOP,
)
# A NON-head cold within bounds is skipped without K/T applying.
self.assertIs(
decide_cp_prefill_affinity(
**cold_in_warm,
is_head=False,
head_defer_count=99,
head_age_s=99.0,
window_used=0,
),
AffinityDecision.SKIP_COLD,
)
def test_add_chunked_req_seeds_true_prefix_into_cp_budget(self):
# C1 (plan doc S1.1-1a): the chunk's carried prefix must count toward
# the CP cached tally so later admission gates see its footprint.
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,
rem_chunk_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,
)
chunked = self.create_prefill_req("chunked", extend_input_len=128)
chunked.prefix_indices = torch.zeros((256,), dtype=torch.int64)
adder.add_chunked_req(chunked)
self.assertEqual(adder.cp_shared_kv_prefill_total_cached_tokens, 256)
self.assertEqual(adder.cp_shared_kv_prefill_total_extend_tokens, 128)
def _mix_chunked_adder(self, *, rem_chunk_tokens, extend_cap):
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 = 100000
return self.create_adder(
self.create_running_batch(),
page_size=64,
rem_input_tokens=8192,
rem_chunk_tokens=rem_chunk_tokens,
enable_cp_shared_kv_prefill_bs_gt1=True,
cp_shared_kv_prefill_max_batch_requests=8,
cp_shared_kv_prefill_max_total_extend_tokens=extend_cap,
)
def test_cp_prefill_mix_chunked_tail_chunk_admits_following_requests(self):
# Plan doc S1: with the flag on, a TAIL chunk (extend below the chunk
# budget, page-aligned carried prefix) leaves extend-cap headroom and
# the following short-extend request co-batches with it.
with envs.SGLANG_CP_PREFILL_MIX_CHUNKED.override(True):
adder = self._mix_chunked_adder(rem_chunk_tokens=4096, extend_cap=4096)
chunked = self.create_prefill_req("chunked", extend_input_len=1280)
chunked.prefix_indices = torch.zeros((4096,), dtype=torch.int64)
self.assertIsNone(adder.add_chunked_req(chunked)) # tail chunk
normal = self.create_prefill_req("normal", extend_input_len=128)
self.assertEqual(
adder.add_one_req(
normal, has_chunked_req=True, truncation_align_size=None
),
AddReqResult.CONTINUE,
)
self.assertEqual(
[req.rid for req in adder.can_run_list], ["chunked", "normal"]
)
def test_cp_prefill_mix_chunked_full_chunk_stays_solo_by_budget(self):
# A FULL chunk consumes the whole (chunk-clamped) extend cap, so the
# first following request is rejected by the extend gate — no special
# code, the budget arithmetic ends the scan.
with envs.SGLANG_CP_PREFILL_MIX_CHUNKED.override(True):
adder = self._mix_chunked_adder(rem_chunk_tokens=256, extend_cap=4096)
chunked = self.create_prefill_req("chunked", extend_input_len=512)
chunked.prefix_indices = torch.zeros((4096,), dtype=torch.int64)
self.assertIs(adder.add_chunked_req(chunked), chunked) # truncated
self.assertEqual(chunked.extend_input_len, 256)
normal = self.create_prefill_req("normal", extend_input_len=128)
self.assertEqual(
adder.add_one_req(
normal, has_chunked_req=True, truncation_align_size=None
),
AddReqResult.OTHER,
)
self.assertEqual([req.rid for req in adder.can_run_list], ["chunked"])
def test_cp_prefill_mix_chunked_non_aligned_prefix_stays_solo(self):
# I1 guard: a chunked prefix that is not a page multiple would break
# the CP page-aligned split in a multi-request batch — keep it solo.
with envs.SGLANG_CP_PREFILL_MIX_CHUNKED.override(True):
adder = self._mix_chunked_adder(rem_chunk_tokens=4096, extend_cap=4096)
chunked = self.create_prefill_req("chunked", extend_input_len=1280)
chunked.prefix_indices = torch.zeros((100,), dtype=torch.int64)
self.assertIsNone(adder.add_chunked_req(chunked))
normal = self.create_prefill_req("normal", extend_input_len=128)
self.assertEqual(
adder.add_one_req(
normal, has_chunked_req=True, truncation_align_size=None
),
AddReqResult.OTHER,
)
self.assertEqual([req.rid for req in adder.can_run_list], ["chunked"])
def test_cp_prefill_total_cached_limit_stops_second_cached_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_total_cached_tokens=4096,
)
first = self.create_prefill_req("first", extend_input_len=64)
first.prefix_indices = torch.arange(4096, dtype=torch.int64)
first.fill_ids = list(range(4096 + 64))
second = self.create_prefill_req("second", extend_input_len=64)
second.prefix_indices = torch.arange(4096, dtype=torch.int64)
second.fill_ids = list(range(4096 + 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.assertEqual(adder.cp_shared_kv_prefill_total_cached_tokens, 4096)
def test_cp_prefill_total_cached_limit_allows_single_oversized_cached_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 = 20000
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,
)
oversized = self.create_prefill_req("oversized", extend_input_len=64)
oversized.prefix_indices = torch.arange(8192, dtype=torch.int64)
oversized.fill_ids = list(range(8192 + 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"])
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()