Make CP HiCache residency owner-lane deterministic

CP shared KV cannot treat capacity as a scalar token count: cache-hit load-back and fresh extend allocation both have to preserve the logical page owner pattern or later direct writes, HiCache reload, and prefix materialization can read the wrong lane. This change moves the critical paths to owner-lane plans, makes owner-lane exhaustion recoverable during prefill scheduling, and routes shared-KV prefix prefetch through prefetch-stream-safe KV getters so HiCache layer-load waits do not attach to the forward stream.

Constraint: CP shared KV correctness depends on page owner lane preservation across allocation, backup, load, eviction, and prefix materialization.

Constraint: Avoid adding CP/global collectives for capacity agreement; derive capacity from deterministic local owner-lane state.

Rejected: Keep SGLANG_DISABLE_TAI_OWNER_SELECT fallback | legacy allocation can silently break owner-lane invariants.

Rejected: Scalar total-token eviction for CP HiCache load-back | total capacity can be sufficient while the required owner lane is exhausted.

Confidence: medium

Scope-risk: broad

Directive: Do not reintroduce silent legacy fallback in owner-lane paths; unexpected owner-lane failure must be warning-level fail-closed or recoverable capacity wait.

Tested: Remote g0034 container PYTHONPATH=python python -m pytest test/registered/unit/mem_cache/test_alloc_pages_with_owners.py test/registered/unit/mem_cache/test_cp_shared_kv_layout.py test/registered/unit/mem_cache/test_cp_hicache_load_back_owner_lanes.py test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py -q -> 95 passed.

Tested: Local py_compile for modified runtime/cache/scheduler modules.

Not-tested: Full CUDA ETE performance trace for cache-hit overlap and MTP accept-rate impact.

