94057c3d3e
PR Test (NPU) / check-changes (push) Has been cancelled
PR Test (NPU) / pr-gate (push) Has been cancelled
PR Test (NPU) / set-image-config (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-4-npu-a3 (push) Has been cancelled
PR Test (NPU) / stage-b-test-16-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-1-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-2-npu-a3 (push) Has been cancelled
PR Test (Arm64) / pr-gate (push) Has been cancelled
PR Test (Arm64) / check-changes (push) Has been cancelled
PR Test (Arm64) / build-test (push) Has been cancelled
PR Test (sgl-router) / gate (push) Has been cancelled
PR Test (sgl-router) / tier-1 — lint (push) Has been cancelled
PR Test (sgl-router) / tier-2 — build + test (push) Has been cancelled
PR Test (sgl-router) / tier-3 — docker (placeholder) (push) Has been cancelled
PR Test (sgl-router) / tier-3 — k8s integration (push) Has been cancelled
PR Test (sgl-router) / tier-3 — e2e (push) Has been cancelled
PR Test (sgl-router) / finish (push) Has been cancelled
PR Test (NPU) / single-node-poc (map[name:qwen3_6_27b_w8a8_1p_in64k_out1k_50ms runner:linux-aarch64-a3-2 test_case:test/registered/ascend/performance/qwen3_6_27b/test_npu_qwen3_6_27b_w8a8_1p_in64k_out1k_50ms.py test_type:perf]) (push) Has been cancelled
PR Test (NPU) / pr-test-npu-finish (push) Has been cancelled
PR Test (Xeon) / pr-gate (push) Has been cancelled
PR Test (Xeon) / check-changes (push) Has been cancelled
PR Test (Xeon) / build-test (, xeon-gnr, base-b-test-cpu) (push) Has been cancelled
PR Test (XPU) / check-changes (push) Has been cancelled
PR Test (XPU) / pr-gate (push) Has been cancelled
PR Test (XPU) / stage-a-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / wait-for-stage-a (push) Has been cancelled
PR Test (XPU) / stage-b-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / finish (push) Has been cancelled
CI Model Inventory / build-inventory (push) Has been cancelled
Lint / lint (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Compilation Check (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Manual Policy (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Request Processing (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Summary (push) Has been cancelled
PR Test (SMG) / build-wheel (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on windows (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (x86_64 - auto) (push) Has been cancelled
PR Test (SMG) / python-unit-tests (push) Has been cancelled
PR Test (SMG) / unit-tests (push) Has been cancelled
PR Test (SMG) / benchmarks (push) Has been cancelled
PR Test (SMG) / chat-completions (push) Has been cancelled
PR Test (SMG) / chat-completions-4gpu (push) Has been cancelled
PR Test (SMG) / e2e (push) Has been cancelled
PR Test (SMG) / docker-build-test (push) Has been cancelled
PR Test (SMG) / k8s-integration (push) Has been cancelled
PR Test (SMG) / finish (push) Has been cancelled
PR Test (SMG) / summarize-benchmarks (push) Has been cancelled
Release SGLang Model Gateway Docker Image / publish (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Build SDist (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Upload to PyPI (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (aarch64, 12.9, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (x86_64, 12.9, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu129 (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (aarch64, 13.0, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (x86_64, 13.0, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu130 (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 700) (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 720) (push) Has been cancelled
Release SGLang Kernels / release-rocm700 (push) Has been cancelled
Release SGLang Kernels / release-rocm720 (push) Has been cancelled
Release SGLang Kernels / build-musa43 (43, 3.10) (push) Has been cancelled
Release SGLang Kernels / release-musa43 (push) Has been cancelled
368 lines
15 KiB
Python
368 lines
15 KiB
Python
from __future__ import annotations
|
|
|
|
from enum import Enum, IntEnum, auto
|
|
from typing import Callable, Dict, List, Optional, Tuple
|
|
|
|
import torch
|
|
|
|
from sglang.srt.environ import envs
|
|
|
|
_FLASHINFER_TIE_BREAK_VALUES = {
|
|
"small": 1,
|
|
"large": 2,
|
|
}
|
|
|
|
|
|
class TopkTransformMethod(IntEnum):
|
|
# Transform topk indices to indices to the page table (page_size = 1)
|
|
PAGED = auto()
|
|
# Transform topk indices to indices to ragged kv (non-paged)
|
|
RAGGED = auto()
|
|
|
|
|
|
class DSATopKBackend(Enum):
|
|
SGL_KERNEL = "sgl-kernel"
|
|
TORCH = "torch"
|
|
FLASHINFER = "flashinfer"
|
|
|
|
def is_sgl_kernel(self) -> bool:
|
|
return self == DSATopKBackend.SGL_KERNEL
|
|
|
|
def is_torch(self) -> bool:
|
|
return self == DSATopKBackend.TORCH
|
|
|
|
def is_flashinfer(self) -> bool:
|
|
return self == DSATopKBackend.FLASHINFER
|
|
|
|
def topk_func(
|
|
self,
|
|
score: torch.Tensor,
|
|
lengths: torch.Tensor,
|
|
topk: int,
|
|
row_starts: Optional[torch.Tensor] = None,
|
|
) -> torch.Tensor:
|
|
if self.is_sgl_kernel():
|
|
from sgl_kernel import fast_topk_v2
|
|
|
|
return fast_topk_v2(score, lengths, topk, row_starts=row_starts)
|
|
if self.is_torch():
|
|
return _topk_unfused(
|
|
score,
|
|
lengths,
|
|
topk,
|
|
row_starts=row_starts,
|
|
topk_op=torch.topk,
|
|
topk_op_kwargs={"dim": -1},
|
|
)
|
|
if self.is_flashinfer():
|
|
import flashinfer
|
|
|
|
return _topk_unfused(
|
|
score,
|
|
lengths,
|
|
topk,
|
|
row_starts=row_starts,
|
|
topk_op=flashinfer.top_k,
|
|
topk_op_kwargs={
|
|
"sorted": False,
|
|
"deterministic": envs.SGLANG_DSA_TOPK_FLASHINFER_DETERMINISTIC.get(),
|
|
"tie_break": _flashinfer_tie_break_value(),
|
|
"dsa_graph_safe": True,
|
|
},
|
|
)
|
|
raise RuntimeError(f"Unsupported {self = }.")
|
|
|
|
def topk_transform(
|
|
self,
|
|
logits: torch.Tensor,
|
|
lengths: torch.Tensor,
|
|
topk: int,
|
|
topk_transform_method: TopkTransformMethod,
|
|
attn_metadata,
|
|
cu_seqlens_q_topk: Optional[torch.Tensor] = None,
|
|
topk_indices_offset: Optional[torch.Tensor] = None,
|
|
row_starts: Optional[torch.Tensor] = None,
|
|
batch_idx_list: Optional[List[int]] = None,
|
|
force_unfused_topk: bool = False,
|
|
) -> torch.Tensor:
|
|
if not envs.SGLANG_DSA_FUSE_TOPK.get() or force_unfused_topk:
|
|
return self.topk_func(logits, lengths, topk, row_starts=row_starts)
|
|
|
|
# Decode-shaped PAGED top-k (plain decode AND spec verify / draft-extend,
|
|
# whose expanded rows match the same shape) routes to the DeepSeek-V4 top-k
|
|
# v2 JIT kernel, which fuses top-k selection and the page-table transform in
|
|
# one launch and consumes the indexer's own page_size>=1 table directly, so
|
|
# no page_size=1 table is materialized. Shared by DeepSeek-V3.2 and GLM DSA.
|
|
# This is a deterministic dispatch on the work shape, not a best-effort
|
|
# attempt: the fused-decode CUDA graph drops the page_size=1 table for
|
|
# exactly this case (see dsa_drop_wide_page_table), so once the shape
|
|
# matches we commit to v2 and never silently fall back to the legacy
|
|
# page_size=1 path from here.
|
|
if (
|
|
envs.SGLANG_OPT_USE_TOPK_V2.get()
|
|
and topk_transform_method == TopkTransformMethod.PAGED
|
|
and row_starts is None
|
|
and batch_idx_list is None
|
|
and 0 < topk <= 2048
|
|
and lengths.shape[0]
|
|
== logits.shape[0]
|
|
== attn_metadata.real_page_table.shape[0]
|
|
):
|
|
return _topk_transform_v2_paged(logits, lengths, topk, attn_metadata)
|
|
|
|
# The legacy transforms below read attn_metadata.page_table_1 (page_size=1),
|
|
# which is always present here: the fold only drops it for the decode case
|
|
# dispatched to v2 above.
|
|
assert attn_metadata.page_table_1 is not None
|
|
|
|
if self.is_sgl_kernel():
|
|
from sgl_kernel import (
|
|
fast_topk_transform_fused,
|
|
fast_topk_transform_ragged_fused,
|
|
)
|
|
|
|
if topk_transform_method == TopkTransformMethod.PAGED:
|
|
page_table_size_1 = (
|
|
attn_metadata.page_table_1[batch_idx_list]
|
|
if batch_idx_list is not None
|
|
else attn_metadata.page_table_1
|
|
)
|
|
return fast_topk_transform_fused(
|
|
score=logits,
|
|
lengths=lengths,
|
|
page_table_size_1=page_table_size_1,
|
|
cu_seqlens_q=cu_seqlens_q_topk,
|
|
topk=topk,
|
|
row_starts=row_starts,
|
|
)
|
|
if topk_transform_method == TopkTransformMethod.RAGGED:
|
|
if topk_indices_offset is None:
|
|
raise RuntimeError(
|
|
"RAGGED topk_transform requires topk_indices_offset; "
|
|
"expected extend-without-speculative metadata."
|
|
)
|
|
return fast_topk_transform_ragged_fused(
|
|
score=logits,
|
|
lengths=lengths,
|
|
topk_indices_offset=topk_indices_offset,
|
|
topk=topk,
|
|
row_starts=row_starts,
|
|
)
|
|
raise RuntimeError(f"Unsupported {topk_transform_method = }.")
|
|
|
|
if self.is_flashinfer():
|
|
import flashinfer
|
|
|
|
if topk_transform_method == TopkTransformMethod.PAGED:
|
|
row_to_batch, local_row_starts = _build_flashinfer_paged_args(
|
|
attn_metadata=attn_metadata,
|
|
row_starts=row_starts,
|
|
cu_seqlens_q_topk=cu_seqlens_q_topk,
|
|
batch_idx_list=batch_idx_list,
|
|
device=logits.device,
|
|
num_rows=logits.shape[0],
|
|
)
|
|
return flashinfer.top_k_page_table_transform(
|
|
logits.contiguous(),
|
|
attn_metadata.page_table_1.contiguous(),
|
|
lengths.contiguous(),
|
|
topk,
|
|
row_to_batch=row_to_batch,
|
|
deterministic=envs.SGLANG_DSA_TOPK_FLASHINFER_DETERMINISTIC.get(),
|
|
tie_break=_flashinfer_tie_break_value(),
|
|
dsa_graph_safe=True,
|
|
row_starts=local_row_starts,
|
|
)
|
|
if topk_transform_method == TopkTransformMethod.RAGGED:
|
|
if topk_indices_offset is None:
|
|
raise RuntimeError(
|
|
"RAGGED topk_transform requires topk_indices_offset; "
|
|
"expected extend-without-speculative metadata."
|
|
)
|
|
return flashinfer.top_k_ragged_transform(
|
|
logits.contiguous(),
|
|
topk_indices_offset.contiguous(),
|
|
lengths.contiguous(),
|
|
topk,
|
|
deterministic=envs.SGLANG_DSA_TOPK_FLASHINFER_DETERMINISTIC.get(),
|
|
tie_break=_flashinfer_tie_break_value(),
|
|
dsa_graph_safe=True,
|
|
row_starts=row_starts,
|
|
)
|
|
raise RuntimeError(f"Unsupported {topk_transform_method = }.")
|
|
|
|
raise RuntimeError(f"Unsupported {self = } for SGLANG_DSA_FUSE_TOPK.")
|
|
|
|
|
|
def _topk_unfused(
|
|
score: torch.Tensor,
|
|
lengths: torch.Tensor,
|
|
topk: int,
|
|
row_starts: Optional[torch.Tensor] = None,
|
|
topk_op: Callable[..., Tuple[torch.Tensor, torch.Tensor]] = torch.topk,
|
|
topk_op_kwargs: Optional[Dict[str, object]] = None,
|
|
) -> torch.Tensor:
|
|
batch_size, max_score_len = score.shape
|
|
topk_indices = score.new_full((batch_size, topk), -1, dtype=torch.int32)
|
|
if batch_size == 0 or topk == 0 or max_score_len == 0:
|
|
return topk_indices
|
|
|
|
if row_starts is None:
|
|
row_starts = torch.zeros_like(lengths, dtype=torch.int32, device=score.device)
|
|
else:
|
|
row_starts = row_starts.to(dtype=torch.int32, device=score.device)
|
|
lengths = lengths.to(dtype=torch.int32, device=score.device)
|
|
|
|
col_indices = torch.arange(max_score_len, dtype=torch.int32, device=score.device)
|
|
col_indices = col_indices.unsqueeze(0)
|
|
row_starts_unsqueezed = row_starts.unsqueeze(1)
|
|
row_ends_unsqueezed = (row_starts + lengths).unsqueeze(1)
|
|
valid_mask = (col_indices >= row_starts_unsqueezed) & (
|
|
col_indices < row_ends_unsqueezed
|
|
)
|
|
|
|
masked_logits = score.masked_fill(~valid_mask, float("-inf"))
|
|
valid_topk = min(topk, max_score_len)
|
|
topk_kwargs = topk_op_kwargs or {}
|
|
topk_scores, topk_col_indices = topk_op(masked_logits, valid_topk, **topk_kwargs)
|
|
topk_local_indices = topk_col_indices.to(torch.int32) - row_starts_unsqueezed
|
|
topk_local_indices = topk_local_indices.masked_fill(
|
|
topk_scores == float("-inf"), -1
|
|
)
|
|
topk_indices[:, :valid_topk] = topk_local_indices
|
|
|
|
return topk_indices
|
|
|
|
|
|
def _topk_transform_v2_paged(
|
|
logits: torch.Tensor,
|
|
lengths: torch.Tensor,
|
|
topk: int,
|
|
attn_metadata,
|
|
) -> torch.Tensor:
|
|
"""Fused top-k + page-table transform via the DeepSeek-V4 v2 JIT kernel.
|
|
|
|
Returns the transformed page indices ``(num_rows, topk)`` int32 (physical
|
|
page_size=1 KV slots, ``-1`` padded) -- identical in meaning to
|
|
``fast_topk_transform_fused`` / ``flashinfer.top_k_page_table_transform``.
|
|
The kernel selects, per row, the top-k of ``logits[row, :lengths[row]]`` and
|
|
maps each selected position ``p`` through the page table as
|
|
``real_page_table[row, p // page_size] * page_size + (p % page_size)``. Feeding
|
|
it the indexer's compact ``real_page_table`` (page_size = pool page size,
|
|
typically 64) yields the same physical slots as gathering the page_size=1
|
|
table, without materializing that wide table.
|
|
|
|
This is a committed contract, not a best-effort path: ``topk_transform`` routes
|
|
here only for the decode-shaped PAGED case, and the fused-decode CUDA graph
|
|
drops the page_size=1 table for exactly this case (see
|
|
``dsa_drop_wide_page_table``). The preconditions below are therefore
|
|
invariants the caller must uphold -- they assert (raise) on violation rather
|
|
than fall back to the slow legacy path (which may not even have a page_size=1
|
|
table to fall back to) or silently paper over bad input (padding, recomputing
|
|
the plan) at the cost of the performance this path exists to deliver.
|
|
|
|
``lengths`` entries must be NON-NEGATIVE: the kernel reads them as
|
|
``uint32_t``, so a negative row length (DP-padded / idle-companion rows)
|
|
reinterprets as ~4e9 tokens and illegal-addresses. Metadata producers clamp
|
|
padded rows to 0 (see ``fused_dsa_draft_extend_metadata`` /
|
|
``seqlens_expand_kernel``); 0 takes the trivial all-(-1) output path.
|
|
"""
|
|
from sglang.jit_kernel.dsv4.topk import topk_transform_512_v2
|
|
from sglang.srt.model_executor.forward_context import get_token_to_kv_pool
|
|
|
|
num_rows = logits.shape[0]
|
|
|
|
# The indexer (DeepGEMM) emits fp32 scores with unit row stride and a 16B-aligned
|
|
# row stride (a multiple of 4), which is exactly the kernel's ABI (it checks
|
|
# score_stride % 4 == 0 with strides {S, 1}). This holds even though the scores
|
|
# may be a padded view (stride(0) > width, so not `is_contiguous()`); assert the
|
|
# real requirement rather than force a contiguous copy of the wide score buffer.
|
|
assert (
|
|
logits.dtype == torch.float32
|
|
and logits.stride(1) == 1
|
|
and logits.stride(0) % 4 == 0
|
|
), f"v2 top-k expects fp32 scores with unit row stride and 16B-aligned score_stride, got {logits.dtype=} {logits.stride()=}"
|
|
assert 0 < topk <= 2048, f"v2 top-k supports 0 < topk <= 2048, got {topk=}"
|
|
|
|
page_table = attn_metadata.real_page_table
|
|
assert page_table.dtype == torch.int32
|
|
lengths_i32 = lengths.to(torch.int32)
|
|
|
|
# The plan is preprocessed once per forward (DSAMetadata.topk_v2_plan,
|
|
# refreshed in-place under CUDA graph) and reused across layers. A missing or
|
|
# mismatched plan means the caller skipped that preprocessing -- fail loudly
|
|
# rather than silently recompute it per layer.
|
|
plan = attn_metadata.topk_v2_plan
|
|
assert (
|
|
plan is not None and plan.shape[0] == num_rows + 1
|
|
), "topk_v2_plan must be preprocessed per forward (see DSAMetadata.topk_v2_plan)"
|
|
|
|
page_size = get_token_to_kv_pool().page_size
|
|
out = logits.new_full((num_rows, topk), -1, dtype=torch.int32)
|
|
topk_transform_512_v2(logits, lengths_i32, page_table, out, page_size, plan)
|
|
return out
|
|
|
|
|
|
def _build_flashinfer_paged_args(
|
|
attn_metadata,
|
|
row_starts: Optional[torch.Tensor],
|
|
cu_seqlens_q_topk: Optional[torch.Tensor],
|
|
batch_idx_list: Optional[List[int]],
|
|
device: torch.device,
|
|
num_rows: int,
|
|
) -> Tuple[Optional[torch.Tensor], Optional[torch.Tensor]]:
|
|
row_to_batch = (
|
|
torch.as_tensor(batch_idx_list, dtype=torch.int32, device=device)
|
|
if batch_idx_list is not None
|
|
else None
|
|
)
|
|
|
|
if (
|
|
row_to_batch is not None
|
|
and cu_seqlens_q_topk is not None
|
|
and row_to_batch.shape[0] != num_rows
|
|
):
|
|
q_lens = (cu_seqlens_q_topk[1:] - cu_seqlens_q_topk[:-1]).to(
|
|
dtype=torch.int32, device=device
|
|
)
|
|
row_to_batch = torch.repeat_interleave(row_to_batch, q_lens)
|
|
|
|
if row_to_batch is None and cu_seqlens_q_topk is not None:
|
|
# Decode-like case (one query row per batch) does not need an explicit mapping.
|
|
# Avoid dynamic tensor construction in this branch to keep CUDA graph capture safe.
|
|
num_batches = cu_seqlens_q_topk.shape[0] - 1
|
|
if not (row_starts is None and num_rows == num_batches):
|
|
q_lens = (cu_seqlens_q_topk[1:] - cu_seqlens_q_topk[:-1]).to(
|
|
dtype=torch.int32, device=device
|
|
)
|
|
row_to_batch = torch.repeat_interleave(
|
|
torch.arange(q_lens.shape[0], dtype=torch.int32, device=device),
|
|
q_lens,
|
|
)
|
|
|
|
if row_starts is not None and row_to_batch is None:
|
|
raise RuntimeError(
|
|
"PAGED topk_transform with row_starts requires cu_seqlens_q metadata."
|
|
)
|
|
|
|
local_row_starts = row_starts
|
|
if local_row_starts is not None and row_to_batch is not None:
|
|
local_row_starts = (
|
|
local_row_starts - attn_metadata.cu_seqlens_k[:-1][row_to_batch]
|
|
)
|
|
|
|
return row_to_batch, local_row_starts
|
|
|
|
|
|
def _flashinfer_tie_break_value() -> int:
|
|
mode = envs.SGLANG_DSA_TOPK_FLASHINFER_TIE_BREAK.get()
|
|
if mode is None:
|
|
return 0
|
|
mode = mode.lower()
|
|
if mode not in _FLASHINFER_TIE_BREAK_VALUES:
|
|
raise RuntimeError(
|
|
"SGLANG_DSA_TOPK_FLASHINFER_TIE_BREAK must be one of "
|
|
f"{tuple(_FLASHINFER_TIE_BREAK_VALUES)} or unset, got {mode!r}."
|
|
)
|
|
return _FLASHINFER_TIE_BREAK_VALUES[mode]
|