diff --git a/python/sglang/srt/layers/attention/dsa/transform_index.py b/python/sglang/srt/layers/attention/dsa/transform_index.py index 10b1068f5..9ae724580 100644 --- a/python/sglang/srt/layers/attention/dsa/transform_index.py +++ b/python/sglang/srt/layers/attention/dsa/transform_index.py @@ -1,3 +1,4 @@ +from itertools import accumulate from typing import List, Optional import torch @@ -6,11 +7,39 @@ import triton.language as tl def transform_index_page_table_prefill(**kwargs): - return transform_index_page_table_prefill_ref(**kwargs) + return transform_index_page_table_prefill_fast(**kwargs) def transform_index_page_table_decode(**kwargs): - return transform_index_page_table_decode_ref(**kwargs) + return transform_index_page_table_decode_fast(**kwargs) + + +def _allocate_prefill_result( + topk_indices: torch.Tensor, + real_num_tokens: int, + output_num_tokens: Optional[int], +) -> torch.Tensor: + topk_num_tokens = topk_indices.shape[0] + if output_num_tokens is None: + output_num_tokens = topk_num_tokens + + assert real_num_tokens <= topk_num_tokens, ( + f"sum(extend_lens_cpu) ({real_num_tokens}) exceeds " + f"topk_indices rows ({topk_num_tokens})" + ) + assert topk_num_tokens <= output_num_tokens, ( + f"topk_indices rows ({topk_num_tokens}) exceeds " + f"output_num_tokens ({output_num_tokens})" + ) + + result = torch.empty( + (output_num_tokens, topk_indices.shape[1]), + dtype=torch.int32, + device=topk_indices.device, + ) + if real_num_tokens < output_num_tokens: + result[real_num_tokens:].fill_(-1) + return result @triton.jit @@ -19,11 +48,11 @@ def transform_index_page_table_decode_kernel( topk_indices_ptr: torch.Tensor, result_ptr: torch.Tensor, page_size: tl.constexpr, - max_seqlen_k: tl.constexpr, + page_table_row_stride: tl.constexpr, ): TOPK: tl.constexpr = 2048 req_id = tl.program_id(0) - page_table_ptr = page_table_ptr + req_id * max_seqlen_k + page_table_ptr = page_table_ptr + req_id * page_table_row_stride topk_indices_ptr = topk_indices_ptr + req_id * TOPK result_ptr = result_ptr + req_id * TOPK @@ -35,6 +64,61 @@ def transform_index_page_table_decode_kernel( tl.store(result_ptr + offset, -1, mask=~mask) +@triton.jit +def transform_index_page_table_prefill_kernel( + page_table_ptr: torch.Tensor, + topk_indices_ptr: torch.Tensor, + cu_seqlens_q_ptr: torch.Tensor, + result_ptr: torch.Tensor, + page_table_stride_0: tl.constexpr, + page_table_stride_1: tl.constexpr, + topk_indices_stride_0: tl.constexpr, + topk_indices_stride_1: tl.constexpr, + result_stride_0: tl.constexpr, + result_stride_1: tl.constexpr, + PAGE_TABLE_IS_EXPANDED: tl.constexpr, + TOPK: tl.constexpr, + BLOCK_Q: tl.constexpr, + BLOCK_TOPK: tl.constexpr, +): + request_id = tl.program_id(0) + query_offsets = tl.program_id(1) * BLOCK_Q + tl.arange(0, BLOCK_Q) + topk_offsets = tl.program_id(2) * BLOCK_TOPK + tl.arange(0, BLOCK_TOPK) + + query_start = tl.load(cu_seqlens_q_ptr + request_id) + query_end = tl.load(cu_seqlens_q_ptr + request_id + 1) + token_indices = query_start + query_offsets + mask = (token_indices[:, None] < query_end) & (topk_offsets[None, :] < TOPK) + + loaded_topk_indices = tl.load( + topk_indices_ptr + + token_indices[:, None] * topk_indices_stride_0 + + topk_offsets[None, :] * topk_indices_stride_1, + mask=mask, + other=-1, + ) + valid_topk_mask = mask & (loaded_topk_indices >= 0) + + if PAGE_TABLE_IS_EXPANDED: + page_table_rows = token_indices + else: + page_table_rows = token_indices * 0 + request_id + loaded_kv_indices = tl.load( + page_table_ptr + + page_table_rows[:, None] * page_table_stride_0 + + loaded_topk_indices * page_table_stride_1, + mask=valid_topk_mask, + other=-1, + ) + tl.store( + result_ptr + + token_indices[:, None] * result_stride_0 + + topk_offsets[None, :] * result_stride_1, + loaded_kv_indices, + mask=mask, + ) + + def transform_index_page_table_decode_fast( page_table: torch.Tensor, topk_indices: torch.Tensor, @@ -54,7 +138,6 @@ def transform_index_page_table_decode_fast( assert page_table.shape[0] == topk_indices.shape[0] assert topk_indices.shape[1] == 2048 qo_len = topk_indices.shape[0] - max_seqlen_k = page_table.shape[1] if result is None: result = torch.empty_like(topk_indices, dtype=torch.int32) # Launch triton kernel @@ -64,7 +147,7 @@ def transform_index_page_table_decode_fast( topk_indices, result, page_size, - max_seqlen_k=max_seqlen_k, + page_table_row_stride=page_table.stride(0), ) return result @@ -74,20 +157,48 @@ def transform_index_page_table_prefill_fast( topk_indices: torch.Tensor, extend_lens_cpu: List[int], page_size: int = 1, + output_num_tokens: Optional[int] = None, + page_table_is_expanded: bool = False, + cu_seqlens_q: Optional[torch.Tensor] = None, ) -> torch.Tensor: - # TODO(baizhou): can be implemented with another triton kernel assert page_size == 1 - result = torch.empty_like(topk_indices, dtype=torch.int32) - assert len(extend_lens_cpu) == page_table.shape[0] - offset = 0 - for i, l in enumerate(extend_lens_cpu): - transform_index_page_table_decode_fast( - page_table[i].unsqueeze(0).expand(l, -1), - topk_indices[offset : offset + l], - result=result[offset : offset + l], + assert topk_indices.shape[1] == 2048 + real_num_tokens = sum(extend_lens_cpu) + result = _allocate_prefill_result(topk_indices, real_num_tokens, output_num_tokens) + if real_num_tokens == 0: + return result + + max_extend_len = max(extend_lens_cpu) + block_q = 1 if max_extend_len == 1 else 2 if max_extend_len == 2 else 4 + block_topk = 256 + if cu_seqlens_q is None: + cu_seqlens_q = torch.tensor( + [0, *accumulate(extend_lens_cpu)], + dtype=torch.int32, + device=topk_indices.device, ) - offset += l - assert offset == topk_indices.shape[0] + grid = ( + cu_seqlens_q.shape[0] - 1, + triton.cdiv(max_extend_len, block_q), + triton.cdiv(topk_indices.shape[1], block_topk), + ) + transform_index_page_table_prefill_kernel[grid]( + page_table, + topk_indices, + cu_seqlens_q, + result, + page_table.stride(0), + page_table.stride(1), + topk_indices.stride(0), + topk_indices.stride(1), + result.stride(0), + result.stride(1), + PAGE_TABLE_IS_EXPANDED=page_table_is_expanded, + TOPK=topk_indices.shape[1], + BLOCK_Q=block_q, + BLOCK_TOPK=block_topk, + num_warps=4, + ) return result @@ -117,10 +228,22 @@ def transform_index_page_table_prefill_ref( topk_indices: torch.Tensor, extend_lens_cpu: List[int], page_size: int = 1, + output_num_tokens: Optional[int] = None, + page_table_is_expanded: bool = False, ) -> torch.Tensor: assert page_size == 1 - result = torch.empty_like(topk_indices, dtype=torch.int32) - assert len(extend_lens_cpu) == page_table.shape[0] + real_num_tokens = sum(extend_lens_cpu) + result = _allocate_prefill_result(topk_indices, real_num_tokens, output_num_tokens) + + if page_table_is_expanded: + if real_num_tokens > 0: + transform_index_page_table_decode_ref( + page_table[:real_num_tokens], + topk_indices[:real_num_tokens], + result=result[:real_num_tokens], + ) + return result + offset = 0 for i, l in enumerate(extend_lens_cpu): transform_index_page_table_decode_ref( @@ -129,7 +252,6 @@ def transform_index_page_table_prefill_ref( result=result[offset : offset + l], ) offset += l - assert offset == topk_indices.shape[0] return result diff --git a/python/sglang/srt/layers/attention/dsa_backend.py b/python/sglang/srt/layers/attention/dsa_backend.py index c716b0ea3..7416bb4f2 100644 --- a/python/sglang/srt/layers/attention/dsa_backend.py +++ b/python/sglang/srt/layers/attention/dsa_backend.py @@ -1898,19 +1898,20 @@ class DeepseekSparseAttnBackend( q_nope = q_all[:, :, : layer.v_head_dim] q_rope = q_all[:, :, layer.v_head_dim :] - # Align topk_indices with q dimensions - # This handles cases where q is padded (TP + partial DP attention) - if topk_indices is not None: - topk_indices = self._pad_topk_indices(topk_indices, q_nope.shape[0]) - # NOTE(dark): here, we use page size = 1 topk_transform_method = self.get_topk_transform_method( forward_batch.forward_mode ) + if self.use_fused_topk: + if topk_indices is not None: + topk_indices = self._pad_topk_indices(topk_indices, q_nope.shape[0]) + page_table_1 = self._get_fused_topk_page_table(topk_indices) else: if topk_transform_method == TopkTransformMethod.RAGGED: + if topk_indices is not None: + topk_indices = self._pad_topk_indices(topk_indices, q_nope.shape[0]) topk_indices_offset = metadata.topk_indices_offset assert topk_indices_offset is not None mask = topk_indices != -1 @@ -1929,6 +1930,12 @@ class DeepseekSparseAttnBackend( topk_indices=topk_indices, extend_lens_cpu=metadata.dsa_extend_seq_lens_list, page_size=1, + output_num_tokens=q_nope.shape[0], + page_table_is_expanded=( + forward_batch.forward_mode.is_target_verify() + or forward_batch.forward_mode.is_draft_extend_v2() + ), + cu_seqlens_q=metadata.cu_seqlens_q, ) # todo hisparse: to cover more backends @@ -2675,11 +2682,9 @@ class DeepseekSparseAttnBackend( else: q_all = q.view(-1, layer.tp_q_head_num, layer.head_dim) - # Align topk_indices with q dimensions - if topk_indices is not None: - topk_indices = self._pad_topk_indices(topk_indices, q.shape[0]) - if self.use_fused_topk: + if topk_indices is not None: + topk_indices = self._pad_topk_indices(topk_indices, q.shape[0]) page_table_1 = self._get_fused_topk_page_table(topk_indices) elif is_prefill: page_table_1 = transform_index_page_table_prefill( @@ -2687,8 +2692,16 @@ class DeepseekSparseAttnBackend( topk_indices=topk_indices, extend_lens_cpu=metadata.dsa_extend_seq_lens_list, page_size=1, + output_num_tokens=q.shape[0], + page_table_is_expanded=( + forward_batch.forward_mode.is_target_verify() + or forward_batch.forward_mode.is_draft_extend_v2() + ), + cu_seqlens_q=metadata.cu_seqlens_q, ) else: + if topk_indices is not None: + topk_indices = self._pad_topk_indices(topk_indices, q.shape[0]) page_table_1 = transform_index_page_table_decode( page_table=metadata.page_table_1, topk_indices=topk_indices, diff --git a/python/sglang/srt/models/deepseek_nextn.py b/python/sglang/srt/models/deepseek_nextn.py index 7e8cc0b9b..35fcedeea 100644 --- a/python/sglang/srt/models/deepseek_nextn.py +++ b/python/sglang/srt/models/deepseek_nextn.py @@ -33,6 +33,7 @@ from sglang.srt.layers.attention.dsa.utils import ( can_dsa_cp_split, dsa_use_prefill_cp, is_dsa_enable_prefill_cp, + is_dsa_prefill_cp_round_robin_split, ) from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ReplicatedLinear @@ -60,6 +61,34 @@ from sglang.srt.models.utils import WeightsMapper from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.utils import BumpAllocator, add_prefix, is_cuda, is_npu + +def _gather_dsa_topk_indices_for_cp( + topk_indices: torch.Tensor, + local_num_tokens: int, + cp_size: int, + forward_batch: ForwardBatch, + stream, +) -> torch.Tensor: + if ( + is_dsa_prefill_cp_round_robin_split() + and topk_indices.shape[0] < local_num_tokens + ): + pad_rows = local_num_tokens - topk_indices.shape[0] + topk_indices = torch.cat( + [ + topk_indices, + topk_indices.new_full((pad_rows, topk_indices.shape[1]), -1), + ], + dim=0, + ) + return cp_all_gather_rerange_output( + topk_indices, + cp_size, + forward_batch, + stream, + ) + + logger = logging.getLogger(__name__) @@ -222,12 +251,21 @@ class DeepseekModelNextN(nn.Module): else: hidden_states = self.eh_proj(eh_input) - if dsa_use_prefill_cp( + use_cp = dsa_use_prefill_cp( forward_batch, self.dsa_enable_prefill_cp - ) or mla_use_prefill_cp(forward_batch, self.mla_enable_prefill_cp): + ) or mla_use_prefill_cp(forward_batch, self.mla_enable_prefill_cp) + if use_cp: hidden_states = cp_split_and_rebuild_data(forward_batch, hidden_states) positions = cp_split_and_rebuild_position(forward_batch, positions) residual = None + seed_buf = ( + forward_batch.spec_info.dsa_seed_topk_capture + if forward_batch.forward_mode.is_extend(include_draft_extend_v2=True) + else None + ) + should_update_dsa_topk_indices = ( + forward_batch.reuse_dsa_topk_indices or seed_buf is not None + ) with get_global_expert_distribution_recorder().disable_this_region(): hidden_states, residual, topk_indices = self.decoder( positions, @@ -241,34 +279,39 @@ class DeepseekModelNextN(nn.Module): else None ), ) - if forward_batch.reuse_dsa_topk_indices: - forward_batch.spec_info.dsa_topk_indices = topk_indices - - # MTP IndexShare: on draft-extend, publish the last-token DSA - # indexer top-k to seed (avoid recomputing in) the draft-decode loop. - if forward_batch.forward_mode.is_extend(include_draft_extend_v2=True): - seed_buf = forward_batch.spec_info.dsa_seed_topk_capture - if seed_buf is not None and topk_indices is not None: - sel = forward_batch.spec_info.dsa_seed_topk_select - src = topk_indices if sel is None else topk_indices[sel] - seed_buf[: src.shape[0]].copy_(src) - if not forward_batch.forward_mode.is_idle(): if residual is not None: hidden_states, _ = self.shared_head.norm(hidden_states, residual) else: hidden_states = self.shared_head.norm(hidden_states) - if dsa_use_prefill_cp( - forward_batch, self.dsa_enable_prefill_cp - ) or mla_use_prefill_cp(forward_batch, self.mla_enable_prefill_cp): - # allgather + rerrange + if use_cp: + local_num_tokens = hidden_states.shape[0] hidden_states = cp_all_gather_rerange_output( hidden_states, self.cp_size, forward_batch, torch.cuda.current_stream(), ) + if should_update_dsa_topk_indices and topk_indices is not None: + topk_indices = _gather_dsa_topk_indices_for_cp( + topk_indices, + local_num_tokens, + self.cp_size, + forward_batch, + torch.cuda.current_stream(), + ) + if should_update_dsa_topk_indices and topk_indices is not None: + if forward_batch.reuse_dsa_topk_indices: + forward_batch.spec_info.dsa_topk_indices = topk_indices + if seed_buf is not None: + sel = forward_batch.spec_info.dsa_seed_topk_select + src = ( + topk_indices[: seed_buf.shape[0]] + if sel is None + else topk_indices[sel] + ) + seed_buf[: src.shape[0]].copy_(src) finally: exit_stack.close() diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 48ec4791d..6d9bb9405 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -5204,7 +5204,7 @@ class ServerArgs: mode = strategy_to_legacy_mode[self.cp_strategy] use_dsa_legacy_aliases = self.enable_dsa_prefill_context_parallel or getattr( - self, "attention_backend", None + self._resolved(), "attention_backend", None ) in ("dsa", "dsv4") if use_dsa_legacy_aliases: self.enable_dsa_prefill_context_parallel = True diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index c5b2ff109..0f40374f5 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -15,7 +15,6 @@ from sglang.srt.hardware_backend.npu.graph_runner.eagle_draft_npu_graph_runner i ) from sglang.srt.hardware_backend.npu.graph_runner.npu_graph_runner import NPUGraphRunner from sglang.srt.kv_canary.runner.canary_manager import context_tuple -from sglang.srt.layers.attention.dsa.utils import dsa_use_prefill_cp from sglang.srt.layers.attention.flashinfer_backend import FlashInferAttnBackend from sglang.srt.layers.attention.tokenspeed_mla_backend import TokenspeedMLABackend from sglang.srt.layers.attention.triton_backend import TritonAttnBackend @@ -825,11 +824,9 @@ class EagleDraftWorker(EagleDraftWorkerBase): # Seed the first draft-decode loop from each request's last prefill # position. Gather last-per-req before the copy (prefill can be long). - # Skipped under context-parallel prefill (token layout wouldn't match). seed_from_extend = ( self.seed_dsa_topk_from_draft_extend and not forward_batch.forward_mode.is_idle() - and not dsa_use_prefill_cp(forward_batch) ) if seed_from_extend: bs = forward_batch.batch_size diff --git a/test/registered/cp/test_glm52_cp_index_share.py b/test/registered/cp/test_glm52_cp_index_share.py new file mode 100644 index 000000000..b64a1d764 --- /dev/null +++ b/test/registered/cp/test_glm52_cp_index_share.py @@ -0,0 +1,63 @@ +import unittest + +from sglang.srt.utils import kill_process_tree +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kits.eval_accuracy_kit import GSM8KMixin +from sglang.test.test_utils import ( + DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + DEFAULT_URL_FOR_TEST, + CustomTestCase, + popen_launch_server, +) + +register_cuda_ci(est_time=400, stage="extra-b", runner_config="4-gpu-b200") + +GLM52_NVFP4_MODEL_PATH = "nvidia/GLM-5.2-NVFP4" + + +class TestGLM52CPInterleave(GSM8KMixin, CustomTestCase): + gsm8k_accuracy_thres = 0.935 + gsm8k_num_examples = 500 + gsm8k_num_threads = 32 + gsm8k_num_shots = 20 + gsm8k_accept_length_thres = 3 + + @classmethod + def setUpClass(cls): + cls.model = GLM52_NVFP4_MODEL_PATH + cls.base_url = DEFAULT_URL_FOR_TEST + cls.process = popen_launch_server( + cls.model, + cls.base_url, + timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, + other_args=[ + "--tp", + "4", + "--attn-cp-size", + "4", + "--enable-prefill-cp", + "--cp-strategy", + "interleave", + "--speculative-algorithm", + "EAGLE", + "--speculative-num-steps", + "3", + "--speculative-eagle-topk", + "1", + "--speculative-num-draft-tokens", + "4", + "--mem-frac", + "0.85", + "--model-loader-extra-config", + '{"enable_multithread_load": true, "num_threads": 64}', + ], + ) + + @classmethod + def tearDownClass(cls): + if hasattr(cls, "process") and cls.process: + kill_process_tree(cls.process.pid) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/kernels/test_dsa_transform_index.py b/test/registered/kernels/test_dsa_transform_index.py new file mode 100644 index 000000000..d32219487 --- /dev/null +++ b/test/registered/kernels/test_dsa_transform_index.py @@ -0,0 +1,240 @@ +import unittest +from unittest.mock import patch + +import torch + +import sglang.srt.layers.attention.dsa.transform_index as transform_index_module +from sglang.srt.layers.attention.dsa.transform_index import ( + transform_index_page_table_decode_fast, + transform_index_page_table_prefill_fast, +) +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=60, stage="base-b", runner_config="1-gpu-large") + +TOPK = 2048 + + +@unittest.skipUnless(torch.cuda.is_available(), "CUDA is required for this test.") +class TestDSATransformIndex(CustomTestCase): + def setUp(self): + super().setUp() + self.device = torch.device("cuda") + + def tearDown(self): + torch.cuda.empty_cache() + super().tearDown() + + def _make_page_table(self, rows: int, context_length: int) -> torch.Tensor: + columns = torch.arange(context_length, dtype=torch.int32, device=self.device) + row_bias = ( + torch.arange(rows, dtype=torch.int32, device=self.device).unsqueeze(1) * 17 + ) + return columns.unsqueeze(0) + row_bias + + def _make_topk(self, rows: int, context_length: int) -> torch.Tensor: + topk = ( + torch.arange(TOPK, dtype=torch.int64, device=self.device) + .remainder(context_length) + .repeat(rows, 1) + ) + if rows > 0: + topk[:, 0] = 0 + topk[:, 1] = context_length - 1 + topk[:, 257::257] = -1 + return topk + + def _expected( + self, + page_table: torch.Tensor, + topk_indices: torch.Tensor, + extend_lens_cpu: list[int], + output_num_tokens: int, + page_table_is_expanded: bool, + ) -> torch.Tensor: + real_num_tokens = sum(extend_lens_cpu) + expected = torch.full( + (output_num_tokens, TOPK), + -1, + dtype=torch.int32, + device=self.device, + ) + if real_num_tokens == 0: + return expected + + if page_table_is_expanded: + source_rows = page_table[:real_num_tokens] + else: + request_ids = torch.repeat_interleave( + torch.arange( + len(extend_lens_cpu), dtype=torch.int64, device=self.device + ), + torch.tensor(extend_lens_cpu, dtype=torch.int64, device=self.device), + ) + source_rows = page_table[request_ids] + + real_topk = topk_indices[:real_num_tokens] + torch.gather( + source_rows, + dim=1, + index=real_topk.clamp(min=0), + out=expected[:real_num_tokens], + ) + expected[:real_num_tokens][real_topk < 0] = -1 + return expected + + def _check_decode_case( + self, + batch_size: int, + context_length: int, + *, + zero_row_stride: bool = False, + provide_result: bool = False, + ) -> None: + if zero_row_stride: + page_table = self._make_page_table(1, context_length).expand(batch_size, -1) + else: + page_table = self._make_page_table(batch_size, context_length) + topk_indices = self._make_topk(batch_size, context_length) + expected = torch.empty( + (batch_size, TOPK), dtype=torch.int32, device=self.device + ) + torch.gather( + page_table, + dim=1, + index=topk_indices.clamp(min=0), + out=expected, + ) + expected[topk_indices < 0] = -1 + result = torch.empty_like(expected) if provide_result else None + + actual = transform_index_page_table_decode_fast( + page_table=page_table, + topk_indices=topk_indices, + result=result, + ) + torch.cuda.synchronize() + if result is not None: + self.assertIs(actual, result) + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + + def _check_case( + self, + extend_lens_cpu: list[int], + context_length: int, + *, + page_table_is_expanded: bool, + topk_padding: int = 0, + output_padding: int = 0, + ) -> None: + real_num_tokens = sum(extend_lens_cpu) + page_table_rows = ( + real_num_tokens if page_table_is_expanded else len(extend_lens_cpu) + ) + topk_num_tokens = real_num_tokens + topk_padding + output_num_tokens = topk_num_tokens + output_padding + page_table = self._make_page_table(page_table_rows, context_length) + topk_indices = self._make_topk(topk_num_tokens, context_length) + expected = self._expected( + page_table, + topk_indices, + extend_lens_cpu, + output_num_tokens, + page_table_is_expanded, + ) + + actual = transform_index_page_table_prefill_fast( + page_table=page_table, + topk_indices=topk_indices, + extend_lens_cpu=extend_lens_cpu, + output_num_tokens=output_num_tokens, + page_table_is_expanded=page_table_is_expanded, + ) + torch.cuda.synchronize() + torch.testing.assert_close(actual, expected, rtol=0, atol=0) + + def test_prefill_uses_dedicated_kernel(self): + extend_lens_cpu = [2, 1] + context_length = 4096 + page_table = self._make_page_table(len(extend_lens_cpu), context_length) + topk_indices = self._make_topk(sum(extend_lens_cpu), context_length) + + with patch.object( + transform_index_module, + "transform_index_page_table_decode_fast", + side_effect=AssertionError("prefill must not launch decode per request"), + ): + transform_index_page_table_prefill_fast( + page_table=page_table, + topk_indices=topk_indices, + extend_lens_cpu=extend_lens_cpu, + ) + + def test_prefill_uses_device_cu_seqlens(self): + extend_lens_cpu = [2, 1] + context_length = 4096 + page_table = self._make_page_table(len(extend_lens_cpu), context_length) + topk_indices = self._make_topk(sum(extend_lens_cpu), context_length) + cu_seqlens_q = torch.tensor([0, 2, 3], dtype=torch.int32, device=self.device) + + with patch.object( + transform_index_module.torch, + "tensor", + side_effect=AssertionError("must reuse device-side metadata"), + ): + transform_index_page_table_prefill_fast( + page_table=page_table, + topk_indices=topk_indices, + extend_lens_cpu=extend_lens_cpu, + cu_seqlens_q=cu_seqlens_q, + ) + + def test_mixed_lengths_padding_and_empty_batch(self): + self._check_case( + [0, 3, 1, 0, 4], + 8192, + page_table_is_expanded=False, + topk_padding=5, + output_padding=7, + ) + self._check_case( + [0, 0], + 16, + page_table_is_expanded=False, + output_padding=8, + ) + + def test_large_batch_size(self): + self._check_case( + [1] * 8192, + 4096, + page_table_is_expanded=False, + ) + + def test_large_context_lengths(self): + for context_length, page_table_is_expanded in ( + (640_000, True), + (1_000_000, False), + ): + with self.subTest( + context_length=context_length, + page_table_is_expanded=page_table_is_expanded, + ): + self._check_case( + [2, 1], + context_length, + page_table_is_expanded=page_table_is_expanded, + ) + + def test_decode_fast_correctness_and_strides(self): + self._check_decode_case(17, 8192, provide_result=True) + self._check_decode_case(17, 8192, zero_row_stride=True) + + def test_decode_fast_extreme_shapes(self): + self._check_decode_case(8192, 4096) + self._check_decode_case(2, 1_000_000) + + +if __name__ == "__main__": + unittest.main()