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