Skip to content
Merged
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
164 changes: 162 additions & 2 deletions python/sglang/kernels/ops/attention/dsv4/candidate_blocks.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,24 @@
"""Per-row candidate block counts and sparse-row lengths for the paged indexer."""
"""Candidate block selection of the two-level indexer: the per-row block counts
and sparse-row lengths, the block keys and the JIT block top-k, and the torch
block ids and masks."""

from typing import Optional, Union

import torch
import torch.nn.functional as F
import triton
import triton.language as tl

from sglang.kernels.jit.utils import is_arch_support_pdl
from sglang.kernels.jit.utils import (
cache_once,
is_arch_support_pdl,
load_jit,
make_cpp_args,
)

from .candidate_table import CANDIDATE_BLOCK_SIZE
from .topk import plan_topk_v2, topk_transform_paged_v2
from .utils import make_name


@triton.jit
Expand Down Expand Up @@ -62,3 +76,149 @@ def candidate_row_lens(
**pdl_kwargs,
)
return nblocks, valid


@cache_once
def _jit_block_amax_module():
args = make_cpp_args(is_arch_support_pdl())
return load_jit(
make_name("block_amax"),
*args,
cuda_files=["deepseek_v4/block_amax.cuh"],
cuda_wrappers=[("amax8_varlen", f"BlockAmaxKernel<{args}>::amax8_varlen")],
)


