diff --git a/python/sglang/kernels/ops/attention/dsv4/candidate_blocks.py b/python/sglang/kernels/ops/attention/dsv4/candidate_blocks.py index a8f0cfe615a0..2df02ef9d0ee 100644 --- a/python/sglang/kernels/ops/attention/dsv4/candidate_blocks.py +++ b/python/sglang/kernels/ops/attention/dsv4/candidate_blocks.py @@ -5,6 +5,9 @@ import triton.language as tl from sglang.kernels.jit.utils import is_arch_support_pdl +from sglang.srt.environ import envs + +_DEEPSELECT_INPUT_ALIGNMENT_BYTES = 1024 @triton.jit @@ -20,7 +23,7 @@ def _candidate_scores_kernel( SCORES, WIDTH: tl.constexpr, STRIDE: tl.constexpr, - BLOCKS: tl.constexpr, + SCORE_STRIDE: tl.constexpr, GROUP: tl.constexpr, GROUP_PAD: tl.constexpr, TILE: tl.constexpr, @@ -39,7 +42,7 @@ def _candidate_scores_kernel( scores = tl.where( (length > 0) & (blocks == (length - 1) // GROUP), float("inf"), scores ) - tl.store(SCORES + row * BLOCKS + blocks, scores, blocks < BLOCKS) + tl.store(SCORES + row * SCORE_STRIDE + blocks, scores, blocks < SCORE_STRIDE) @triton.jit @@ -72,14 +75,20 @@ def _publish_candidate_mask_kernel( WIDTH: tl.constexpr, GROUP: tl.constexpr, TOPK: tl.constexpr, + INDEX_STRIDE: tl.constexpr, + VALUE_STRIDE: tl.constexpr, TILE: tl.constexpr, ): row = tl.program_id(0).to(tl.int64) i = tl.program_id(1) * TILE + tl.arange(0, TILE) - selected = tl.load(INDICES + row * TOPK + i // GROUP, i < TOPK * GROUP, 0) - score = tl.load(VALUES + row * TOPK + i // GROUP, i < TOPK * GROUP, -float("inf")) + selected = tl.load(INDICES + row * INDEX_STRIDE + i // GROUP, i < TOPK * GROUP, 0) + score = tl.load( + VALUES + row * VALUE_STRIDE + i // GROUP, + i < TOPK * GROUP, + -float("inf"), + ) cols = selected * GROUP + i % GROUP - # torch.topk returns unique block indices: each output position has one writer. + # Top-K returns unique block indices: each output position has one writer. tl.store( KEEP + row * WIDTH + cols, score > -float("inf"), @@ -95,10 +104,12 @@ def candidate_block_logits( block_size: int, published: torch.Tensor | None, ) -> tuple[torch.Tensor, torch.Tensor | None]: - """Keep torch.topk's block selection, including its tie behavior. + """Select candidate blocks and publish their token-level visibility mask. - Without ``published`` (a source) mask the unread tail while scoring blocks; - with it (a consumer) apply visibility and the published mask in one pass. + A source masks the unread tail while reducing each block. A consumer masks + visibility and the published candidates in one pass, without copying the + capacity-sized logits before each masked_fill. DeepSelect can choose a + different valid subset when finite block scores tie. """ rows, width = logits.shape output = torch.empty((rows, width), dtype=torch.float32, device=logits.device) @@ -117,33 +128,67 @@ def candidate_block_logits( return output, None blocks = triton.cdiv(width, block_size) - scores = torch.empty((rows, blocks), dtype=torch.float32, device=logits.device) + use_deepselect = envs.SGLANG_OPT_DSV41_DEEPSELECT_CANDIDATE_TOPK.get() + deep_select = None + if use_deepselect: + if torch.cuda.get_device_capability(logits.device) != (9, 0): + raise RuntimeError( + "SGLANG_OPT_DSV41_DEEPSELECT_CANDIDATE_TOPK only supports SM90" + ) + try: + import deep_select + except ImportError as exc: + raise RuntimeError( + "SGLANG_OPT_DSV41_DEEPSELECT_CANDIDATE_TOPK requires the " + "deep_select package" + ) from exc + score_alignment = _DEEPSELECT_INPUT_ALIGNMENT_BYTES // torch.float32.itemsize + score_stride = triton.cdiv(blocks, score_alignment) * score_alignment + scores = torch.empty( + (rows, score_stride if use_deepselect else blocks), + dtype=torch.float32, + device=logits.device, + ) group_pad = triton.next_power_of_2(block_size) tile = max(1, 1024 // group_pad) - _candidate_scores_kernel[(rows, triton.cdiv(blocks, tile))]( + _candidate_scores_kernel[ + (rows, triton.cdiv(score_stride if use_deepselect else blocks, tile)) + ]( logits, seq_lens, output, scores, width, logits.stride(0), - blocks, + scores.stride(0), block_size, group_pad, tile, ) # Publication only needs membership; sorting the selected pairs is unused. - top = scores.topk(min(topk_blocks, blocks), dim=-1, sorted=False) + selected = min(topk_blocks, blocks) + if use_deepselect: + top_values, top_indices = deep_select.topk( + scores, + selected, + indices_type=torch.int32, + return_value=True, + ) + else: + top = scores.topk(selected, dim=-1, sorted=False) + top_values, top_indices = top.values, top.indices keep = torch.zeros((rows, width), dtype=torch.bool, device=logits.device) _publish_candidate_mask_kernel[ - (rows, triton.cdiv(top.indices.shape[1] * block_size, 256)) + (rows, triton.cdiv(top_indices.shape[1] * block_size, 256)) ]( - top.indices, - top.values, + top_indices, + top_values, keep, width, block_size, - top.indices.shape[1], + top_indices.shape[1], + top_indices.stride(0), + top_values.stride(0), 256, num_warps=4, ) diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 44be0f1f46ac..73bdade2d4d7 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -1499,6 +1499,8 @@ class Envs: # Run the DeepSeek-V4.1 ratio-1/2 prefill indexer on the torch path instead # of the DeepGEMM dense fp4 logits kernel (test oracle / fallback). SGLANG_DSV41_TORCH_PREFILL_INDEXER = EnvBool(False) + # Use the optional DeepSelect SM90 kernel for candidate-block FP32 Top-K. + SGLANG_OPT_DSV41_DEEPSELECT_CANDIDATE_TOPK = EnvBool(False) # Keep the DeepSeek-V4.1 engram tables in host memory (layout below) and gather # rows from the GPU instead of sharding them over HBM. SGLANG_ENABLE_DSV41_ENGRAM_HOST_TABLE = EnvBool(False) diff --git a/python/sglang/srt/layers/attention/dsv4/indexer.py b/python/sglang/srt/layers/attention/dsv4/indexer.py index 6403dc4876ac..5d951d3b095b 100644 --- a/python/sglang/srt/layers/attention/dsv4/indexer.py +++ b/python/sglang/srt/layers/attention/dsv4/indexer.py @@ -1197,6 +1197,25 @@ def select_candidate_blocks( topk_blocks best-scoring blocks per query. Unreachable positions are already -inf in logits, so an all -inf block means not reachable yet; the block holding the query's newest position is always kept.""" + if ( + envs.SGLANG_OPT_DSV41_DEEPSELECT_CANDIDATE_TOPK.get() + and torch.is_tensor(compress_lens) + and logits.dim() == 2 + ): + from sglang.kernels.ops.attention.dsv4.candidate_blocks import ( + candidate_block_logits, + ) + + _, keep = candidate_block_logits( + logits, + compress_lens.reshape(-1), + topk_blocks=topk_blocks, + block_size=block_size, + published=None, + ) + assert keep is not None + return keep + width = logits.size(-1) scores = F.pad(logits, (0, -width % block_size), value=-torch.inf) scores = scores.unflatten(-1, (-1, block_size)).amax(dim=-1) diff --git a/test/registered/kernel/attention/dsv4/test_deepselect_candidate_blocks.py b/test/registered/kernel/attention/dsv4/test_deepselect_candidate_blocks.py new file mode 100644 index 000000000000..1cc8c1e369e2 --- /dev/null +++ b/test/registered/kernel/attention/dsv4/test_deepselect_candidate_blocks.py @@ -0,0 +1,68 @@ +import unittest + +import torch + +from sglang.kernels.ops.attention.dsv4.candidate_blocks import candidate_block_logits +from sglang.srt.environ import envs +from sglang.srt.layers.attention.dsv4.indexer import select_candidate_blocks +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=20, stage="base-b-kernel-unit", runner_config="1-gpu-large") + + +def _deepselect_sm90_available() -> bool: + if not torch.cuda.is_available() or torch.cuda.get_device_capability() != (9, 0): + return False + try: + import deep_select # noqa: F401 + except ImportError: + return False + return True + + +@unittest.skipUnless( + _deepselect_sm90_available(), "requires the optional DeepSelect SM90 package" +) +class TestDeepSelectCandidateBlocks(CustomTestCase): + def test_candidate_publication_matches_torch(self): + torch.manual_seed(913) + rows, width, block_size, topk_blocks = 6, 131079, 8, 513 + logits = torch.randn(rows, width + 11, device="cuda")[:, :width] + lengths = torch.tensor( + [0, 1, 17, width // 2, width - 3, width], + device="cuda", + dtype=torch.int32, + ) + + with envs.SGLANG_OPT_DSV41_DEEPSELECT_CANDIDATE_TOPK.override(False): + ref_logits, ref_keep = candidate_block_logits( + logits, + lengths, + topk_blocks=topk_blocks, + block_size=block_size, + published=None, + ) + with envs.SGLANG_OPT_DSV41_DEEPSELECT_CANDIDATE_TOPK.override(True): + got_logits, got_keep = candidate_block_logits( + logits, + lengths, + topk_blocks=topk_blocks, + block_size=block_size, + published=None, + ) + + torch.testing.assert_close(got_logits, ref_logits, rtol=0, atol=0) + torch.testing.assert_close(got_keep, ref_keep, rtol=0, atol=0) + with envs.SGLANG_OPT_DSV41_DEEPSELECT_CANDIDATE_TOPK.override(True): + selector_keep = select_candidate_blocks( + ref_logits, + lengths[:, None], + topk_blocks=topk_blocks, + block_size=block_size, + ) + torch.testing.assert_close(selector_keep, ref_keep, rtol=0, atol=0) + + +if __name__ == "__main__": + unittest.main()