From a6ee64d237a2c19d067ed9998245c83b44de6290 Mon Sep 17 00:00:00 2001 From: YAMY <74099316+YAMY1234@users.noreply.github.com> Date: Thu, 2 Jul 2026 22:54:40 -0700 Subject: [PATCH] [DeepSeek-V4] Add an opt-in non-paged indexer for long-context prefill (#29619) --- python/sglang/srt/environ.py | 1 + .../layers/attention/deepseek_v4_backend.py | 13 +- .../srt/layers/attention/dsv4/indexer.py | 252 +++++++++++++++--- .../srt/layers/attention/dsv4/metadata.py | 23 +- .../srt/mem_cache/deepseek_v4_memory_pool.py | 21 +- .../unit/layers/test_dsv4_nonpaged_indexer.py | 208 +++++++++++++++ 6 files changed, 468 insertions(+), 50 deletions(-) create mode 100644 test/registered/unit/layers/test_dsv4_nonpaged_indexer.py diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 4d4f6527c..6cf331007 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -866,6 +866,7 @@ class Envs: SGLANG_OPT_FUSE_MHC_POST_PRE = EnvBool(False) SGLANG_OPT_USE_TILELANG_INDEXER = EnvBool(False) SGLANG_OPT_USE_AITER_INDEXER = EnvBool(False) + SGLANG_OPT_DSV4_NONPAGED_INDEXER = EnvBool(False) SGLANG_OPT_USE_JIT_INDEXER_METADATA = EnvBool(True) SGLANG_OPT_USE_ONLINE_COMPRESS = EnvBool(False) SGLANG_EXPERIMENTAL_ONLINE_C128_MTP = EnvBool(False) diff --git a/python/sglang/srt/layers/attention/deepseek_v4_backend.py b/python/sglang/srt/layers/attention/deepseek_v4_backend.py index c6c21a1b6..a913d8e40 100644 --- a/python/sglang/srt/layers/attention/deepseek_v4_backend.py +++ b/python/sglang/srt/layers/attention/deepseek_v4_backend.py @@ -543,11 +543,17 @@ class DeepseekV4AttnBackend( online_state_slot_offset=online_c128_state_slot_offset, ) - def init_forward_metadata_indexer(self, core_attn_metadata: DSV4AttnMetadata): + def init_forward_metadata_indexer( + self, + core_attn_metadata: DSV4AttnMetadata, + *, + use_prefill_cuda_graph: bool = False, + ): return PagedIndexerMetadata( page_size=self.page_size, page_table=core_attn_metadata.page_table, c4_seq_lens=core_attn_metadata.c4_topk_lengths_raw, + use_prefill_cuda_graph=use_prefill_cuda_graph, ) def init_forward_metadata_decode( @@ -630,7 +636,10 @@ class DeepseekV4AttnBackend( is_prefill=True, ) indexer_metadata = ( - self.init_forward_metadata_indexer(core_attn_metadata) + self.init_forward_metadata_indexer( + core_attn_metadata, + use_prefill_cuda_graph=use_prefill_cuda_graph, + ) if need_compress else None ) diff --git a/python/sglang/srt/layers/attention/dsv4/indexer.py b/python/sglang/srt/layers/attention/dsv4/indexer.py index eceba4e18..575b26e8b 100644 --- a/python/sglang/srt/layers/attention/dsv4/indexer.py +++ b/python/sglang/srt/layers/attention/dsv4/indexer.py @@ -17,10 +17,21 @@ from sglang.jit_kernel.dsv4 import ( from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config from sglang.srt.environ import envs from sglang.srt.layers.attention.dsv4.compressor import Compressor -from sglang.srt.layers.attention.dsv4.metadata import PagedIndexerMetadata +from sglang.srt.layers.attention.dsv4.metadata import ( + NonPagedIndexerPlan, + PagedIndexerMetadata, +) +from sglang.srt.layers.dp_attention import get_attention_cp_size from sglang.srt.layers.linear import ReplicatedLinear +from sglang.srt.model_executor.forward_batch_info import ForwardMode +from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context import ( + is_in_breakable_cuda_graph, +) +from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import ( + is_in_tc_piecewise_cuda_graph, +) from sglang.srt.state_capturer.indexer_topk import get_global_indexer_capturer -from sglang.srt.utils import add_prefix, is_hip +from sglang.srt.utils import add_prefix, is_cuda, is_hip from sglang.srt.utils.common import is_sm120_supported if TYPE_CHECKING: @@ -360,10 +371,9 @@ class C4IndexerBackendMixin: c4_indexer: C4Indexer, positions: torch.Tensor, forward_batch: ForwardBatch, - token_to_kv_pool: DeepSeekV4TokenToKVPool, alt_streams: Optional[List[torch.cuda.Stream]] = None, q_lora_ready: Optional[torch.cuda.Event] = None, - ) -> Tuple[IndexerQuery, torch.Tensor, torch.Tensor]: + ) -> Tuple[IndexerQuery, torch.Tensor]: if TYPE_CHECKING: assert isinstance(self, CompressorBackendMixin) @@ -382,9 +392,6 @@ class C4IndexerBackendMixin: layer_id=c4_indexer.layer_id, compressor=c4_indexer.compressor, ) - c4_indexer_kv_cache = token_to_kv_pool.get_index_k_with_scale_buffer( - layer_id=c4_indexer.layer_id, - ) # The weight projection is small and fast; compute it on its own # stream, then have the Q stream wait on it before launching the big @@ -401,7 +408,7 @@ class C4IndexerBackendMixin: q, weights = c4_indexer.compute_q(q_lora, positions, weights) current_stream.wait_stream(stream_q) - return q, weights, c4_indexer_kv_cache + return q, weights def _forward_prepare_normal( self, @@ -410,9 +417,8 @@ class C4IndexerBackendMixin: c4_indexer: C4Indexer, positions: torch.Tensor, forward_batch: ForwardBatch, - token_to_kv_pool: DeepSeekV4TokenToKVPool, skip_compressor: bool = False, - ) -> Tuple[IndexerQuery, torch.Tensor, torch.Tensor]: + ) -> Tuple[IndexerQuery, torch.Tensor]: if TYPE_CHECKING: assert isinstance(self, CompressorBackendMixin) @@ -425,10 +431,159 @@ class C4IndexerBackendMixin: layer_id=c4_indexer.layer_id, compressor=c4_indexer.compressor, ) - c4_indexer_kv_cache = token_to_kv_pool.get_index_k_with_scale_buffer( + return q, weights + + def _can_use_nonpaged_indexer( + self, + *, + c4_indexer: C4Indexer, + forward_batch: ForwardBatch, + indexer_metadata: PagedIndexerMetadata, + ) -> bool: + if not envs.SGLANG_OPT_DSV4_NONPAGED_INDEXER.get(): + return False + # This path calls CUDA DeepGEMM and assumes the CUDA FP8+FP32 packed + # indexer cache layout. Explicitly reject HIP, NPU, and other devices. + if not is_cuda() or is_hip(): + return False + # The gather plan is built from eager, child-local ForwardBatch metadata. + # Rewritten, TBO-split, and graph-backed batches must use the paged path. + if ( + forward_batch.forward_mode != ForwardMode.EXTEND + or forward_batch._original_forward_mode is not None + or forward_batch.tbo_parent_token_range is not None + or forward_batch.batch_size != 1 + or indexer_metadata.use_prefill_cuda_graph + ): + return False + if ( + c4_indexer.use_fp4_indexer + or envs.SGLANG_OPT_USE_TILELANG_INDEXER.get() + or envs.SGLANG_OPT_USE_AITER_INDEXER.get() + or envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.get() + ): + return False + if ( + get_attention_cp_size() != 1 + or self.hisparse_coordinator is not None + or is_in_tc_piecewise_cuda_graph() + or is_in_breakable_cuda_graph() + ): + return False + return not torch.cuda.is_current_stream_capturing() + + def _get_nonpaged_indexer_plan( + self, + *, + c4_indexer: C4Indexer, + forward_batch: ForwardBatch, + indexer_metadata: PagedIndexerMetadata, + page_table: torch.Tensor, + c4_seq_lens: torch.Tensor, + query_rows: int, + ) -> Optional[NonPagedIndexerPlan]: + if not self._can_use_nonpaged_indexer( + c4_indexer=c4_indexer, + forward_batch=forward_batch, + indexer_metadata=indexer_metadata, + ): + return None + if indexer_metadata.nonpaged_plan is not None: + return indexer_metadata.nonpaged_plan + + if ( + forward_batch.seq_lens is None + or forward_batch.seq_lens_cpu is None + or forward_batch.extend_seq_lens_cpu is None + or forward_batch.extend_seq_lens is None + or forward_batch.extend_start_loc is None + or forward_batch.extend_num_tokens is None + ): + return None + + def to_cpu_int_list(values) -> Optional[List[int]]: + if isinstance(values, torch.Tensor): + if values.device.type != "cpu": + return None + values = values.tolist() + return [int(value) for value in values] + + extend_lens_cpu = to_cpu_int_list(forward_batch.extend_seq_lens_cpu) + seq_lens_cpu = to_cpu_int_list(forward_batch.seq_lens_cpu) + if ( + extend_lens_cpu is None + or seq_lens_cpu is None + or len(extend_lens_cpu) != 1 + or len(seq_lens_cpu) != 1 + or extend_lens_cpu[0] <= 0 + ): + return None + + actual_queries = extend_lens_cpu[0] + if ( + actual_queries != query_rows + or int(forward_batch.extend_num_tokens) != query_rows + or forward_batch.seq_lens.numel() != 1 + or forward_batch.extend_seq_lens.numel() != 1 + or forward_batch.extend_start_loc.numel() != 1 + or page_table.dim() != 2 + or page_table.shape[0] < query_rows + or c4_seq_lens.numel() < query_rows + ): + return None + + final_c4_len = seq_lens_cpu[0] // 4 + if final_c4_len <= 0: + return None + + request_page_table = page_table[:1].contiguous() + ke = c4_seq_lens[:query_rows].reshape(-1).to(torch.int32).contiguous() + gather_seq_lens = ke[-1:] + ks = torch.zeros_like(ke) + c4_page_size = indexer_metadata.c4_page_size + max_seqlen_k = (final_c4_len + c4_page_size - 1) // c4_page_size * c4_page_size + plan = NonPagedIndexerPlan( + page_table=request_page_table, + gather_seq_lens=gather_seq_lens, + ks=ks, + ke=ke, + seq_len_sum=final_c4_len, + max_seq_len=final_c4_len, + max_seqlen_k=max_seqlen_k, + query_rows=query_rows, + ) + indexer_metadata.nonpaged_plan = plan + return plan + + @staticmethod + def _forward_nonpaged_indexer( + *, + q_indexer: torch.Tensor, + weights: torch.Tensor, + c4_indexer: C4Indexer, + token_to_kv_pool: DeepSeekV4TokenToKVPool, + plan: NonPagedIndexerPlan, + ) -> torch.Tensor: + import deep_gemm + + k_u8, scale_u8 = token_to_kv_pool.get_index_k_scale_buffer( layer_id=c4_indexer.layer_id, + seq_len_tensor=plan.gather_seq_lens, + page_indices=plan.page_table, + seq_len_sum=plan.seq_len_sum, + max_seq_len=plan.max_seq_len, + ) + k_fp8 = k_u8.view(FP8_DTYPE) + k_scale = scale_u8.view(torch.float32).squeeze(-1) + return deep_gemm.fp8_mqa_logits( + q_indexer[: plan.query_rows], + (k_fp8, k_scale), + weights[: plan.query_rows], + plan.ks, + plan.ke, + clean_logits=False, + max_seqlen_k=plan.max_seqlen_k, ) - return q, weights, c4_indexer_kv_cache def forward_c4_indexer( self, @@ -465,35 +620,27 @@ class C4IndexerBackendMixin: positions = positions[:num_queries] if enable_multi_stream: - q_indexer, weights, c4_indexer_kv_cache = ( - self._forward_prepare_multi_stream( - x=x, - q_lora=q_lora, - c4_indexer=c4_indexer, - positions=positions, - forward_batch=forward_batch, - token_to_kv_pool=token_to_kv_pool, - alt_streams=alt_streams, - q_lora_ready=q_lora_ready, - ) + q_indexer, weights = self._forward_prepare_multi_stream( + x=x, + q_lora=q_lora, + c4_indexer=c4_indexer, + positions=positions, + forward_batch=forward_batch, + alt_streams=alt_streams, + q_lora_ready=q_lora_ready, ) else: assert q_lora_ready is None - q_indexer, weights, c4_indexer_kv_cache = self._forward_prepare_normal( + q_indexer, weights = self._forward_prepare_normal( x=x, q_lora=q_lora, c4_indexer=c4_indexer, positions=positions, forward_batch=forward_batch, - token_to_kv_pool=token_to_kv_pool, skip_compressor=skip_compressor, ) - assert len(c4_indexer_kv_cache.shape) == 2 - block_kv = 64 - num_heads_kv = 1 use_fp4_indexer = c4_indexer.use_fp4_indexer - head_dim_with_sf = 68 if use_fp4_indexer else 132 if use_fp4_indexer: q_fp4, q_sf = q_indexer @@ -504,9 +651,6 @@ class C4IndexerBackendMixin: assert len(q_indexer.shape) == 3 q = q_indexer.unsqueeze(1) - c4_indexer_kv_cache = c4_indexer_kv_cache.view( - c4_indexer_kv_cache.shape[0], block_kv, num_heads_kv, head_dim_with_sf - ) assert len(weights.shape) == 3 weights = weights.squeeze(2) if use_fp4_indexer: @@ -550,16 +694,42 @@ class C4IndexerBackendMixin: _use_aiter = envs.SGLANG_OPT_USE_AITER_INDEXER.get() and not use_fp4_indexer if _c4sl.dim() == 1 and not _use_tilelang and not _use_aiter: _c4sl = _c4sl.unsqueeze(-1) - logits = fn( - q, - c4_indexer_kv_cache, - weights, - _c4sl, - page_table, - indexer_metadata.deep_gemm_metadata, - indexer_metadata.max_c4_seq_len, - False, + nonpaged_plan = self._get_nonpaged_indexer_plan( + c4_indexer=c4_indexer, + forward_batch=forward_batch, + indexer_metadata=indexer_metadata, + page_table=page_table, + c4_seq_lens=c4_seq_lens, + query_rows=query_rows, ) + if nonpaged_plan is not None: + assert isinstance(q_indexer, torch.Tensor) + logits = self._forward_nonpaged_indexer( + q_indexer=q_indexer, + weights=weights, + c4_indexer=c4_indexer, + token_to_kv_pool=token_to_kv_pool, + plan=nonpaged_plan, + ) + else: + c4_indexer_kv_cache = token_to_kv_pool.get_index_k_with_scale_buffer( + layer_id=c4_indexer.layer_id, + ) + assert c4_indexer_kv_cache.dim() == 2 + head_dim_with_sf = 68 if use_fp4_indexer else 132 + c4_indexer_kv_cache = c4_indexer_kv_cache.view( + c4_indexer_kv_cache.shape[0], 64, 1, head_dim_with_sf + ) + logits = fn( + q, + c4_indexer_kv_cache, + weights, + _c4sl, + page_table, + indexer_metadata.deep_gemm_metadata, + indexer_metadata.max_c4_seq_len, + False, + ) assert indexer_metadata.page_table is core_metadata.page_table if self.debug_use_external_c4_sparse_indices: diff --git a/python/sglang/srt/layers/attention/dsv4/metadata.py b/python/sglang/srt/layers/attention/dsv4/metadata.py index e1ce33752..d26d1ccf5 100644 --- a/python/sglang/srt/layers/attention/dsv4/metadata.py +++ b/python/sglang/srt/layers/attention/dsv4/metadata.py @@ -95,13 +95,29 @@ def copy_metadata( ), f"{provided_fields - all_fields=}, {all_fields - provided_fields=}" +@dataclass +class NonPagedIndexerPlan: + page_table: torch.Tensor + gather_seq_lens: torch.Tensor + ks: torch.Tensor + ke: torch.Tensor + seq_len_sum: int + max_seq_len: int + max_seqlen_k: int + query_rows: int + + @dataclass class PagedIndexerMetadata: page_size: int page_table: torch.Tensor c4_seq_lens: torch.Tensor + use_prefill_cuda_graph: bool = False deep_gemm_metadata: Any = field(init=False, repr=False) topk_metadata: torch.Tensor = field(init=False, repr=False) + nonpaged_plan: Optional[NonPagedIndexerPlan] = field( + init=False, repr=False, default=None + ) def __post_init__(self): if ( @@ -156,18 +172,19 @@ class PagedIndexerMetadata: def copy_(self, other: PagedIndexerMetadata): if is_hip(): copy_fields = ["page_table", "c4_seq_lens"] - assign_fields = ["deep_gemm_metadata"] + assign_fields = ["deep_gemm_metadata", "nonpaged_plan"] else: copy_fields = ["page_table", "c4_seq_lens", "deep_gemm_metadata"] - assign_fields = [] + assign_fields = ["nonpaged_plan"] copy_fields += ["topk_metadata"] copy_metadata( src=other, dst=self, - check_eq_fields=["page_size"], + check_eq_fields=["page_size", "use_prefill_cuda_graph"], copy_fields=copy_fields, assign_fields=assign_fields, ) + self.nonpaged_plan = None def maybe_copy_inplace(dst, *, src) -> None: diff --git a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py index d8b0d756a..f3bfef414 100644 --- a/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py +++ b/python/sglang/srt/mem_cache/deepseek_v4_memory_pool.py @@ -321,12 +321,19 @@ class DeepSeekV4IndexerPool(KVCache): def get_index_k_scale_buffer( self, layer_id: int, - seq_len: int, + seq_len_tensor: torch.Tensor, page_indices: torch.Tensor, + seq_len_sum: int, + max_seq_len: int, ) -> Tuple[torch.Tensor, torch.Tensor]: buf = self.index_k_with_scale_buffer[layer_id] return index_buf_accessor.GetKAndS.execute( - self, buf, seq_len=seq_len, page_indices=page_indices + self, + buf, + page_indices=page_indices, + seq_len_tensor=seq_len_tensor, + seq_len_sum=seq_len_sum, + max_seq_len=max_seq_len, ) def set_index_k_scale_buffer( @@ -1055,14 +1062,20 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool): def get_index_k_scale_buffer( self, layer_id: int, - seq_len: int, + seq_len_tensor: torch.Tensor, page_indices: torch.Tensor, + seq_len_sum: int, + max_seq_len: int, ) -> Tuple[torch.Tensor, torch.Tensor]: self.wait_layer_transfer(layer_id) compress_ratio, compress_layer_id, _ = self.layer_mapping[layer_id] assert compress_ratio == 4, f"only c4 has indexer, got {compress_ratio = }" return self.c4_indexer_kv_pool.get_index_k_scale_buffer( - compress_layer_id, seq_len, page_indices + compress_layer_id, + seq_len_tensor, + page_indices, + seq_len_sum, + max_seq_len, ) def set_index_k_scale_buffer( diff --git a/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py b/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py new file mode 100644 index 000000000..d70059c5d --- /dev/null +++ b/test/registered/unit/layers/test_dsv4_nonpaged_indexer.py @@ -0,0 +1,208 @@ +import sys +import unittest +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import torch + +from sglang.srt.environ import envs +from sglang.srt.layers.attention.dsv4.indexer import FP8_DTYPE, C4IndexerBackendMixin +from sglang.srt.layers.attention.dsv4.metadata import NonPagedIndexerPlan +from sglang.srt.model_executor.forward_batch_info import ForwardMode +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + +_INDEXER = "sglang.srt.layers.attention.dsv4.indexer" + + +class TestDSV4NonPagedIndexer(CustomTestCase): + def _is_eligible(self, **overrides): + backend = SimpleNamespace(hisparse_coordinator=None) + c4_indexer = SimpleNamespace(use_fp4_indexer=overrides.get("fp4", False)) + forward_batch = SimpleNamespace( + forward_mode=overrides.get("mode", ForwardMode.EXTEND), + _original_forward_mode=overrides.get("original_mode"), + tbo_parent_token_range=overrides.get("tbo"), + batch_size=overrides.get("batch_size", 1), + ) + metadata = SimpleNamespace( + use_prefill_cuda_graph=overrides.get("prefill_graph", False) + ) + with ( + envs.SGLANG_OPT_DSV4_NONPAGED_INDEXER.override( + overrides.get("enabled", True) + ), + envs.SGLANG_OPT_USE_TILELANG_INDEXER.override(False), + envs.SGLANG_OPT_USE_AITER_INDEXER.override(False), + envs.SGLANG_FP8_PAGED_MQA_LOGITS_TORCH.override(False), + patch(f"{_INDEXER}.is_cuda", return_value=True), + patch(f"{_INDEXER}.is_hip", return_value=False), + patch(f"{_INDEXER}.get_attention_cp_size", return_value=1), + patch( + f"{_INDEXER}.is_in_tc_piecewise_cuda_graph", + return_value=overrides.get("piecewise_graph", False), + ), + patch(f"{_INDEXER}.is_in_breakable_cuda_graph", return_value=False), + patch("torch.cuda.is_current_stream_capturing", return_value=False), + ): + return C4IndexerBackendMixin._can_use_nonpaged_indexer( + backend, + c4_indexer=c4_indexer, + forward_batch=forward_batch, + indexer_metadata=metadata, + ) + + def test_eligibility_is_fail_closed(self): + self.assertIs(envs.SGLANG_OPT_DSV4_NONPAGED_INDEXER.default, False) + self.assertTrue(self._is_eligible()) + for case in ( + {"enabled": False}, + {"mode": ForwardMode.DECODE}, + {"original_mode": ForwardMode.DECODE}, + {"batch_size": 2}, + {"batch_size": 20_000}, + {"tbo": (1, 2)}, + {"prefill_graph": True}, + {"piecewise_graph": True}, + {"fp4": True}, + ): + with self.subTest(case=case): + self.assertFalse(self._is_eligible(**case)) + + def test_single_request_plan_contract(self): + backend = SimpleNamespace(_can_use_nonpaged_indexer=lambda **_: True) + c4_indexer = SimpleNamespace(use_fp4_indexer=False) + query_rows = 4 + batch = SimpleNamespace( + seq_lens=torch.tensor([262], dtype=torch.int32), + seq_lens_cpu=[262], + extend_seq_lens_cpu=[query_rows], + extend_seq_lens=torch.tensor([query_rows], dtype=torch.int32), + extend_start_loc=torch.tensor([0], dtype=torch.int32), + extend_num_tokens=query_rows, + ) + metadata = SimpleNamespace(nonpaged_plan=None, c4_page_size=64) + page_table = torch.tensor([[3, 1]], dtype=torch.int32).repeat(query_rows, 1) + c4_seq_lens = torch.tensor([62, 63, 64, 65], dtype=torch.int32) + + def build_plan(): + return C4IndexerBackendMixin._get_nonpaged_indexer_plan( + backend, + c4_indexer=c4_indexer, + forward_batch=batch, + indexer_metadata=metadata, + page_table=page_table, + c4_seq_lens=c4_seq_lens, + query_rows=query_rows, + ) + + plan = build_plan() + self.assertEqual( + (plan.seq_len_sum, plan.max_seqlen_k, plan.query_rows), + (65, 128, query_rows), + ) + torch.testing.assert_close(plan.page_table, page_table[:1]) + torch.testing.assert_close(plan.ke, c4_seq_lens) + torch.testing.assert_close(plan.gather_seq_lens, c4_seq_lens[-1:]) + + metadata.nonpaged_plan = None + batch.extend_seq_lens_cpu = [2, 2] + self.assertIsNone(build_plan()) + + def test_extreme_plan_metadata_is_bounded_and_fail_closed(self): + backend = SimpleNamespace(_can_use_nonpaged_indexer=lambda **_: True) + c4_indexer = SimpleNamespace(use_fp4_indexer=False) + query_rows = 4 + batch = SimpleNamespace( + seq_lens=torch.tensor([500_000], dtype=torch.int32), + seq_lens_cpu=[500_000], + extend_seq_lens_cpu=[query_rows], + extend_seq_lens=torch.tensor([query_rows], dtype=torch.int32), + extend_start_loc=torch.tensor([0], dtype=torch.int32), + extend_num_tokens=query_rows, + ) + metadata = SimpleNamespace(nonpaged_plan=None, c4_page_size=64) + page_table = torch.zeros((query_rows, 1), dtype=torch.int32) + c4_seq_lens = torch.tensor( + [124_997, 124_998, 124_999, 125_000], dtype=torch.int32 + ) + + def build_plan(): + return C4IndexerBackendMixin._get_nonpaged_indexer_plan( + backend, + c4_indexer=c4_indexer, + forward_batch=batch, + indexer_metadata=metadata, + page_table=page_table, + c4_seq_lens=c4_seq_lens, + query_rows=query_rows, + ) + + plan = build_plan() + self.assertEqual(plan.seq_len_sum, 125_000) + self.assertEqual(plan.max_seq_len, 125_000) + self.assertEqual(plan.max_seqlen_k, 125_056) + + metadata.nonpaged_plan = None + batch.seq_lens = torch.tensor([500_000, 200], dtype=torch.int32) + batch.seq_lens_cpu = [500_000, 200] + batch.extend_seq_lens_cpu = [2, 2] + batch.extend_seq_lens = torch.tensor([2, 2], dtype=torch.int32) + batch.extend_start_loc = torch.tensor([0, 2], dtype=torch.int32) + self.assertIsNone(build_plan()) + + def test_nonpaged_dispatch_uses_gathered_kv_contract(self): + query_rows = 4 + plan = NonPagedIndexerPlan( + page_table=torch.tensor([[3, 1]], dtype=torch.int32), + gather_seq_lens=torch.tensor([65], dtype=torch.int32), + ks=torch.zeros(query_rows, dtype=torch.int32), + ke=torch.tensor([62, 63, 64, 65], dtype=torch.int32), + seq_len_sum=65, + max_seq_len=65, + max_seqlen_k=128, + query_rows=query_rows, + ) + q_indexer = torch.zeros((6, 2, 128), dtype=torch.uint8).view(FP8_DTYPE) + weights = torch.ones((6, 2), dtype=torch.float32) + k_u8 = torch.zeros((65, 128), dtype=torch.uint8) + scale_u8 = torch.zeros((65, 4), dtype=torch.uint8) + token_to_kv_pool = MagicMock() + token_to_kv_pool.get_index_k_scale_buffer.return_value = (k_u8, scale_u8) + c4_indexer = SimpleNamespace(layer_id=17) + expected = MagicMock(name="logits") + deep_gemm = SimpleNamespace(fp8_mqa_logits=MagicMock(return_value=expected)) + + with patch.dict(sys.modules, {"deep_gemm": deep_gemm}): + actual = C4IndexerBackendMixin._forward_nonpaged_indexer( + q_indexer=q_indexer, + weights=weights, + c4_indexer=c4_indexer, + token_to_kv_pool=token_to_kv_pool, + plan=plan, + ) + + self.assertIs(actual, expected) + token_to_kv_pool.get_index_k_scale_buffer.assert_called_once_with( + layer_id=17, + seq_len_tensor=plan.gather_seq_lens, + page_indices=plan.page_table, + seq_len_sum=65, + max_seq_len=65, + ) + call = deep_gemm.fp8_mqa_logits.call_args + torch.testing.assert_close(call.args[0], q_indexer[:query_rows]) + torch.testing.assert_close(call.args[1][0], k_u8.view(FP8_DTYPE)) + torch.testing.assert_close( + call.args[1][1], scale_u8.view(torch.float32).squeeze(-1) + ) + torch.testing.assert_close(call.args[2], weights[:query_rows]) + torch.testing.assert_close(call.args[3], plan.ks) + torch.testing.assert_close(call.args[4], plan.ke) + self.assertEqual(call.kwargs, {"clean_logits": False, "max_seqlen_k": 128}) + + +if __name__ == "__main__": + unittest.main()