Co-authored-by: OmX <omx@oh-my-codex.dev>
This commit is contained in:
laoyao0822
2026-05-28 05:54:23 +08:00
co-authored by OmX
parent 2c94b8de23
commit ff33446787
12 changed files with 2913 additions and 136 deletions
@@ -0,0 +1,318 @@
import functools
import json
import sys
import types
import unittest
import torch
if "pybase64" not in sys.modules:
pybase64_stub = types.ModuleType("pybase64")
pybase64_stub.b64encode = lambda *args, **kwargs: b""
pybase64_stub.b64decode = lambda *args, **kwargs: b""
sys.modules["pybase64"] = pybase64_stub
if "orjson" not in sys.modules:
orjson_stub = types.ModuleType("orjson")
orjson_stub.loads = lambda data, *args, **kwargs: json.loads(
data.decode() if isinstance(data, (bytes, bytearray)) else data
)
orjson_stub.dumps = lambda obj, *args, **kwargs: json.dumps(obj).encode()
sys.modules["orjson"] = orjson_stub
sgl_kernel_stub = sys.modules.setdefault("sgl_kernel", types.ModuleType("sgl_kernel"))
sgl_kernel_stub.__file__ = getattr(sgl_kernel_stub, "__file__", "sgl_kernel_stub.py")
sgl_kernel_stub.__path__ = getattr(sgl_kernel_stub, "__path__", [])
if not hasattr(sgl_kernel_stub, "__getattr__"):
def _sgl_kernel_getattr(name):
if name.startswith("__"):
raise AttributeError(name)
fn = lambda *args, **kwargs: None
setattr(sgl_kernel_stub, name, fn)
return fn
sgl_kernel_stub.__getattr__ = _sgl_kernel_getattr
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(sgl_kernel_stub, _name):
setattr(sgl_kernel_stub, _name, lambda *args, **kwargs: None)
quantization_stub = sys.modules.setdefault(
"sgl_kernel.quantization", 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",
):
if not hasattr(quantization_stub, _name):
setattr(quantization_stub, _name, lambda *args, **kwargs: None)
kvcacheio_stub = sys.modules.setdefault(
"sgl_kernel.kvcacheio", types.ModuleType("sgl_kernel.kvcacheio")
)
for _name in (
"transfer_kv_all_layer",
"transfer_kv_all_layer_direct_lf_pf",
"transfer_kv_all_layer_lf_pf",
"transfer_kv_all_layer_lf_ph",
"transfer_kv_all_layer_mla",
"transfer_kv_all_layer_mla_lf_pf",
"transfer_kv_direct",
"transfer_kv_per_layer",
"transfer_kv_per_layer_direct_pf_lf",
"transfer_kv_per_layer_mla",
"transfer_kv_per_layer_mla_pf_lf",
"transfer_kv_per_layer_pf_lf",
"transfer_kv_per_layer_ph_lf",
):
if not hasattr(kvcacheio_stub, _name):
setattr(kvcacheio_stub, _name, lambda *args, **kwargs: None)
_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.mem_cache.allocator import CPSharedPagedTokenToKVPoolAllocator
from sglang.srt.mem_cache.hiradix_cache import (
CpHiCacheNodeMetadata,
HiRadixCache,
)
from sglang.srt.mem_cache.radix_cache import RadixKey, TreeNode, get_child_key
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=1, suite="stage-a-test-cpu")
class _FakeLayout:
cp_size = 4
cp_rank = 0
class _FakeEvictionStrategy:
def get_priority(self, node):
return getattr(node, "priority", 0)
class _FakeController:
write_policy = "write_through"
has_draft_hicache = False
cp_shared_kv_layout = _FakeLayout()
def __init__(self, allocator):
self.allocator = allocator
self.load_calls = 0
self.evicted_device_indices = []
self.ack_load_queue = []
self.force_load_none = False
def load_cp(self, nodes_to_load, node_id=-1):
page_owners = []
for node in nodes_to_load:
page_owners.extend(node.cp_hicache.page_owners.tolist())
_, _, deficits = self.allocator.compute_owner_lane_stats(page_owners)
if any(deficits):
raise AssertionError(
f"load_cp called before owner-lane deficits were evicted: {deficits}"
)
self.load_calls += 1
if self.force_load_none:
return None
return self.allocator.alloc_pages_with_owners(page_owners)
def evict_device(self, indices):
self.evicted_device_indices.append(indices.clone())
self.allocator.free(indices)
return int(indices.numel())
def _make_allocator(page_size=4, cp_size=4):
allocator = CPSharedPagedTokenToKVPoolAllocator(
logical_size=page_size * 16,
physical_size=page_size * 4,
page_size=page_size,
dtype=torch.bfloat16,
device="cpu",
kvcache=None,
need_sort=False,
cp_size=cp_size,
cp_rank=0,
)
allocator.release_pages = torch.empty((0,), dtype=torch.int64)
return allocator
def _metadata(page_owners, page_size=4):
logical_len = len(page_owners) * page_size
owned_positions = []
for page_idx, owner in enumerate(page_owners):
if owner == 0:
owned_positions.extend(
range(page_idx * page_size, (page_idx + 1) * page_size)
)
return CpHiCacheNodeMetadata(
logical_len=logical_len,
owned_positions=torch.tensor(owned_positions, dtype=torch.int64),
host_indices=torch.arange(len(owned_positions), dtype=torch.int64),
page_owners=torch.tensor(page_owners, dtype=torch.int8),
page_size=page_size,
)
def _make_node(node_id, key_start, page_owners, *, value=None, priority=0):
page_size = 4
node = TreeNode(id=node_id, priority=priority)
node.key = RadixKey(token_ids=list(range(key_start, key_start + len(page_owners) * page_size)))
node.value = value
node.host_len = len(page_owners) * page_size
node.cp_hicache = _metadata(page_owners, page_size=page_size)
return node
def _attach_child(cache, parent, child):
child.parent = parent
parent.children[cache.get_child_key_fn(child.key)] = child
def _make_cache(allocator):
cache = HiRadixCache.__new__(HiRadixCache)
cache.disable = False
cache._uses_cp_hicache = True
cache.token_to_kv_pool_allocator = allocator
cache.page_size = allocator.page_size
cache.cache_controller = _FakeController(allocator)
cache.root_node = TreeNode(id=0, priority=-999)
cache.root_node.key = RadixKey(token_ids=[])
cache.root_node.value = []
cache.root_node.lock_ref = 1
cache.evictable_leaves = set()
cache.evictable_host_leaves = set()
cache.evictable_size_ = 0
cache.protected_size_ = 0
cache.ongoing_load_back = {}
cache.ongoing_write_through = {}
cache.pending_host_backups = {}
cache.load_back_threshold = 0
cache.metrics_collector = None
cache.eviction_strategy = _FakeEvictionStrategy()
cache.get_child_key_fn = functools.partial(get_child_key, page_size=allocator.page_size)
cache.enable_kv_cache_events = False
return cache
class TestCpHiCacheLoadBackOwnerLanes(CustomTestCase):
def test_load_back_plan_reports_owner_lane_vectors(self):
allocator = _make_allocator()
allocator.free_pages = torch.tensor([1, 2], dtype=torch.int64)
allocator.release_pages = torch.tensor([4], dtype=torch.int64)
cache = _make_cache(allocator)
node = _make_node(10, 100, [0, 1, 1, 3])
plan = cache._build_cp_load_back_plan([node], node_id=node.id)
self.assertEqual(plan.page_owners, [0, 1, 1, 3])
self.assertEqual(plan.required_by_owner, [1, 2, 0, 1])
self.assertEqual(plan.available_by_owner, [1, 1, 0, 1])
self.assertEqual(plan.deficit_by_owner, [0, 1, 0, 0])
self.assertEqual(plan.host_hit_len, 16)
def test_load_back_plan_fails_closed_without_cp_metadata(self):
allocator = _make_allocator()
cache = _make_cache(allocator)
node = TreeNode(id=11)
node.host_len = 4
node.cp_hicache = None
with self.assertRaisesRegex(RuntimeError, "missing cp_hicache metadata"):
cache._build_cp_load_back_plan([node], node_id=node.id)
def test_load_back_evicts_owner_lane_deficit_before_allocating(self):
allocator = _make_allocator()
# Only owner lane 1 is initially available. The load-back target needs
# lanes [0, 1], so calling load_cp before targeted eviction is a bug.
allocator.free_pages = torch.tensor([2], dtype=torch.int64)
cache = _make_cache(allocator)
victim = _make_node(
20,
200,
[0],
value=torch.tensor([4, 5, 6, 7], dtype=torch.int64),
priority=0,
)
_attach_child(cache, cache.root_node, victim)
cache.evictable_leaves.add(victim)
cache.evictable_size_ = len(victim.key)
target = _make_node(21, 300, [0, 1], value=None, priority=10)
_attach_child(cache, cache.root_node, target)
loaded = cache.load_back(target, mem_quota=100)
self.assertIsNotNone(loaded)
self.assertEqual(cache.cache_controller.load_calls, 1)
self.assertEqual(len(cache.cache_controller.evicted_device_indices), 1)
self.assertEqual(cache.cache_controller.evicted_device_indices[0].tolist(), [4, 5, 6, 7])
self.assertEqual(target.value.tolist(), loaded.tolist())
self.assertIn(target.id, cache.ongoing_load_back)
def test_load_back_failure_leaves_node_unassigned_and_unlocked(self):
allocator = _make_allocator()
allocator.free_pages = torch.tensor([1, 2], dtype=torch.int64)
cache = _make_cache(allocator)
cache.cache_controller.force_load_none = True
target = _make_node(30, 400, [0, 1], value=None)
_attach_child(cache, cache.root_node, target)
loaded = cache.load_back(target, mem_quota=100)
self.assertIsNone(loaded)
self.assertIsNone(target.value)
self.assertEqual(target.lock_ref, 0)
self.assertNotIn(target.id, cache.ongoing_load_back)
def test_load_back_lock_is_released_only_by_loading_ack(self):
class ReadyEvent:
def query(self):
return True
allocator = _make_allocator()
allocator.free_pages = torch.tensor([1, 2], dtype=torch.int64)
cache = _make_cache(allocator)
target = _make_node(31, 500, [0, 1], value=None)
_attach_child(cache, cache.root_node, target)
loaded = cache.load_back(target, mem_quota=100)
self.assertIsNotNone(loaded)
self.assertEqual(target.lock_ref, 1)
self.assertIn(target.id, cache.ongoing_load_back)
cache.cache_controller.ack_load_queue.append((None, ReadyEvent(), [target.id]))
cache.loading_check()
self.assertEqual(target.lock_ref, 0)
self.assertNotIn(target.id, cache.ongoing_load_back)
self.assertEqual(cache.cache_controller.ack_load_queue, [])
if __name__ == "__main__":
unittest.main()
@@ -4,6 +4,18 @@ from unittest.mock import patch
import numpy as np
import torch
_sgl_kernel_lib = torch.library.Library("sgl_kernel", "FRAGMENT")
try:
_sgl_kernel_lib.define(
"moe_fused_gate(Tensor input_tensor, Tensor? bias, int num_expert_group, "
"int topk_group, int topk, int num_fused_shared_experts, "
"float routed_scaling_factor, bool apply_routed_scaling_factor_on_output) "
"-> (Tensor, Tensor)"
)
except RuntimeError as exc:
if "already" not in str(exc).lower() and "duplicate" not in str(exc).lower():
raise
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -435,6 +447,75 @@ class TestCPSharedPagedAllocator(CustomTestCase):
self.assertEqual(allocator.calls, 1)
self.assertEqual(out.numel(), page_size * 8)
def test_compute_owner_alloc_skips_aggregate_evict_before_owner_attempt(self):
from types import SimpleNamespace
from sglang.srt.mem_cache import common
page_size = 64
class FakeAllocator:
def __init__(self):
self.page_size = page_size
self.cp_size = 4
self.calls = 0
def available_size(self):
# Aggregate capacity looks insufficient, but the owner-aware
# allocator can satisfy the request from the right lanes.
# Aggregate eviction before this attempt would evict unrelated
# lanes and destroy cache locality.
return 0
def alloc_extend_compute_owner(
self,
_prefix_lens,
_prefix_lens_cpu,
_seq_lens,
_seq_lens_cpu,
_last_loc,
extend_num_tokens,
_page_compute_owners,
):
self.calls += 1
return torch.arange(extend_num_tokens, dtype=torch.int64)
def alloc_extend(self, *_args, **_kwargs):
raise AssertionError("legacy allocation should not be used")
class FakeTreeCache:
def __init__(self, allocator):
self.token_to_kv_pool_allocator = allocator
def is_chunk_cache(self):
return False
def evict(self, _params):
raise AssertionError(
"aggregate eviction should not run before owner-aware allocation"
)
allocator = FakeAllocator()
server_args = SimpleNamespace(
enable_nsa_prefill_cp_shared_kv=True,
enable_nsa_prefill_context_parallel=True,
nsa_prefill_cp_mode="in-seq-split",
)
with patch.object(common, "get_global_server_args", return_value=server_args):
out = common.alloc_paged_token_slots_extend(
tree_cache=FakeTreeCache(allocator),
prefix_lens=torch.tensor([0], dtype=torch.int64),
prefix_lens_cpu=torch.tensor([0], dtype=torch.int64),
seq_lens=torch.tensor([page_size * 8], dtype=torch.int64),
seq_lens_cpu=torch.tensor([page_size * 8], dtype=torch.int64),
last_loc=torch.tensor([-1], dtype=torch.int64),
extend_num_tokens=page_size * 8,
)
self.assertEqual(allocator.calls, 1)
self.assertEqual(out.numel(), page_size * 8)
def test_compute_owner_lane_eviction_recovers_exhausted_owner_lane(self):
from sglang.srt.mem_cache.allocator import CPSharedPagedTokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import EvictResult
@@ -508,6 +589,253 @@ class TestCPSharedPagedAllocator(CustomTestCase):
)
self.assertIsNotNone(locs)
def test_compute_owner_lane_eviction_passes_deficits_to_tree_cache(self):
from sglang.srt.mem_cache.allocator import CPSharedPagedTokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import EvictResult
from sglang.srt.mem_cache.common import _evict_for_compute_owner_lanes
page_size = 64
cp_size = 4
allocator = CPSharedPagedTokenToKVPoolAllocator(
logical_size=page_size * 16,
physical_size=page_size * 4,
page_size=page_size,
dtype=torch.bfloat16,
device="cpu",
kvcache=None,
need_sort=False,
cp_size=cp_size,
cp_rank=0,
)
allocator.free_pages = torch.empty((0,), dtype=torch.int64)
allocator.release_pages = torch.empty((0,), dtype=torch.int64)
class FakeTreeCache:
def __init__(self):
self.owner_deficits = []
def is_chunk_cache(self):
return False
def evictable_size(self):
return page_size * cp_size
def evict(self, params):
deficits = getattr(params, "owner_lane_deficits", None)
self.owner_deficits.append(deficits)
if deficits == [1, 0, 0, 0]:
allocator.free(torch.tensor([page_size], dtype=torch.int64))
else:
allocator.free(torch.tensor([page_size * 2], dtype=torch.int64))
return EvictResult(num_tokens_evicted=page_size)
tree_cache = FakeTreeCache()
_evict_for_compute_owner_lanes(
tree_cache=tree_cache,
allocator=allocator,
page_compute_owners=[0],
)
self.assertEqual(tree_cache.owner_deficits[0], [1, 0, 0, 0])
locs = allocator.alloc_extend_compute_owner(
prefix_lens=torch.tensor([0], dtype=torch.int64),
prefix_lens_cpu=torch.tensor([0], dtype=torch.int64),
seq_lens=torch.tensor([page_size], dtype=torch.int64),
seq_lens_cpu=torch.tensor([page_size], dtype=torch.int64),
last_loc=torch.tensor([-1], dtype=torch.int64),
extend_num_tokens=page_size,
page_compute_owners=[0],
)
self.assertIsNotNone(locs)
def test_compute_owner_capacity_wait_reports_owner_lane_deficits(self):
from types import SimpleNamespace
from sglang.srt.mem_cache import common
from sglang.srt.mem_cache.base_prefix_cache import EvictResult
page_size = 64
class FakeAllocator:
def __init__(self):
self.page_size = page_size
self.cp_size = 4
def available_size(self):
return self.page_size * 4
def allocator_state_str(self):
return "allocator_state_for_test"
def alloc_extend_compute_owner(self, *_args, **_kwargs):
return None
def compute_owner_lane_stats(self, _page_compute_owners):
return [2, 2, 2, 2], [2, 2, 2, 0], [0, 0, 0, 2]
def alloc_extend(self, *_args, **_kwargs):
raise AssertionError("legacy allocation should not be used")
class FakeTreeCache:
def __init__(self):
self.token_to_kv_pool_allocator = FakeAllocator()
self.evict_params = []
def is_chunk_cache(self):
return False
def evictable_size(self):
return page_size * 8
def evict(self, params):
self.evict_params.append(params)
return EvictResult(num_tokens_evicted=0)
def pretty_print(self):
raise AssertionError("recoverable capacity wait should not dump tree")
server_args = SimpleNamespace(
enable_nsa_prefill_cp_shared_kv=True,
enable_nsa_prefill_context_parallel=True,
nsa_prefill_cp_mode="in-seq-split",
)
tree_cache = FakeTreeCache()
with patch.object(common, "get_global_server_args", return_value=server_args):
with self.assertRaises(common.KVCapacityWaitError) as cm:
common.alloc_paged_token_slots_extend(
tree_cache=tree_cache,
prefix_lens=torch.tensor([0], dtype=torch.int64),
prefix_lens_cpu=torch.tensor([0], dtype=torch.int64),
seq_lens=torch.tensor([page_size * 8], dtype=torch.int64),
seq_lens_cpu=torch.tensor([page_size * 8], dtype=torch.int64),
last_loc=torch.tensor([-1], dtype=torch.int64),
extend_num_tokens=page_size * 8,
)
err = cm.exception
self.assertEqual(err.required_by_owner, [2, 2, 2, 2])
self.assertEqual(err.available_by_owner, [2, 2, 2, 0])
self.assertEqual(err.deficit_by_owner, [0, 0, 0, 2])
self.assertIn("owner_lane_exhausted", str(err))
self.assertIn("deficit_by_owner=[0, 0, 0, 2]", str(err))
self.assertIn(
[0, 0, 0, 2],
[params.owner_lane_deficits for params in tree_cache.evict_params],
)
def test_alloc_for_extend_releases_req_slots_on_recoverable_capacity_wait(self):
from types import SimpleNamespace
from sglang.srt.mem_cache import common
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
req_to_token_pool = ReqToTokenPool(
size=1,
max_context_len=128,
device="cpu",
enable_memory_saver=False,
)
req = SimpleNamespace(
req_pool_idx=None,
is_chunked=0,
kv_committed_len=0,
prefix_indices=torch.empty((0,), dtype=torch.int64),
)
batch = SimpleNamespace(
maybe_evict_swa=lambda: None,
reqs=[req],
prefix_lens=[0],
extend_lens=[64],
device="cpu",
req_to_token_pool=req_to_token_pool,
tree_cache=SimpleNamespace(page_size=64),
seq_lens=torch.tensor([64], dtype=torch.int64),
seq_lens_cpu=torch.tensor([64], dtype=torch.int64),
extend_num_tokens=64,
)
wait_error = common.KVCapacityWaitError(
"owner_lane_exhausted for test",
required_by_owner=[1, 0],
available_by_owner=[0, 0],
deficit_by_owner=[1, 0],
)
with patch.object(
common,
"alloc_paged_token_slots_extend",
side_effect=wait_error,
):
with self.assertRaises(common.KVCapacityWaitError):
common.alloc_for_extend(batch)
self.assertIsNone(req.req_pool_idx)
self.assertEqual(req_to_token_pool.available_size(), 1)
def test_scheduler_capacity_wait_rollback_releases_adder_locks(self):
from types import SimpleNamespace
from sglang.srt.managers.scheduler import Scheduler
class FakeTreeCache:
def __init__(self):
self.dec_calls = []
def supports_swa(self):
return False
def is_tree_cache(self):
return True
def dec_lock_ref(self, node, params=None):
self.dec_calls.append((node, params))
scheduler = object.__new__(Scheduler)
scheduler.tree_cache = FakeTreeCache()
req = SimpleNamespace(last_node="scheduled-node", swa_uuid_for_lock=None)
chunked_req = SimpleNamespace(last_node="chunked-node", swa_uuid_for_lock=None)
scheduler._release_prefill_adder_locks(
[req, chunked_req],
skip_req=chunked_req,
)
self.assertEqual(
[(node, params) for node, params in scheduler.tree_cache.dec_calls],
[("scheduled-node", None)],
)
def test_scheduler_capacity_wait_rollback_releases_swa_lock_uuid(self):
from types import SimpleNamespace
from sglang.srt.managers.scheduler import Scheduler
class FakeTreeCache:
def __init__(self):
self.dec_calls = []
def supports_swa(self):
return True
def is_tree_cache(self):
return True
def dec_lock_ref(self, node, params=None):
self.dec_calls.append((node, params))
scheduler = object.__new__(Scheduler)
scheduler.tree_cache = FakeTreeCache()
req = SimpleNamespace(last_node="scheduled-node", swa_uuid_for_lock=17)
scheduler._release_prefill_adder_locks([req])
self.assertEqual(scheduler.tree_cache.dec_calls[0][0], "scheduled-node")
self.assertEqual(
scheduler.tree_cache.dec_calls[0][1].swa_uuid_for_lock,
17,
)
self.assertIsNone(req.swa_uuid_for_lock)
if __name__ == "__main__":
unittest.main()
@@ -1,9 +1,30 @@
import unittest
import sys
import types
from types import SimpleNamespace
from unittest.mock import patch
import torch
_sgl_kernel_lib = torch.library.Library("sgl_kernel", "FRAGMENT")
try:
_sgl_kernel_lib.define(
"moe_fused_gate(Tensor input_tensor, Tensor? bias, int num_expert_group, "
"int topk_group, int topk, int num_fused_shared_experts, "
"float routed_scaling_factor, bool apply_routed_scaling_factor_on_output) "
"-> (Tensor, Tensor)"
)
except RuntimeError as exc:
if "already" not in str(exc).lower() and "duplicate" not in str(exc).lower():
raise
flash_attn_stub = sys.modules.setdefault(
"sgl_kernel.flash_attn", types.ModuleType("sgl_kernel.flash_attn")
)
for _name in ("flash_attn_varlen_func", "flash_attn_with_kvcache"):
if not hasattr(flash_attn_stub, _name):
setattr(flash_attn_stub, _name, lambda *args, **kwargs: None)
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=1, suite="stage-a-test-cpu")
@@ -14,7 +35,7 @@ def _identity_all_reduce(buffer, *args, **kwargs):
class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
def test_mla_prefetch_materializes_on_current_stream_and_defers_reduce_to_launch(
def test_mla_prefetch_materializes_and_reduces_on_prefetch_stream(
self,
):
from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch
@@ -48,8 +69,13 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
def __init__(self):
self.kv_cache = torch.zeros((32, 1, 2), dtype=torch.float32)
self.prefetch_getter_streams = []
def get_key_buffer(self, layer_id):
raise AssertionError("prefetch must not call blocking get_key_buffer")
def get_key_buffer_for_prefetch(self, layer_id, stream):
self.prefetch_getter_streams.append((layer_id, stream.name))
return self.kv_cache
active_stream = ["current"]
@@ -89,23 +115,28 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
"_all_reduce_materialized_buffer_async",
side_effect=record_reduce,
):
pool = FakePool()
prefetcher.start_next_layer_prefix(
next_layer_id=1,
token_to_kv_pool=FakePool(),
token_to_kv_pool=pool,
)
self.assertEqual(calls, [("materialize", "current")])
self.assertEqual(prefetch_stream.waited, [])
self.assertEqual(pool.prefetch_getter_streams, [(1, "prefetch")])
self.assertEqual(
calls,
[("materialize", "prefetch"), ("reduce", "prefetch", "prefetch")],
)
self.assertEqual(prefetch_stream.waited, ["current"])
prefetcher.launch_pending_reduce()
self.assertEqual(
calls,
[("materialize", "current"), ("reduce", "prefetch", "prefetch")],
[("materialize", "prefetch"), ("reduce", "prefetch", "prefetch")],
)
self.assertEqual(prefetch_stream.waited, ["current"])
def test_index_prefetch_materializes_on_current_stream_and_defers_reduce_to_launch(
def test_index_prefetch_materializes_and_reduces_on_prefetch_stream(
self,
):
from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch
@@ -139,8 +170,15 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
def __init__(self):
self.page_buffer = torch.zeros((16, 3), dtype=torch.uint8)
self.prefetch_getter_streams = []
def get_index_k_with_scale_buffer(self, layer_id):
raise AssertionError(
"prefetch must not call blocking get_index_k_with_scale_buffer"
)
def get_index_k_with_scale_buffer_for_prefetch(self, layer_id, stream):
self.prefetch_getter_streams.append((layer_id, stream.name))
return self.page_buffer
active_stream = ["current"]
@@ -179,22 +217,55 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
"_all_reduce_materialized_buffer_async",
side_effect=record_reduce,
):
pool = FakePool()
prefetcher.start_next_layer_prefix(
next_layer_id=1,
token_to_kv_pool=FakePool(),
token_to_kv_pool=pool,
)
self.assertEqual(calls, [("materialize", "current")])
self.assertEqual(prefetch_stream.waited, [])
self.assertEqual(pool.prefetch_getter_streams, [(1, "prefetch")])
self.assertEqual(
calls,
[("materialize", "prefetch"), ("reduce", "prefetch", "prefetch")],
)
self.assertEqual(prefetch_stream.waited, ["current"])
prefetcher.launch_pending_reduce()
self.assertEqual(
calls,
[("materialize", "current"), ("reduce", "prefetch", "prefetch")],
[("materialize", "prefetch"), ("reduce", "prefetch", "prefetch")],
)
self.assertEqual(prefetch_stream.waited, ["current"])
def test_mla_pool_prefetch_getter_orders_layer_transfer_on_prefetch_stream(self):
from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool
class FakeCounter:
def __init__(self):
self.calls = []
def wait_until(self, threshold):
raise AssertionError("prefetch getter must not wait on current stream")
def wait_until_on_stream(self, threshold, stream):
self.calls.append((threshold, stream.name))
class FakeStream:
name = "prefetch"
pool = MLATokenToKVPool.__new__(MLATokenToKVPool)
pool.start_layer = 2
pool.layer_transfer_counter = FakeCounter()
pool.store_dtype = torch.float32
pool.dtype = torch.float32
pool.kv_buffer = [torch.ones((4, 1), dtype=torch.float32) for _ in range(3)]
out = pool.get_key_buffer_for_prefetch(3, FakeStream())
self.assertIs(out, pool.kv_buffer[1])
self.assertEqual(pool.layer_transfer_counter.calls, [(1, "prefetch")])
def test_all_reduce_uses_group_fast_path_for_float_buffers(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
@@ -1782,12 +1853,12 @@ class TestCpSharedKVTaiMaterializeIntegration(unittest.TestCase):
self.assertIs(dense_buffer, fallback_buffer)
self.assertIs(dense_pages, fallback_pages)
logger.info.assert_called_once()
logger.warning.assert_called_once()
self.assertIn(
"CP shared KV index prefetch fallback",
logger.info.call_args.args[0],
"[CP_SHARED_KV_FALLBACK][index_prefetch]",
logger.warning.call_args.args[0],
)
self.assertIn("consume_miss", logger.info.call_args.args[1])
self.assertIn("consume_miss", logger.warning.call_args.args[1])
def test_index_prefetch_first_layer_miss_does_not_log_fallback(self):
from sglang.srt.layers.attention.nsa import nsa_indexer
@@ -1835,7 +1906,7 @@ class TestCpSharedKVTaiMaterializeIntegration(unittest.TestCase):
logical_page_table=logical_pages,
)
logger.info.assert_not_called()
logger.warning.assert_not_called()
def test_index_prefetch_create_skip_logs_fallback_when_enabled(self):
from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch
@@ -1847,7 +1918,7 @@ class TestCpSharedKVTaiMaterializeIntegration(unittest.TestCase):
), patch.object(
prefetch.torch.cuda, "is_available", return_value=False
), patch.object(
prefetch.logger, "info"
prefetch.logger, "warning"
) as logger:
result = prefetch.CpSharedKVIndexPrefetcher.maybe_create(
forward_batch=SimpleNamespace(),
@@ -1858,7 +1929,7 @@ class TestCpSharedKVTaiMaterializeIntegration(unittest.TestCase):
self.assertIsNone(result)
logger.assert_called_once()
self.assertIn(
"CP shared KV index prefetch fallback",
"[CP_SHARED_KV_FALLBACK][index_prefetch]",
logger.call_args.args[0],
)
self.assertIn("cuda_unavailable_or_stream_capturing", logger.call_args.args[1])