Files
sglang/test/registered/unit/managers/test_prefill_adder.py
laoyao0822 d14c02b0dc Count evictable device cache when gating HiCache load-back
HiCache host hits can be skipped before load-back when the quota gate only counts immediately free KV allocator space. Under CP shared-KV pressure most reusable capacity may be represented as evictable radix-cache leaves, so the gate can incorrectly reject a host hit and leave prefill with cached-token zero despite host residency. Count device evictable cache in the quota estimate while leaving actual owner-lane allocation and eviction checks in the load path.

Constraint: CP HiCache load-back still has to respect owner-lane allocation and allocator eviction semantics.

Rejected: Force load-back regardless of quota | would bypass the scheduler pressure signal and increase OOM risk.

Rejected: Treat cache-hit zero as a transfer issue | logs showed host hits were found but skipped by quota before transfer.

Confidence: medium

Scope-risk: moderate

Directive: Do not remove evictable cache from load-back capacity accounting without checking CP HiCache host-hit behavior under device pressure.

Tested: git diff --check

Tested: remote g0034 container pytest -q test/registered/unit/managers/test_prefill_adder.py test/registered/unit/managers/test_hicache_controller_cp.py test/registered/unit/mem_cache/test_cp_hicache_metadata.py test/registered/unit/mem_cache/test_alloc_pages_with_owners.py (90 passed, 3 warnings)

Not-tested: Full ETE GLM5 CP+HiCache+EAGLE pressure run after this quota change

Co-authored-by: OmX <omx@oh-my-codex.dev>
2026-05-27 02:41:45 +08:00

556 lines
21 KiB
Python

import sys
import types
import unittest
from types import SimpleNamespace
from unittest.mock import MagicMock
import torch
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.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_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 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)
if __name__ == "__main__":
unittest.main()