From 35b7589e1abcc1d68b487b61c773cdb420229037 Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Wed, 16 Sep 2026 15:15:18 -0700 Subject: [PATCH] dsv4.1: candidate indexer library (#39671) Co-authored-by: DarkSharpness <76582120+DarkSharpness@users.noreply.github.com> Co-authored-by: Xiaoyu Zhang <1182563586@qq.com> --- .../ops/attention/dsv4/candidate_blocks.py | 145 +--------- .../ops/attention/dsv4/indexer_postprocess.py | 74 ----- .../attention/dsv4/candidate_indexer.py | 118 ++++++++ .../dsv4/candidate_indexer_deep_gemm.py | 255 ++++++++++++++++++ .../srt/layers/attention/dsv4/indexer.py | 56 ++++ .../layers/deep_gemm_wrapper/configurer.py | 17 ++ 6 files changed, 447 insertions(+), 218 deletions(-) delete mode 100644 python/sglang/kernels/ops/attention/dsv4/indexer_postprocess.py create mode 100644 python/sglang/srt/layers/attention/dsv4/candidate_indexer.py create mode 100644 python/sglang/srt/layers/attention/dsv4/candidate_indexer_deep_gemm.py diff --git a/python/sglang/kernels/ops/attention/dsv4/candidate_blocks.py b/python/sglang/kernels/ops/attention/dsv4/candidate_blocks.py index a8f0cfe61..b70d0f2d2 100644 --- a/python/sglang/kernels/ops/attention/dsv4/candidate_blocks.py +++ b/python/sglang/kernels/ops/attention/dsv4/candidate_blocks.py @@ -1,4 +1,4 @@ -"""Candidate-block scores and visibility masking for paged indexer logits.""" +"""Per-row candidate block counts and sparse-row lengths for the paged indexer.""" import torch import triton @@ -7,149 +7,6 @@ import triton.language as tl from sglang.kernels.jit.utils import is_arch_support_pdl -@triton.jit -def _maximum_with_nan(a, b): - return tl.maximum(a, b, propagate_nan=tl.PropagateNan.ALL) - - -@triton.jit -def _candidate_scores_kernel( - X, - LENS, - OUT, - SCORES, - WIDTH: tl.constexpr, - STRIDE: tl.constexpr, - BLOCKS: tl.constexpr, - GROUP: tl.constexpr, - GROUP_PAD: tl.constexpr, - TILE: tl.constexpr, -): - row = tl.program_id(0).to(tl.int64) - blocks = tl.program_id(1) * TILE + tl.arange(0, TILE) - offsets = tl.arange(0, GROUP_PAD) - cols = blocks[:, None] * GROUP + offsets[None, :] - length = tl.load(LENS + row) - in_bounds = (cols < WIDTH) & (offsets[None, :] < GROUP) - values = tl.load( - X + row * STRIDE + cols, in_bounds & (cols < length), other=-float("inf") - ).to(tl.float32) - tl.store(OUT + row * WIDTH + cols, values, in_bounds) - scores = tl.reduce(values, axis=1, combine_fn=_maximum_with_nan) - scores = tl.where( - (length > 0) & (blocks == (length - 1) // GROUP), float("inf"), scores - ) - tl.store(SCORES + row * BLOCKS + blocks, scores, blocks < BLOCKS) - - -@triton.jit -def _candidate_mask_kernel( - X, - LENS, - KEEP, - OUT, - WIDTH: tl.constexpr, - STRIDE: tl.constexpr, - KEEP_STRIDE: tl.constexpr, - KEEP_COL_STRIDE: tl.constexpr, - TILE: tl.constexpr, -): - row = tl.program_id(0).to(tl.int64) - cols = tl.program_id(1) * TILE + tl.arange(0, TILE) - visible = (cols < WIDTH) & (cols < tl.load(LENS + row)) - keep = tl.load(KEEP + row * KEEP_STRIDE + cols * KEEP_COL_STRIDE, visible, other=0) - values = tl.load(X + row * STRIDE + cols, visible & keep, other=-float("inf")).to( - tl.float32 - ) - tl.store(OUT + row * WIDTH + cols, values, cols < WIDTH) - - -@triton.jit -def _publish_candidate_mask_kernel( - INDICES, - VALUES, - KEEP, - WIDTH: tl.constexpr, - GROUP: tl.constexpr, - TOPK: 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")) - cols = selected * GROUP + i % GROUP - # torch.topk returns unique block indices: each output position has one writer. - tl.store( - KEEP + row * WIDTH + cols, - score > -float("inf"), - (i < TOPK * GROUP) & (cols < WIDTH), - ) - - -def candidate_block_logits( - logits: torch.Tensor, - seq_lens: torch.Tensor, - *, - topk_blocks: int, - block_size: int, - published: torch.Tensor | None, -) -> tuple[torch.Tensor, torch.Tensor | None]: - """Keep torch.topk's block selection, including its tie behavior. - - Without ``published`` (a source) mask the unread tail while scoring blocks; - with it (a consumer) apply visibility and the published mask in one pass. - """ - rows, width = logits.shape - output = torch.empty((rows, width), dtype=torch.float32, device=logits.device) - if published is not None: - _candidate_mask_kernel[(rows, triton.cdiv(width, 4096))]( - logits, - seq_lens, - published, - output, - width, - logits.stride(0), - published.stride(0), - published.stride(1), - 4096, - ) - return output, None - - blocks = triton.cdiv(width, block_size) - scores = torch.empty((rows, 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))]( - logits, - seq_lens, - output, - scores, - width, - logits.stride(0), - blocks, - 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) - 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)) - ]( - top.indices, - top.values, - keep, - width, - block_size, - top.indices.shape[1], - 256, - num_warps=4, - ) - return output, keep - - @triton.jit def _candidate_row_lens_kernel( LENS, diff --git a/python/sglang/kernels/ops/attention/dsv4/indexer_postprocess.py b/python/sglang/kernels/ops/attention/dsv4/indexer_postprocess.py deleted file mode 100644 index f9983326e..000000000 --- a/python/sglang/kernels/ops/attention/dsv4/indexer_postprocess.py +++ /dev/null @@ -1,74 +0,0 @@ -"""Filter selected indexer scores and map logical positions to KV slots.""" - -import torch -import triton -import triton.language as tl - - -@triton.jit -def _filter_topk_pages( - SCORES, - INDICES, - PAGES, - OUT, - RAW, - WIDTH: tl.constexpr, - TOPK: tl.constexpr, - PAGE_SIZE: tl.constexpr, - SS: tl.constexpr, - SI: tl.constexpr, - SP: tl.constexpr, - SO: tl.constexpr, - SR: tl.constexpr, - WRITE_RAW: tl.constexpr, - BLOCK: tl.constexpr, -): - row = tl.program_id(0).to(tl.int64) - col = tl.program_id(1) * BLOCK + tl.arange(0, BLOCK) - index = tl.load(INDICES + row * SI + col, col < TOPK, -1).to(tl.int64) - in_bounds = (index >= 0) & (index < WIDTH) & (col < TOPK) - score = tl.load(SCORES + row * SS + index, in_bounds, -float("inf")) - # This comparison also rejects NaN; +inf remains a valid score. - valid = in_bounds & (score > -float("inf")) - page = tl.load(PAGES + row * SP + index // PAGE_SIZE, valid, 0) - slot = (page * PAGE_SIZE).to(tl.int64) + index % PAGE_SIZE - tl.store(OUT + row * SO + col, tl.where(valid, slot, -1), col < TOPK) - if WRITE_RAW: - tl.store(RAW + row * SR + col, tl.where(valid, index, -1), col < TOPK) - - -def filter_topk_pages( - scores: torch.Tensor, - indices: torch.Tensor, - page_table: torch.Tensor, - page_indices: torch.Tensor, - page_size: int, - raw_indices: torch.Tensor | None = None, -) -> None: - """Preserve top-k order, write -1 for invalid scores, and map valid slots.""" - rows, topk = indices.shape - assert scores.ndim == page_table.ndim == page_indices.ndim == 2 - assert scores.shape[0] == page_table.shape[0] == page_indices.shape[0] == rows - assert page_indices.shape[1] == topk and scores.shape[1] > 0 - assert page_table.shape[1] * page_size >= scores.shape[1] - assert all(t.stride(1) == 1 for t in (scores, indices, page_table, page_indices)) - if raw_indices is not None: - assert raw_indices.shape == indices.shape and raw_indices.stride(1) == 1 - _filter_topk_pages[(rows, triton.cdiv(topk, 256))]( - scores, - indices, - page_table, - page_indices, - raw_indices, - scores.shape[1], - topk, - page_size, - scores.stride(0), - indices.stride(0), - page_table.stride(0), - page_indices.stride(0), - raw_indices.stride(0) if raw_indices is not None else 0, - raw_indices is not None, - 256, - num_warps=4, - ) diff --git a/python/sglang/srt/layers/attention/dsv4/candidate_indexer.py b/python/sglang/srt/layers/attention/dsv4/candidate_indexer.py new file mode 100644 index 000000000..9ee739a75 --- /dev/null +++ b/python/sglang/srt/layers/attention/dsv4/candidate_indexer.py @@ -0,0 +1,118 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import TYPE_CHECKING, List, Optional, Union + +import torch +import torch.nn.functional as F + +from sglang.srt.layers.attention.dsv4.metadata import PagedIndexerMetadata +from sglang.srt.runtime_context import get_platform + +if TYPE_CHECKING: + from sglang.srt.layers.attention.dsv4.candidate_indexer_deep_gemm import ( + DeepGemmCandidateIndexer, + ) + + +class CandidateMetadata: + """Base of an implementation's published state on + ``DSV4Metadata.candidate_metadata``.""" + + +@dataclass(frozen=True) +class IndexerInputs: + """One index-source layer's operands on the paged fp4 decode path (one query + row per request, or per draft token under verify).""" + + q_fp4: torch.Tensor # [rows, 1, heads, 64] int8, packed fp4 + q_sf: torch.Tensor # [rows, 1, heads] int32, packed ue8m0 + k_cache: torch.Tensor # [pages, page_size, 1, 68] uint8, the layer's index-K pool + weights: torch.Tensor # [rows, heads] bf16/fp32 head weights + metadata: PagedIndexerMetadata # this ratio's lengths, page table and plans + # [rows] int, one request id per query row, the rows of one request + # consecutive (verify: its draft tokens); None = every row its own request + request_ids: Optional[torch.Tensor] = None + + @property + def num_rows(self) -> int: + return self.q_fp4.shape[0] + + +def make_candidate_indexer( + topk_blocks: int, block_size: int +) -> Optional[DeepGemmCandidateIndexer]: + """The paged fp4 decode path's two-level indexer; None on Hopper, whose decode + indexer selects through masks inline.""" + if topk_blocks <= 0 or get_platform().device_sm < 100: + return None + from sglang.srt.layers.deep_gemm_wrapper.configurer import ( + DEEPGEMM_PAGED_SPARSE_MQA_LOGITS, + ) + + if not DEEPGEMM_PAGED_SPARSE_MQA_LOGITS: + raise RuntimeError( + "the candidate indexer needs DeepGEMM's paged sparse MQA logits " + "(sgl-deep-gemm >= 0.2.0 with SGLANG_ENABLE_JIT_DEEPGEMM on)" + ) + from sglang.srt.layers.attention.dsv4.candidate_indexer_deep_gemm import ( + DeepGemmCandidateIndexer, + ) + + return DeepGemmCandidateIndexer(topk_blocks, block_size) + + +# TODO(candidate): Hopper decode and prefill still select through these masks +# inline in the backend; move them behind the protocol as publish/select_prefill. +@dataclass +class CandidateMasks(CandidateMetadata): + mask: Optional[torch.Tensor] = None # decode: [rows, width] bool + request_masks: Optional[List[torch.Tensor]] = None # prefill: [rows_b, lc_b] each + + +def published_masks(candidate) -> CandidateMasks: + assert isinstance(candidate, CandidateMasks), "candidate masks missing" + return candidate + + +def mask_topk_scores( + scores: torch.Tensor, + indices: torch.Tensor, + offsets: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Keep masked indexer scores out of attention even when top-k underfills.""" + columns = indices.to(torch.int64) + if offsets is not None: + columns = columns - offsets[:, None] + selected_scores = scores.gather(1, columns.clamp(0, scores.shape[1] - 1)) + valid = ( + (columns >= 0) & (columns < scores.shape[1]) & (selected_scores > -torch.inf) + ) + return indices.masked_fill(~valid, -1) + + +def select_candidate_blocks( + logits: torch.Tensor, + compress_lens: Union[torch.Tensor, int], + topk_blocks: int, + block_size: int, +) -> torch.Tensor: + """Level one of the two-level top-k: a bool mask over positions keeping the + 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.""" + 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) + num_blocks = scores.size(-1) + + last = (compress_lens - 1) // block_size + scores = scores.masked_fill( + torch.arange(num_blocks, device=logits.device) == last, torch.inf + ) + + top = scores.topk(min(topk_blocks, num_blocks), dim=-1) + keep = torch.zeros_like(scores, dtype=torch.bool).scatter_( + -1, top.indices, top.values > -torch.inf + ) + return keep.repeat_interleave(block_size, dim=-1)[..., :width] diff --git a/python/sglang/srt/layers/attention/dsv4/candidate_indexer_deep_gemm.py b/python/sglang/srt/layers/attention/dsv4/candidate_indexer_deep_gemm.py new file mode 100644 index 000000000..41d4dd591 --- /dev/null +++ b/python/sglang/srt/layers/attention/dsv4/candidate_indexer_deep_gemm.py @@ -0,0 +1,255 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Optional + +import torch + +from sglang.kernels.ops.attention.dsv4.candidate_blocks import candidate_row_lens +from sglang.kernels.ops.attention.dsv4.candidate_table import ( + amax8_varlen, + sort_candidate_blocks, +) +from sglang.kernels.ops.attention.dsv4.topk import ( + plan_topk_v2, + topk_transform_bf16_small, + topk_transform_paged_v2, +) +from sglang.srt.layers.attention.dsv4.candidate_indexer import ( + CandidateMetadata, + IndexerInputs, +) +from sglang.srt.layers.attention.dsv4.indexer import ( + deep_gemm_fp4_paged_mqa_logits, +) + +CANDIDATE_BLOCK_SIZE = 8 # positions per block; DeepGEMM accepts 8 or 16 + + +@dataclass +class SparseBlockTable(CandidateMetadata): + # [rows, topk_blocks] int32: ascending logical block ids, valid for the first + # min(topk_blocks, ceil(seq_len / 8)) entries of a row; DeepGEMM reads only those + blocks: torch.Tensor + # DeepGEMM's schedule metadata (uint8) for them + schedule: torch.Tensor + # [rows, topk_blocks] int32: the same blocks as pool slots / 8, so a consumer's + # top-k maps column j of the sparse row to slot phys_blocks[b, j // 8] * 8 + j % 8 + # with the plain page-table transform at page size 8 + phys_blocks: torch.Tensor + # [rows] int32: length of each row of the sparse logits: the published blocks + # laid out block by block, the newest possibly partial (`candidate_row_lens`) + valid_lens: torch.Tensor + # recorded on the side stream once the fields above are complete + ready: torch.cuda.Event + + +def amax_topk_blocks( + logits: torch.Tensor, + seq_lens: torch.Tensor, + nblocks: torch.Tensor, + topk_blocks: int, + max_seq_len: Optional[int] = None, +) -> torch.Tensor: + """Per row the ``topk_blocks`` blocks of 8 positions with the largest block + maximum among its first ``seq_lens[b]`` positions, the newest block always + included: block ids in no particular order, ``-1`` past the row's count. + ``nblocks`` is ``ceil(seq_lens / 8)`` as int32.""" + rows = logits.shape[0] + block = CANDIDATE_BLOCK_SIZE + if max_seq_len is None: + max_seq_len = logits.shape[1] + # NOTE: plan cannot be the previous kernel of topk_transform_paged_v2 + plan = plan_topk_v2(nblocks) + # block maxima, the newest block +inf; the top-k reads each row up to nblocks + # only, so nothing past a row's keys is initialised (v2 needs stride % 4 == 0) + keys = logits.new_empty(rows, -(-max_seq_len // (4 * block)) * 4) + amax8_varlen(logits, seq_lens, out=keys) + blocks = torch.empty(rows, topk_blocks, dtype=torch.int32, device=logits.device) + topk_transform_paged_v2(keys, nblocks, None, blocks, 1, plan) + return blocks + + +def build_sparse_indexer_schedule( + blocks: torch.Tensor, + seq_lens: torch.Tensor, + page_table: torch.Tensor, + page_size: int, + q_dtype: torch.dtype, + request_ids: torch.Tensor, +) -> torch.Tensor: + """DeepGEMM's schedule for the published blocks: ``seq_lens`` ``[rows]`` + int32, ``page_table`` ``[rows, pages]`` int32 at the index pool's page size. + ``request_ids`` ``[rows]`` int32 lets DeepGEMM pair two rows of a request on + one KV pass; each row keeps its own block list and output layout, and paired + rows must share their page-table row.""" + import deep_gemm + + return deep_gemm.get_paged_sparse_mqa_logits_metadata( + seq_lens.contiguous(), + page_table, + request_ids, + page_size, + blocks, + q_dtype, + CANDIDATE_BLOCK_SIZE, + ) + + +def sparse_logits( + q_fp4: torch.Tensor, + q_sf: torch.Tensor, + k_cache: torch.Tensor, + weights: torch.Tensor, + table: SparseBlockTable, +) -> torch.Tensor: + """bf16 logits ``[rows, topk_blocks * 8]`` of the published blocks: ``q_fp4`` + ``[rows, 1, heads, 64]`` int8 with ``q_sf`` ``[rows, 1, heads]`` int32 (packed + ue8m0), ``k_cache`` ``[pages, page_size, 1, 68]`` uint8 whose page stride is + a multiple of 512 bytes, ``weights`` ``[rows, heads]`` bf16.""" + import deep_gemm + + return deep_gemm.fp8_fp4_paged_sparse_mqa_logits( + (q_fp4, q_sf), + k_cache, + weights, + table.schedule, + table.blocks.shape[1], + CANDIDATE_BLOCK_SIZE, + ) + + +def topk_transform_sparse( + logits: torch.Tensor, + valid_lens: torch.Tensor, + table: SparseBlockTable, + page_indices: torch.Tensor, +) -> None: + """Top-``k`` (``k = page_indices.shape[1]``) of every row of the sparse + ``logits`` (bf16 ``[rows, topk_blocks * 8]``) within its first ``valid_lens[b]`` + columns, written as pool slots, ``-1`` where a row has fewer than ``k`` valid + columns, in no particular order.""" + topk_transform_bf16_small( + logits, valid_lens, table.phys_blocks, page_indices, CANDIDATE_BLOCK_SIZE + ) + + +# TODO(dark): support publish prefill/select prefill +# TODO(dark): support fusion of publish + topk of publish layer +class DeepGemmCandidateIndexer: + def __init__(self, topk_blocks: int, block_size: int): + assert block_size == CANDIDATE_BLOCK_SIZE, block_size + self.topk_blocks = topk_blocks + self.block_size = block_size + self.alt_stream = torch.cuda.Stream() + self._row_ids: Optional[torch.Tensor] = None + self._retired_row_ids: list = [] # captured graphs keep reading the buffers they saw + + def _request_ids( + self, request_ids: Optional[torch.Tensor], rows: int, device: torch.device + ) -> torch.Tensor: + """int32 ``[rows]``; None means every row is its own request, served from a + cached ``arange`` (grown in steps of 8192) so that case costs no launch.""" + if request_ids is not None: + # the scheduler keeps request indices as int64; one small cast per publish + return request_ids[:rows].to(torch.int32).contiguous() + buf = self._row_ids + if buf is None or buf.numel() < rows: + assert not torch.cuda.is_current_stream_capturing(), ( + f"row-id buffer grows to {rows} rows inside a CUDA graph capture; " + "warm up with the largest row count first" + ) + if buf is not None: + self._retired_row_ids.append(buf) + size = max(8192, -(-rows // 8192) * 8192) + buf = self._row_ids = torch.arange(size, dtype=torch.int32, device=device) + return buf[:rows] + + def publish_decode( + self, + inputs: IndexerInputs, + page_indices: torch.Tensor, + raw_indices: Optional[torch.Tensor] = None, + ) -> SparseBlockTable: + """The publishing layer's own plain top-k into ``page_indices``, plus the + block table for the consumers; the backend stores it on the forward + metadata.""" + metadata = inputs.metadata + seq_lens = metadata.compressed_seq_lens.reshape(-1) + logits = deep_gemm_fp4_paged_mqa_logits( + (inputs.q_fp4, inputs.q_sf), + inputs.k_cache, + inputs.weights, + metadata.compressed_seq_lens, + metadata.page_table, + metadata.deep_gemm_metadata, + metadata.max_compressed_seq_len, + ) + main_stream = torch.cuda.current_stream() + self.alt_stream.wait_stream(main_stream) + # The block-selection chain reads logits after the main stream moves on. + logits.record_stream(self.alt_stream) + # TODO(candidate): one kernel for both selections below (dense logits read once) + topk_transform_paged_v2( + logits, + seq_lens, + metadata.page_table, + page_indices, + metadata.compressed_page_size, + metadata.topk_metadata, + out_raw_indices=raw_indices, + ) + # per row: block count for the block top-k, sparse-row length for the consumers + with torch.cuda.stream(self.alt_stream): + nblocks, row_valid_lens = candidate_row_lens(seq_lens, self.topk_blocks) + blocks = amax_topk_blocks(logits, seq_lens, nblocks, self.topk_blocks) + # in place: ascending, INT32_MAX padded, plus the blocks as pool slots / 8 + phys_blocks = sort_candidate_blocks( + blocks, + seq_lens, + metadata.page_table, + metadata.compressed_page_size, + ) + schedule = build_sparse_indexer_schedule( + blocks, + seq_lens, + metadata.page_table, + metadata.compressed_page_size, + inputs.q_fp4.dtype, + self._request_ids(inputs.request_ids, inputs.num_rows, blocks.device), + ) + # select_decode reads these on the main stream + for t in (blocks, schedule, phys_blocks, row_valid_lens): + t.record_stream(main_stream) + ready = torch.cuda.Event() + ready.record(self.alt_stream) + return SparseBlockTable( + blocks=blocks, + schedule=schedule, + phys_blocks=phys_blocks, + valid_lens=row_valid_lens, + ready=ready, + ) + + def _scores(self, table: SparseBlockTable, inputs: IndexerInputs) -> torch.Tensor: + return sparse_logits( + inputs.q_fp4, + inputs.q_sf, + inputs.k_cache, + inputs.weights.to(torch.bfloat16), + table, + ) + + def select_decode( + self, + candidate_metadata: SparseBlockTable, + inputs: IndexerInputs, + page_indices: torch.Tensor, + raw_indices: Optional[torch.Tensor] = None, + ) -> None: + assert raw_indices is None + table = candidate_metadata + torch.cuda.current_stream().wait_event(table.ready) + logits = self._scores(table, inputs) + # decode carries no raw_indices; the kernel writes slots only + topk_transform_sparse(logits, table.valid_lens, table, page_indices) diff --git a/python/sglang/srt/layers/attention/dsv4/indexer.py b/python/sglang/srt/layers/attention/dsv4/indexer.py index a778b2307..77416f4e7 100644 --- a/python/sglang/srt/layers/attention/dsv4/indexer.py +++ b/python/sglang/srt/layers/attention/dsv4/indexer.py @@ -439,6 +439,62 @@ def topk_transform_flashinfer_fused( ) +def deep_gemm_fp4_paged_mqa_logits( + q_fp4: Tuple[torch.Tensor, torch.Tensor], + k_cache: torch.Tensor, + weights: torch.Tensor, + seq_lens: torch.Tensor, + page_table: torch.Tensor, + deep_gemm_metadata, + max_seq_len: int, +) -> torch.Tensor: + """DeepGEMM paged fp4 logits; no hadamard, the reference does not apply one.""" + from deep_gemm import fp8_fp4_paged_mqa_logits + + sl = seq_lens.to(torch.int32) + if sl.dim() == 1: + sl = sl.unsqueeze(-1) + return fp8_fp4_paged_mqa_logits( + q_fp4, + k_cache, + weights, + sl, + page_table, + deep_gemm_metadata, + max_seq_len, + False, + ) + + +def topk_transform_paged_from_metadata( + logits: torch.Tensor, + metadata, + page_indices: torch.Tensor, + raw_indices: Optional[torch.Tensor] = None, +) -> None: + """Pool slots into ``page_indices`` (``-1`` past the valid count) and, when given, + positions into ``raw_indices``; ``metadata`` is a ``PagedIndexerMetadata``.""" + if metadata.use_topk_v2: + topk_transform_paged_v2( + logits, + metadata.compressed_seq_lens, + metadata.page_table, + page_indices, + metadata.compressed_page_size, + metadata.topk_metadata, + raw_indices, + ) + else: + topk_transform_paged( + logits, + metadata.compressed_seq_lens, + metadata.page_table, + page_indices, + metadata.compressed_page_size, + raw_indices, + ) + + class C4IndexerBackendMixin: def __init__(self): super().__init__() diff --git a/python/sglang/srt/layers/deep_gemm_wrapper/configurer.py b/python/sglang/srt/layers/deep_gemm_wrapper/configurer.py index a21073f5c..65261e33c 100644 --- a/python/sglang/srt/layers/deep_gemm_wrapper/configurer.py +++ b/python/sglang/srt/layers/deep_gemm_wrapper/configurer.py @@ -43,3 +43,20 @@ DEEPGEMM_SCALE_UE8M0 = ENABLE_JIT_DEEPGEMM and ( get_platform().is_sm100 or get_device_sm() == 120 ) DEEPGEMM_NEED_TMA_ALIGNED_SCALES = not (DEEPGEMM_SCALE_UE8M0 or _is_musa) + + +def _supports_paged_sparse_mqa_logits() -> bool: + if not DEEPGEMM_BLACKWELL: + return False + import deep_gemm + + return all( + callable(getattr(deep_gemm, name, None)) + for name in ( + "get_paged_sparse_mqa_logits_metadata", + "fp8_fp4_paged_sparse_mqa_logits", + ) + ) + + +DEEPGEMM_PAGED_SPARSE_MQA_LOGITS = _supports_paged_sparse_mqa_logits()