[DSA] Drop the redundant 512 from the top-k transform entry-point names (#36831)
Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 5
parent
00689c0c94
commit
b6c06e1efb
@@ -32,7 +32,7 @@ from .moe import (
|
||||
silu_and_mul_contig_post_quant,
|
||||
silu_and_mul_masked_post_quant,
|
||||
)
|
||||
from .topk import plan_topk_v2, topk_transform_512, topk_transform_512_v2
|
||||
from .topk import plan_topk_v2, topk_transform_paged, topk_transform_paged_v2
|
||||
from .utils import make_name
|
||||
|
||||
__all__ = [
|
||||
@@ -54,8 +54,8 @@ __all__ = [
|
||||
"linear_bf16_fp32",
|
||||
"get_paged_mqa_logits_metadata",
|
||||
"triton_create_paged_compress_data",
|
||||
"topk_transform_512",
|
||||
"topk_transform_512_v2",
|
||||
"topk_transform_paged",
|
||||
"topk_transform_paged_v2",
|
||||
"plan_topk_v2",
|
||||
"hash_topk",
|
||||
"mega_moe_pre_dispatch",
|
||||
|
||||
@@ -47,7 +47,7 @@ def _jit_topk_v2_module():
|
||||
)
|
||||
|
||||
|
||||
def topk_transform_512(
|
||||
def topk_transform_paged(
|
||||
scores: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
page_tables: torch.Tensor,
|
||||
@@ -76,7 +76,7 @@ _PLAN_METADATA_INTS_PER_BATCH = 2
|
||||
|
||||
|
||||
def plan_topk_v2(seq_lens: torch.Tensor, static_threshold: int = 0) -> torch.Tensor:
|
||||
"""Preprocess the per-batch routing plan for :func:`topk_transform_512_v2`.
|
||||
"""Preprocess the per-batch routing plan for :func:`topk_transform_paged_v2`.
|
||||
|
||||
IMPORTANT: every entry of ``seq_lens`` must be NON-NEGATIVE. The device
|
||||
kernel reads the int32 buffer as ``uint32_t``, so a negative length (e.g.
|
||||
@@ -108,7 +108,7 @@ def topk_transform_ragged_v2(
|
||||
With the production convention ``out_offsets == row_starts`` that is the
|
||||
column index itself, i.e. the token's slot in the batch's flattened KV.
|
||||
|
||||
Unlike :func:`topk_transform_512_v2` this needs no page table and no plan
|
||||
Unlike :func:`topk_transform_paged_v2` this needs no page table and no plan
|
||||
(the cluster path only pays off for very few rows, and prefill has many).
|
||||
|
||||
IMPORTANT: ``scores`` is written in place -- the <= 3 columns ahead of each
|
||||
@@ -129,7 +129,7 @@ def topk_transform_ragged_v2(
|
||||
module.topk_transform_ragged(scores, seq_lens, row_starts, out_offsets, out_indices)
|
||||
|
||||
|
||||
def topk_transform_512_v2(
|
||||
def topk_transform_paged_v2(
|
||||
scores: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
page_tables: Optional[torch.Tensor],
|
||||
|
||||
@@ -300,7 +300,7 @@ def handle_model_specific_adjustments(server_args: Any):
|
||||
run_post_process_pass(server_args, _deepseek_moe_quant_resolution)
|
||||
if get_platform().is_hip:
|
||||
if is_deepseek_dsa(hf_config):
|
||||
# The fused top-k v2 kernel (topk_transform_512_v2) is a
|
||||
# The fused top-k v2 kernel (topk_transform_paged_v2) is a
|
||||
# CUDA/Hopper-only path: its JIT source includes
|
||||
# <cooperative_groups.h> and uses cg::this_cluster()
|
||||
# (thread-block clusters), neither of which exists on ROCm,
|
||||
|
||||
@@ -306,7 +306,7 @@ def _topk_transform_v2_paged(
|
||||
padded rows to 0 (see ``fused_dsa_draft_extend_metadata`` /
|
||||
``seqlens_expand_kernel``); 0 takes the trivial all-(-1) output path.
|
||||
"""
|
||||
from sglang.kernels.ops.attention.dsv4.topk import topk_transform_512_v2
|
||||
from sglang.kernels.ops.attention.dsv4.topk import topk_transform_paged_v2
|
||||
|
||||
num_rows = logits.shape[0]
|
||||
|
||||
@@ -335,7 +335,7 @@ def _topk_transform_v2_paged(
|
||||
|
||||
page_size = attn_metadata.page_size
|
||||
out = logits.new_empty((num_rows, topk), dtype=torch.int32)
|
||||
topk_transform_512_v2(logits, lengths, page_table, out, page_size, plan)
|
||||
topk_transform_paged_v2(logits, lengths, page_table, out, page_size, plan)
|
||||
return out
|
||||
|
||||
|
||||
|
||||
@@ -19,8 +19,8 @@ import torch.nn.functional as F
|
||||
from sglang.kernels.ops.attention.dsv4 import (
|
||||
fused_q_indexer_rope_hadamard_fp4_quant,
|
||||
fused_q_indexer_rope_hadamard_quant,
|
||||
topk_transform_512,
|
||||
topk_transform_512_v2,
|
||||
topk_transform_paged,
|
||||
topk_transform_paged_v2,
|
||||
)
|
||||
from sglang.kernels.ops.quantization.fp8_kernel import is_fp8_fnuz
|
||||
from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config
|
||||
@@ -260,7 +260,7 @@ def fp8_paged_mqa_logits_torch_sm120(
|
||||
return logits
|
||||
|
||||
|
||||
def _topk_transform_512_vectorized(
|
||||
def _topk_transform_vectorized(
|
||||
scores: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
page_tables: torch.Tensor,
|
||||
@@ -348,7 +348,7 @@ def _topk_transform_512_vectorized(
|
||||
out_raw_indices.copy_(raw_indices)
|
||||
|
||||
|
||||
def topk_transform_512_pytorch_vectorized(
|
||||
def topk_transform_pytorch_vectorized(
|
||||
scores: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
page_tables: torch.Tensor,
|
||||
@@ -356,11 +356,11 @@ def topk_transform_512_pytorch_vectorized(
|
||||
page_size: int,
|
||||
out_raw_indices: Optional[torch.Tensor] = None,
|
||||
) -> None:
|
||||
"""Vectorized PyTorch fallback for topk_transform_512.
|
||||
"""Vectorized PyTorch fallback for topk_transform_paged.
|
||||
All helper tensors (arange, zeros) are cached to avoid device-tensor
|
||||
creation during HIP/CUDA graph capture."""
|
||||
|
||||
_topk_transform_512_vectorized(
|
||||
_topk_transform_vectorized(
|
||||
scores,
|
||||
seq_lens,
|
||||
page_tables,
|
||||
@@ -372,7 +372,7 @@ def topk_transform_512_pytorch_vectorized(
|
||||
)
|
||||
|
||||
|
||||
def topk_transform_512_flashinfer_unfused(
|
||||
def topk_transform_flashinfer_unfused(
|
||||
scores: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
page_tables: torch.Tensor,
|
||||
@@ -386,7 +386,7 @@ def topk_transform_512_flashinfer_unfused(
|
||||
_flashinfer_tie_break_value,
|
||||
)
|
||||
|
||||
_topk_transform_512_vectorized(
|
||||
_topk_transform_vectorized(
|
||||
scores,
|
||||
seq_lens,
|
||||
page_tables,
|
||||
@@ -404,7 +404,7 @@ def topk_transform_512_flashinfer_unfused(
|
||||
)
|
||||
|
||||
|
||||
def topk_transform_512_flashinfer_fused(
|
||||
def topk_transform_flashinfer_fused(
|
||||
scores: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
page_tables: torch.Tensor,
|
||||
@@ -438,9 +438,9 @@ class C4IndexerBackendMixin:
|
||||
self.debug_use_external_c4_sparse_indices: bool = False
|
||||
self.dsa_topk_backend: DSATopKBackend = DSATopKBackend.SGL_KERNEL
|
||||
self.flashinfer_topk_transform: Callable[..., None] = (
|
||||
topk_transform_512_flashinfer_fused
|
||||
topk_transform_flashinfer_fused
|
||||
if envs.SGLANG_DSA_FUSE_TOPK.get()
|
||||
else topk_transform_512_flashinfer_unfused
|
||||
else topk_transform_flashinfer_unfused
|
||||
)
|
||||
|
||||
def _forward_prepare_multi_stream(
|
||||
@@ -853,7 +853,7 @@ class C4IndexerBackendMixin:
|
||||
raw_indices = core_metadata.c4_sparse_raw_indices
|
||||
|
||||
if self.dsa_topk_backend.is_torch():
|
||||
topk_transform_512_pytorch_vectorized(
|
||||
topk_transform_pytorch_vectorized(
|
||||
logits,
|
||||
c4_seq_lens,
|
||||
page_table,
|
||||
@@ -871,7 +871,7 @@ class C4IndexerBackendMixin:
|
||||
raw_indices,
|
||||
)
|
||||
elif self.dsa_topk_backend.should_use_topk_v2() and raw_indices is None:
|
||||
topk_transform_512_v2(
|
||||
topk_transform_paged_v2(
|
||||
logits,
|
||||
c4_seq_lens,
|
||||
page_table,
|
||||
@@ -880,7 +880,7 @@ class C4IndexerBackendMixin:
|
||||
indexer_metadata.topk_metadata,
|
||||
)
|
||||
else:
|
||||
topk_transform_512(
|
||||
topk_transform_paged(
|
||||
logits,
|
||||
c4_seq_lens,
|
||||
page_table,
|
||||
|
||||
Reference in New Issue
Block a user