Keep current reuse independent of async CP prefetch
Async MLA/index prefetch is a scheduling optimization, not the correctness contract for target current reuse. Tiny cache-hit suffixes can skip async prefetcher creation while target partial-current reuse still composes page-slot prefix materialization with current KV rows synchronously. CP HiCache radix/device accounting now treats retained valid-tail pages as physical page spans so allocator state stays consistent when logical cache keys are shorter than the retained page. Constraint: CP shared KV ownership and HiCache residency are page-granular while request-visible cache lengths remain valid-token lengths. Constraint: Async prefetch can hang or regress on large-prefix tiny-extend traffic and must not be required for current reuse. Rejected: Treat missing prefetcher as fail-fast for target partial-current reuse | disabled useful current reuse and broke tiny-prefix/tiny-suffix traffic. Rejected: Keep async prefetcher object with synchronous consume mode | conflates prefetch object existence with current-layer correctness and hides fallback semantics. Confidence: medium Scope-risk: moderate Directive: Do not make current-only or target partial-current reuse depend on MLA/index prefetcher creation; prefetcher objects mean async next-layer work exists. Tested: Remote g0034 container py_compile for touched modules. Tested: Remote g0034 PYTHONPATH=python python -m pytest -q test/registered/unit/mem_cache/test_cp_shared_kv_runtime.py -> 73 passed. Tested: Remote g0034 PYTHONPATH=python python -m pytest -q test/registered/unit/mem_cache/test_cp_hicache_metadata.py -> 90 passed. Not-tested: Latest full ETE traffic run with GLM-5.1 CP HiCache after this commit. Not-tested: CUDA kernel-level performance impact of synchronous no-prefetch partial-current compose.
This commit is contained in:
@@ -666,7 +666,133 @@ class FakeTokenAllocator:
|
||||
return "FakeTokenAllocator"
|
||||
|
||||
|
||||
class RecordingTokenAllocator:
|
||||
def __init__(self):
|
||||
self.freed = []
|
||||
|
||||
def free(self, free_index):
|
||||
if free_index.numel() > 0:
|
||||
self.freed.append(free_index.detach().cpu().tolist())
|
||||
|
||||
|
||||
class RecordingDeviceAllocator:
|
||||
def __init__(self):
|
||||
self.freed = []
|
||||
|
||||
def free(self, free_index):
|
||||
self.freed.append(free_index.detach().cpu().tolist())
|
||||
|
||||
|
||||
class TestHiRadixCacheCPBackup(CustomTestCase):
|
||||
def _minimal_cp_hiradix_cache(self, *, page_size=64):
|
||||
cache = HiRadixCache.__new__(HiRadixCache)
|
||||
cache.disable = False
|
||||
cache._uses_cp_hicache = True
|
||||
cache.page_size = page_size
|
||||
cache.is_eagle = False
|
||||
cache.enable_storage = False
|
||||
cache.enable_kv_cache_events = False
|
||||
cache.evictable_size_ = 0
|
||||
cache.protected_size_ = 0
|
||||
cache.evictable_leaves = set()
|
||||
cache.evictable_host_leaves = set()
|
||||
cache.get_child_key_fn = lambda key: tuple(key.token_ids[:page_size])
|
||||
cache.key_match_fn = lambda key0, key1: _key_match_paged(
|
||||
key0, key1, page_size
|
||||
)
|
||||
cache.cache_controller = types.SimpleNamespace(write_policy="write_back")
|
||||
cache._record_store_event = lambda node: None
|
||||
cache._record_remove_event = lambda node: None
|
||||
|
||||
root = TreeNode()
|
||||
root.key = RadixKey(token_ids=[], extra_key=None)
|
||||
root.value = torch.empty((0,), dtype=torch.int64)
|
||||
root.children = {}
|
||||
root.parent = None
|
||||
cache.root_node = root
|
||||
return cache
|
||||
|
||||
def test_cp_eagle_finished_cache_preserves_retained_tail_page(self):
|
||||
cache = HiRadixCache.__new__(HiRadixCache)
|
||||
cache.disable_finished_insert = False
|
||||
cache.disable = False
|
||||
cache.is_eagle = True
|
||||
cache.page_size = 64
|
||||
cache._uses_cp_hicache = True
|
||||
cache.req_to_token_pool = types.SimpleNamespace(
|
||||
req_to_token=torch.arange(128, dtype=torch.int64).view(1, 128)
|
||||
)
|
||||
allocator = RecordingTokenAllocator()
|
||||
cache.token_to_kv_pool_allocator = allocator
|
||||
cache.insert = lambda params: types.SimpleNamespace(prefix_len=0)
|
||||
cache.dec_lock_ref = lambda node: None
|
||||
|
||||
req = types.SimpleNamespace(
|
||||
origin_input_ids=[10, 11, 12, 13],
|
||||
output_ids=[],
|
||||
req_pool_idx=0,
|
||||
extra_key=None,
|
||||
cache_protected_len=0,
|
||||
last_node=object(),
|
||||
cp_hicache_prepared_backup=None,
|
||||
pop_committed_kv_cache=lambda: 4,
|
||||
)
|
||||
|
||||
cache.cache_finished_req(req)
|
||||
|
||||
self.assertEqual(allocator.freed, [])
|
||||
|
||||
def test_cp_valid_tail_device_accounting_uses_physical_page_span(self):
|
||||
cache = self._minimal_cp_hiradix_cache(page_size=64)
|
||||
|
||||
result = cache.insert(
|
||||
InsertParams(
|
||||
key=RadixKey(token_ids=[1, 2, 3], extra_key=None),
|
||||
value=torch.arange(3, dtype=torch.int64),
|
||||
)
|
||||
)
|
||||
|
||||
self.assertEqual(result.prefix_len, 0)
|
||||
self.assertEqual(cache.evictable_size_, 64)
|
||||
node = next(iter(cache.root_node.children.values()))
|
||||
|
||||
cache.inc_node_lock_ref(node)
|
||||
self.assertEqual(cache.evictable_size_, 0)
|
||||
self.assertEqual(cache.protected_size_, 64)
|
||||
|
||||
cache.dec_node_lock_ref(node)
|
||||
self.assertEqual(cache.evictable_size_, 64)
|
||||
self.assertEqual(cache.protected_size_, 0)
|
||||
|
||||
def test_cp_valid_tail_regular_evict_reports_and_subtracts_physical_page(self):
|
||||
cache = self._minimal_cp_hiradix_cache(page_size=64)
|
||||
device_allocator = RecordingDeviceAllocator()
|
||||
cache.cache_controller.mem_pool_device_allocator = device_allocator
|
||||
cache.insert(
|
||||
InsertParams(
|
||||
key=RadixKey(token_ids=[1, 2, 3], extra_key=None),
|
||||
value=torch.arange(3, dtype=torch.int64),
|
||||
)
|
||||
)
|
||||
node = next(iter(cache.root_node.children.values()))
|
||||
|
||||
num_evicted = cache._evict_regular(node)
|
||||
|
||||
self.assertEqual(num_evicted, 64)
|
||||
self.assertEqual(cache.evictable_size_, 0)
|
||||
self.assertEqual(device_allocator.freed, [[0, 1, 2]])
|
||||
self.assertEqual(cache.root_node.children, {})
|
||||
|
||||
def test_cp_partial_split_floors_unbacked_valid_tail_to_page_boundary(self):
|
||||
cache = self._minimal_cp_hiradix_cache(page_size=64)
|
||||
node = TreeNode()
|
||||
node.host_len = 0
|
||||
node.cp_hicache = None
|
||||
|
||||
self.assertEqual(cache._cp_floor_backed_partial_split_len(node, 3), 0)
|
||||
self.assertEqual(cache._cp_floor_backed_partial_split_len(node, 70), 64)
|
||||
self.assertEqual(cache._cp_floor_backed_partial_split_len(node, 128), 128)
|
||||
|
||||
def test_session_aware_cache_forwards_cp_hicache_prepare(self):
|
||||
calls = []
|
||||
|
||||
|
||||
@@ -7,16 +7,45 @@ 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
|
||||
|
||||
|
||||
def _define_sgl_kernel_stub(schema: str) -> None:
|
||||
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
|
||||
|
||||
|
||||
_define_sgl_kernel_stub(
|
||||
"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)"
|
||||
)
|
||||
_define_sgl_kernel_stub(
|
||||
"sgl_per_token_quant_fp8(Tensor input, Tensor output_q, Tensor output_s) -> ()"
|
||||
)
|
||||
_define_sgl_kernel_stub(
|
||||
"sgl_per_token_group_quant_fp8(Tensor input, Tensor output_q, Tensor output_s, "
|
||||
"int group_size, float eps, float fp8_min, float fp8_max, bool scale_ue8m0) -> ()"
|
||||
)
|
||||
_define_sgl_kernel_stub(
|
||||
"sgl_per_token_group_quant_8bit(Tensor input, Tensor output_q, Tensor output_s, "
|
||||
"int group_size, float eps, float fp8_min, float fp8_max, bool scale_ue8m0) -> ()"
|
||||
)
|
||||
_define_sgl_kernel_stub(
|
||||
"fp8_scaled_mm(Tensor mat_a, Tensor mat_b, Tensor scales_a, Tensor scales_b, "
|
||||
"ScalarType out_dtype, Tensor? bias) -> Tensor"
|
||||
)
|
||||
_define_sgl_kernel_stub(
|
||||
"fp8_blockwise_scaled_mm(Tensor mat_a, Tensor mat_b, Tensor scales_a, "
|
||||
"Tensor scales_b, ScalarType out_dtype) -> Tensor"
|
||||
)
|
||||
_define_sgl_kernel_stub(
|
||||
"int8_scaled_mm(Tensor mat_a, Tensor mat_b, Tensor scales_a, Tensor scales_b, "
|
||||
"ScalarType out_dtype, Tensor? bias) -> Tensor"
|
||||
)
|
||||
|
||||
flash_attn_stub = sys.modules.setdefault(
|
||||
"sgl_kernel.flash_attn", types.ModuleType("sgl_kernel.flash_attn")
|
||||
@@ -32,6 +61,72 @@ if not hasattr(sgl_kernel_stub, "__path__"):
|
||||
sgl_kernel_stub.__path__ = []
|
||||
if not hasattr(sgl_kernel_stub, "flash_attn"):
|
||||
sgl_kernel_stub.flash_attn = flash_attn_stub
|
||||
quantization_stub = sys.modules.setdefault(
|
||||
"sgl_kernel.quantization", types.ModuleType("sgl_kernel.quantization")
|
||||
)
|
||||
if not hasattr(sgl_kernel_stub, "quantization"):
|
||||
sgl_kernel_stub.quantization = quantization_stub
|
||||
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)
|
||||
if not hasattr(sgl_kernel_stub, "sgl_per_token_quant_fp8"):
|
||||
sgl_kernel_stub.sgl_per_token_quant_fp8 = lambda *args, **kwargs: None
|
||||
if not hasattr(sgl_kernel_stub, "sgl_per_token_group_quant_fp8"):
|
||||
sgl_kernel_stub.sgl_per_token_group_quant_fp8 = lambda *args, **kwargs: None
|
||||
if not hasattr(sgl_kernel_stub, "sgl_per_token_group_quant_int8"):
|
||||
sgl_kernel_stub.sgl_per_token_group_quant_int8 = lambda *args, **kwargs: None
|
||||
for _name in (
|
||||
"concat_mla_absorb_q",
|
||||
"gelu_and_mul",
|
||||
"silu_and_mul",
|
||||
"moe_align_block_size",
|
||||
"moe_sum",
|
||||
"moe_sum_reduce",
|
||||
"moe_fused_gate",
|
||||
"kimi_k2_moe_fused_gate",
|
||||
"topk_softmax",
|
||||
"topk_sigmoid",
|
||||
"fast_topk_transform_fused",
|
||||
"fast_topk_transform_ragged_fused",
|
||||
"fast_topk_v2",
|
||||
"fused_add_rmsnorm",
|
||||
"gemma_fused_add_rmsnorm",
|
||||
"gemma_rmsnorm",
|
||||
"sgl_per_token_group_quant_8bit",
|
||||
"fp8_blockwise_scaled_mm",
|
||||
"fp8_scaled_mm",
|
||||
"int8_scaled_mm",
|
||||
"gptq_gemm",
|
||||
"gptq_shuffle",
|
||||
"qserve_w4a8_per_chn_gemm",
|
||||
"qserve_w4a8_per_group_gemm",
|
||||
"awq_dequantize",
|
||||
"fused_experts",
|
||||
"apply_shuffle_mul_sum",
|
||||
"es_fp8_blockwise_scaled_grouped_mm",
|
||||
"es_sm100_mxfp8_blockscaled_grouped_mm",
|
||||
"es_sm100_mxfp8_blockscaled_grouped_quant",
|
||||
"fp8_blockwise_scaled_grouped_mm",
|
||||
"prepare_moe_input",
|
||||
"shuffle_rows",
|
||||
"cutlass_w4a8_moe_mm",
|
||||
"get_cutlass_w4a8_moe_mm_data",
|
||||
"merge_state_v2",
|
||||
"cutlass_mla_decode",
|
||||
"cutlass_mla_get_workspace_size",
|
||||
"causal_conv1d_fwd",
|
||||
"causal_conv1d_update",
|
||||
"rmsnorm",
|
||||
):
|
||||
if not hasattr(sgl_kernel_stub, _name):
|
||||
setattr(sgl_kernel_stub, _name, lambda *args, **kwargs: None)
|
||||
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
|
||||
@@ -252,7 +347,25 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
|
||||
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
|
||||
index_accessor_stub = types.ModuleType(
|
||||
"sglang.srt.layers.attention.nsa.index_buf_accessor"
|
||||
)
|
||||
|
||||
class _IndexAccessorOp:
|
||||
@staticmethod
|
||||
def execute(*args, **kwargs):
|
||||
raise AssertionError("index accessor is not used by this test")
|
||||
|
||||
for _name in ("GetK", "GetS", "GetKAndS", "SetKAndS"):
|
||||
setattr(index_accessor_stub, _name, _IndexAccessorOp)
|
||||
|
||||
with patch.dict(
|
||||
sys.modules,
|
||||
{
|
||||
"sglang.srt.layers.attention.nsa.index_buf_accessor": index_accessor_stub
|
||||
},
|
||||
):
|
||||
from sglang.srt.mem_cache.memory_pool import MLATokenToKVPool
|
||||
|
||||
class FakeCounter:
|
||||
def __init__(self):
|
||||
@@ -483,6 +596,10 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
|
||||
)
|
||||
|
||||
forward_batch.spec_info = TargetSpecInfo()
|
||||
forward_batch.cp_shared_kv_mla_prefetcher = object()
|
||||
self.assertTrue(runtime.should_reuse_current_extend_kv(forward_batch))
|
||||
|
||||
forward_batch.cp_shared_kv_mla_prefetcher = None
|
||||
self.assertTrue(runtime.should_reuse_current_extend_kv(forward_batch))
|
||||
|
||||
forward_batch.spec_info = DraftSpecInfo()
|
||||
@@ -490,6 +607,10 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
|
||||
forward_batch.seq_lens_cpu = torch.tensor([56], dtype=torch.int32)
|
||||
self.assertTrue(runtime.should_reuse_current_extend_kv(forward_batch))
|
||||
|
||||
forward_batch.spec_info = TargetSpecInfo()
|
||||
forward_batch.cp_shared_kv_mla_prefetcher = None
|
||||
self.assertTrue(runtime.should_reuse_current_extend_kv(forward_batch))
|
||||
|
||||
def test_runtime_fallback_helpers_use_standard_warning_marker(self):
|
||||
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
|
||||
|
||||
@@ -609,6 +730,49 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
|
||||
self.assertEqual(current_mask.tolist(), [[False, True, True, False, False]])
|
||||
self.assertEqual(mixed_locs.tolist(), [[4, 12, 13, -1, -1]])
|
||||
|
||||
def test_materialize_prefix_and_reuse_current_kv_page_slots_without_prefetcher(
|
||||
self,
|
||||
):
|
||||
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
|
||||
|
||||
page_size = 4
|
||||
layout = CpSharedKVLayout(page_size=page_size, cp_size=1, cp_rank=0)
|
||||
kv_cache = torch.arange(0, 32, dtype=torch.float32).view(32, 1, 1)
|
||||
logical_locs = torch.tensor([[4, 8, 20, 21, 22, 23]], dtype=torch.int64)
|
||||
current_locs = torch.tensor([20, 21], dtype=torch.int64)
|
||||
current_kv = torch.arange(100, 102, dtype=torch.float32).view(2, 1, 1)
|
||||
remap_logical_pages = torch.tensor([[1, 2, 5]], dtype=torch.int64)
|
||||
slot_remap = runtime.build_shared_token_kv_slot_remap(
|
||||
kv_cache=kv_cache,
|
||||
logical_locs=logical_locs,
|
||||
remap_logical_pages=remap_logical_pages,
|
||||
layout=layout,
|
||||
page_size=page_size,
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
runtime, "_all_reduce_materialized_buffer_range", _identity_all_reduce
|
||||
):
|
||||
mixed_kv, mixed_locs = (
|
||||
runtime.materialize_prefix_and_reuse_current_kv_page_slots(
|
||||
kv_cache=kv_cache,
|
||||
logical_locs=logical_locs,
|
||||
current_kv_cache=current_kv,
|
||||
current_locs=current_locs,
|
||||
slot_remap=slot_remap,
|
||||
layout=layout,
|
||||
page_size=page_size,
|
||||
prefix_pages=2,
|
||||
)
|
||||
)
|
||||
|
||||
self.assertEqual(list(mixed_kv.shape), [16, 1, 1])
|
||||
self.assertTrue(torch.equal(mixed_kv[4:8], kv_cache[4:8]))
|
||||
self.assertTrue(torch.equal(mixed_kv[8:12], kv_cache[8:12]))
|
||||
self.assertTrue(torch.equal(mixed_kv[12:14], current_kv))
|
||||
self.assertEqual(mixed_locs.tolist(), [[4, 8, 12, 13, -1, -1]])
|
||||
|
||||
def test_mla_prefetch_consume_prefix_with_current_skips_suffix_materialize(self):
|
||||
from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
|
||||
@@ -1240,6 +1404,103 @@ class TestCpSharedKVRuntimeHelpers(unittest.TestCase):
|
||||
32,
|
||||
)
|
||||
|
||||
def test_mla_prefetch_min_async_extend_tokens_defaults_to_one_page_and_can_override(
|
||||
self,
|
||||
):
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
|
||||
|
||||
envs.SGLANG_CP_SHARED_KV_MLA_PREFETCH_MIN_EXTEND_TOKENS.clear()
|
||||
self.assertEqual(
|
||||
runtime.cp_shared_kv_mla_prefetch_min_async_extend_tokens(page_size=64),
|
||||
64,
|
||||
)
|
||||
self.assertEqual(
|
||||
runtime.cp_shared_kv_mla_prefetch_min_async_extend_tokens(page_size=None),
|
||||
0,
|
||||
)
|
||||
|
||||
with envs.SGLANG_CP_SHARED_KV_MLA_PREFETCH_MIN_EXTEND_TOKENS.override(0):
|
||||
self.assertEqual(
|
||||
runtime.cp_shared_kv_mla_prefetch_min_async_extend_tokens(page_size=64),
|
||||
0,
|
||||
)
|
||||
|
||||
with envs.SGLANG_CP_SHARED_KV_MLA_PREFETCH_MIN_EXTEND_TOKENS.override(128):
|
||||
self.assertEqual(
|
||||
runtime.cp_shared_kv_mla_prefetch_min_async_extend_tokens(page_size=64),
|
||||
128,
|
||||
)
|
||||
|
||||
def test_mla_and_index_prefetch_skip_tiny_extend_even_with_large_prefix(self):
|
||||
from sglang.srt.layers.attention.nsa import cp_shared_kv_prefetch as prefetch
|
||||
|
||||
class Mode:
|
||||
def is_context_parallel_extend(self):
|
||||
return True
|
||||
|
||||
forward_batch = SimpleNamespace(
|
||||
uses_cp_shared_kv=True,
|
||||
hisparse_coordinator=None,
|
||||
forward_mode=Mode(),
|
||||
batch_size=1,
|
||||
token_to_kv_pool=SimpleNamespace(page_size=64, start_layer=0),
|
||||
cp_shared_kv_layout=SimpleNamespace(cp_size=8, cp_rank=0),
|
||||
extend_prefix_lens_cpu=[16320],
|
||||
extend_seq_lens_cpu=[16],
|
||||
)
|
||||
metadata = SimpleNamespace(
|
||||
real_page_table=torch.arange(256, dtype=torch.int64),
|
||||
page_table_1=torch.zeros((1, 16336), dtype=torch.int32),
|
||||
)
|
||||
stream = object()
|
||||
kv_cache = torch.zeros((4096, 2), dtype=torch.float32)
|
||||
remap = SimpleNamespace(
|
||||
slot_logical_pages=torch.arange(1, 257, dtype=torch.int64),
|
||||
page_inverse=torch.arange(0, 257, dtype=torch.int64),
|
||||
dense_num_pages=257,
|
||||
)
|
||||
|
||||
with patch.object(
|
||||
prefetch, "cp_shared_kv_mla_prefetch_enabled", return_value=True
|
||||
), patch.object(
|
||||
prefetch, "cp_shared_kv_debug_enabled", return_value=False
|
||||
), patch.object(
|
||||
prefetch.torch.cuda, "is_available", return_value=True
|
||||
), patch.object(
|
||||
prefetch, "_is_cuda_stream_capturing", return_value=False
|
||||
), patch.object(
|
||||
prefetch, "is_nsa_prefill_cp_in_seq_split", return_value=True
|
||||
), patch.object(
|
||||
prefetch, "get_attention_cp_group", return_value=SimpleNamespace(pynccl_comm=object())
|
||||
), patch.object(
|
||||
prefetch.torch.cuda, "Stream", return_value=stream
|
||||
), patch.object(
|
||||
prefetch, "_prefetch_pool_get_key_buffer", return_value=kv_cache
|
||||
) as mla_getter, patch.object(
|
||||
prefetch, "get_or_build_shared_token_kv_slot_remap", return_value=remap
|
||||
) as token_remap, patch.object(
|
||||
prefetch,
|
||||
"_prefetch_pool_get_index_buffer",
|
||||
side_effect=AssertionError("index getter should not be reached"),
|
||||
) as index_getter:
|
||||
mla_result = prefetch.CpSharedKVMlaPrefetcher.maybe_create(
|
||||
forward_batch=forward_batch,
|
||||
metadata=metadata,
|
||||
topk_transform_is_paged=True,
|
||||
)
|
||||
index_result = prefetch.CpSharedKVIndexPrefetcher.maybe_create(
|
||||
forward_batch=forward_batch,
|
||||
metadata=metadata,
|
||||
topk_transform_is_paged=True,
|
||||
)
|
||||
|
||||
self.assertIsNone(mla_result)
|
||||
self.assertIsNone(index_result)
|
||||
mla_getter.assert_not_called()
|
||||
token_remap.assert_not_called()
|
||||
index_getter.assert_not_called()
|
||||
|
||||
def test_fused_mla_store_uses_tai_kernel_when_enabled(self):
|
||||
from sglang.srt.layers.attention.nsa import cp_shared_kv_runtime as runtime
|
||||
from sglang.srt.mem_cache.cp_shared_kv_layout import CpSharedKVLayout
|
||||
|
||||
Reference in New Issue
Block a user