Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
77 changes: 61 additions & 16 deletions python/sglang/kernels/ops/attention/dsv4/candidate_blocks.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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,
Expand All @@ -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
Expand Down Expand Up @@ -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"),
Expand All @@ -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)
Expand All @@ -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,
)
Expand Down
2 changes: 2 additions & 0 deletions python/sglang/srt/environ.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
19 changes: 19 additions & 0 deletions python/sglang/srt/layers/attention/dsv4/indexer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Original file line number Diff line number Diff line change
@@ -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()
Loading