专题 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>
1206 lines
47 KiB
Python
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()
|