def amax8_varlen(
scores: torch.Tensor,
seq_lens: torch.Tensor,
topk: int = 0,
*,
max_seqlen: int = 0,
out: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Level-one keys of the two-level indexer: ``out[b, i]`` is the max of
``scores[b, 8 i : 8 i + 8]`` for ``i < ceil(seq_lens[b] / 8)``, the last of
them ``+inf`` (the newest block is always selected), nothing written past
that count. Rows with at most ``topk`` blocks are skipped (every block is
selected anyway); ``topk=0`` never skips. ``out`` is allocated as
``[rows, ceil(max_seqlen / 8)]`` when not given, ``max_seqlen`` defaulting to
the width of ``scores``; every ``seq_lens[b]`` must fit in ``8 * out.shape[1]``.
fp32 only for now; ``scores`` rows must be 32-byte aligned (stride a multiple
of 8). Returns ``out``.
"""
if out is None:
num_tokens, max_len = scores.shape
if max_seqlen == 0:
max_seqlen = max_len
out = scores.new_empty(num_tokens, (max_seqlen + 7) // 8)
_jit_block_amax_module().amax8_varlen(scores, seq_lens, out, topk)
return out


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 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 _candidate_block_topk(
logits: torch.Tensor,
compress_lens: Union[torch.Tensor, int],
topk_blocks: int,
block_size: int,
) -> torch.return_types.topk:
width = logits.size(-1)
padding = -width % block_size
scores = F.pad(logits, (0, padding), value=-torch.inf) if padding else logits
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
)

return scores.topk(min(topk_blocks, num_blocks), dim=-1)


def select_candidate_block_ids(
logits: torch.Tensor,
compress_lens: Union[torch.Tensor, int],
topk_blocks: int,
block_size: int,
) -> torch.Tensor:
top = _candidate_block_topk(
logits=logits,
compress_lens=compress_lens,
topk_blocks=topk_blocks,
block_size=block_size,
)
return top.indices.to(torch.int32).masked_fill_(~(top.values > -torch.inf), -1)


def candidate_block_mask(
blocks: torch.Tensor, width: int, block_size: int
) -> torch.Tensor:
num_blocks = (width + block_size - 1) // block_size
keep = torch.zeros(
(*blocks.shape[:-1], num_blocks + 1), dtype=torch.bool, device=blocks.device
)
keep.scatter_(-1, blocks.to(torch.int64).masked_fill(blocks < 0, num_blocks), True)
return keep[..., :num_blocks].repeat_interleave(block_size, dim=-1)[..., :width]


def select_candidate_blocks(
logits: torch.Tensor,
compress_lens: Union[torch.Tensor, int],
topk_blocks: int,
block_size: int,
) -> torch.Tensor:
top = _candidate_block_topk(
logits=logits,
compress_lens=compress_lens,
topk_blocks=topk_blocks,
block_size=block_size,
)
width = logits.shape[-1]
num_blocks = (width + block_size - 1) // block_size
keep = torch.zeros(
(*logits.shape[:-1], num_blocks), dtype=torch.bool, device=logits.device
).scatter_(-1, top.indices, top.values > -torch.inf)
return keep.repeat_interleave(block_size, dim=-1)[..., :width]
68 changes: 29 additions & 39 deletions python/sglang/kernels/ops/attention/dsv4/candidate_table.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
"""Candidate-table kernels of the two-level indexer: level-one block keys and
the sorted, page-transformed block table."""
"""The sparse table of the two-level indexer from a row's selected blocks: the
sorted, page-transformed block table and DeepGEMM's schedule for it."""

from __future__ import annotations

Expand All @@ -19,43 +19,7 @@
if TYPE_CHECKING:
pass


@cache_once
def _jit_block_amax_module():
args = make_cpp_args(is_arch_support_pdl())
return load_jit(
make_name("block_amax"),
*args,
cuda_files=["deepseek_v4/block_amax.cuh"],
cuda_wrappers=[("amax8_varlen", f"BlockAmaxKernel<{args}>::amax8_varlen")],
)


def amax8_varlen(
scores: torch.Tensor,
seq_lens: torch.Tensor,
topk: int = 0,
*,
max_seqlen: int = 0,
out: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Level-one keys of the two-level indexer: ``out[b, i]`` is the max of
``scores[b, 8 i : 8 i + 8]`` for ``i < ceil(seq_lens[b] / 8)``, the last of
them ``+inf`` (the newest block is always selected), nothing written past
that count. Rows with at most ``topk`` blocks are skipped (every block is
selected anyway); ``topk=0`` never skips. ``out`` is allocated as
``[rows, ceil(max_seqlen / 8)]`` when not given, ``max_seqlen`` defaulting to
the width of ``scores``; every ``seq_lens[b]`` must fit in ``8 * out.shape[1]``.
fp32 only for now; ``scores`` rows must be 32-byte aligned (stride a multiple
of 8). Returns ``out``.
"""
if out is None:
num_tokens, max_len = scores.shape
if max_seqlen == 0:
max_seqlen = max_len
out = scores.new_empty(num_tokens, (max_seqlen + 7) // 8)
_jit_block_amax_module().amax8_varlen(scores, seq_lens, out, topk)
return out
CANDIDATE_BLOCK_SIZE = 8 # positions per block; DeepGEMM accepts 8 or 16


@cache_once
Expand Down Expand Up @@ -91,3 +55,29 @@ def sort_candidate_blocks(
blocks, seq_lens, page_table, out_pages, page_size
)
return out_pages


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,
)
126 changes: 126 additions & 0 deletions python/sglang/kernels/ops/attention/dsv4/index_logits.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,126 @@
"""The index logits of the DeepSeek V4.1 ratio-1/2 index layers, by the layout of
the index K they read: the flattened K of a prefill chunk (tile by tile under a
memory budget), the paged pool, or the published blocks of a sparse table."""

from __future__ import annotations

from typing import Iterator, Tuple

import torch

from sglang.srt.layers.attention.mqa_logits_utils import (
mqa_logits_row_bytes,
mqa_logits_rows_per_chunk,
)
from sglang.srt.utils.common import ceil_align

from .candidate_table import CANDIDATE_BLOCK_SIZE


def flat_index_logits_rows_per_tile(
rows: int, width: int, *, heads: int, budget_bytes: int
) -> int:
"""Rows per fp32 logits tile within ``budget_bytes``, at the kernel's row
alignment."""
row_alignment = 128 // heads
rows_per_chunk = mqa_logits_rows_per_chunk(
num_rows=ceil_align(rows, row_alignment),
row_bytes=mqa_logits_row_bytes(width),
budget_bytes=budget_bytes,
)
if rows_per_chunk is None:
return rows
return max(row_alignment, rows_per_chunk // row_alignment * row_alignment)


def flat_index_logits_tiles(
*,
q: tuple[torch.Tensor, torch.Tensor],
kv: tuple[torch.Tensor, torch.Tensor],
weights: torch.Tensor,
starts: torch.Tensor,
lengths: torch.Tensor,
context_lengths: list[int],
budget_bytes: int,
width_align: int = 4,
) -> Iterator[tuple[slice, torch.Tensor]]:
"""``(rows, logits)`` per row tile, each fp32 tile within ``budget_bytes``:
``logits[i, j]`` scores query row ``rows.start + i`` against ``kv[starts + j]``,
garbage past the row's ``lengths``; the width is ``max(context_lengths)``
aligned to ``width_align``."""
from deep_gemm import fp8_fp4_mqa_logits

rows = q[0].shape[0]
width = ceil_align(max(context_lengths, default=0), width_align)
if rows == 0 or width == 0:
return
rows_per_chunk = flat_index_logits_rows_per_tile(
rows, width, heads=q[0].shape[1], budget_bytes=budget_bytes
)
for offset in range(0, rows, rows_per_chunk):
tile = slice(offset, min(offset + rows_per_chunk, rows))
tile_starts = starts[tile]
yield (
tile,
fp8_fp4_mqa_logits(
(q[0][tile], q[1][tile]),
kv,
weights[tile],
tile_starts,
tile_starts + lengths[tile],
False,
width,
),
)


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 sparse_logits(
q_fp4: torch.Tensor,
q_sf: torch.Tensor,
k_cache: torch.Tensor,
weights: torch.Tensor,
schedule: torch.Tensor,
topk_blocks: int,
) -> 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, ``schedule`` the
``build_sparse_indexer_schedule`` of the ``topk_blocks`` published blocks."""
import deep_gemm

return deep_gemm.fp8_fp4_paged_sparse_mqa_logits(
(q_fp4, q_sf),
k_cache,
weights,
schedule,
topk_blocks,
CANDIDATE_BLOCK_SIZE,
)
Loading
Loading