[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_contig_post_quant,
|
||||||
silu_and_mul_masked_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
|
from .utils import make_name
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
@@ -54,8 +54,8 @@ __all__ = [
|
|||||||
"linear_bf16_fp32",
|
"linear_bf16_fp32",
|
||||||
"get_paged_mqa_logits_metadata",
|
"get_paged_mqa_logits_metadata",
|
||||||
"triton_create_paged_compress_data",
|
"triton_create_paged_compress_data",
|
||||||
"topk_transform_512",
|
"topk_transform_paged",
|
||||||
"topk_transform_512_v2",
|
"topk_transform_paged_v2",
|
||||||
"plan_topk_v2",
|
"plan_topk_v2",
|
||||||
"hash_topk",
|
"hash_topk",
|
||||||
"mega_moe_pre_dispatch",
|
"mega_moe_pre_dispatch",
|
||||||
|
|||||||
@@ -47,7 +47,7 @@ def _jit_topk_v2_module():
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def topk_transform_512(
|
def topk_transform_paged(
|
||||||
scores: torch.Tensor,
|
scores: torch.Tensor,
|
||||||
seq_lens: torch.Tensor,
|
seq_lens: torch.Tensor,
|
||||||
page_tables: 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:
|
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
|
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.
|
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
|
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.
|
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).
|
(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
|
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)
|
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,
|
scores: torch.Tensor,
|
||||||
seq_lens: torch.Tensor,
|
seq_lens: torch.Tensor,
|
||||||
page_tables: Optional[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)
|
run_post_process_pass(server_args, _deepseek_moe_quant_resolution)
|
||||||
if get_platform().is_hip:
|
if get_platform().is_hip:
|
||||||
if is_deepseek_dsa(hf_config):
|
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
|
# CUDA/Hopper-only path: its JIT source includes
|
||||||
# <cooperative_groups.h> and uses cg::this_cluster()
|
# <cooperative_groups.h> and uses cg::this_cluster()
|
||||||
# (thread-block clusters), neither of which exists on ROCm,
|
# (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`` /
|
padded rows to 0 (see ``fused_dsa_draft_extend_metadata`` /
|
||||||
``seqlens_expand_kernel``); 0 takes the trivial all-(-1) output path.
|
``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]
|
num_rows = logits.shape[0]
|
||||||
|
|
||||||
@@ -335,7 +335,7 @@ def _topk_transform_v2_paged(
|
|||||||
|
|
||||||
page_size = attn_metadata.page_size
|
page_size = attn_metadata.page_size
|
||||||
out = logits.new_empty((num_rows, topk), dtype=torch.int32)
|
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
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -19,8 +19,8 @@ import torch.nn.functional as F
|
|||||||
from sglang.kernels.ops.attention.dsv4 import (
|
from sglang.kernels.ops.attention.dsv4 import (
|
||||||
fused_q_indexer_rope_hadamard_fp4_quant,
|
fused_q_indexer_rope_hadamard_fp4_quant,
|
||||||
fused_q_indexer_rope_hadamard_quant,
|
fused_q_indexer_rope_hadamard_quant,
|
||||||
topk_transform_512,
|
topk_transform_paged,
|
||||||
topk_transform_512_v2,
|
topk_transform_paged_v2,
|
||||||
)
|
)
|
||||||
from sglang.kernels.ops.quantization.fp8_kernel import is_fp8_fnuz
|
from sglang.kernels.ops.quantization.fp8_kernel import is_fp8_fnuz
|
||||||
from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config
|
from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config
|
||||||
@@ -260,7 +260,7 @@ def fp8_paged_mqa_logits_torch_sm120(
|
|||||||
return logits
|
return logits
|
||||||
|
|
||||||
|
|
||||||
def _topk_transform_512_vectorized(
|
def _topk_transform_vectorized(
|
||||||
scores: torch.Tensor,
|
scores: torch.Tensor,
|
||||||
seq_lens: torch.Tensor,
|
seq_lens: torch.Tensor,
|
||||||
page_tables: torch.Tensor,
|
page_tables: torch.Tensor,
|
||||||
@@ -348,7 +348,7 @@ def _topk_transform_512_vectorized(
|
|||||||
out_raw_indices.copy_(raw_indices)
|
out_raw_indices.copy_(raw_indices)
|
||||||
|
|
||||||
|
|
||||||
def topk_transform_512_pytorch_vectorized(
|
def topk_transform_pytorch_vectorized(
|
||||||
scores: torch.Tensor,
|
scores: torch.Tensor,
|
||||||
seq_lens: torch.Tensor,
|
seq_lens: torch.Tensor,
|
||||||
page_tables: torch.Tensor,
|
page_tables: torch.Tensor,
|
||||||
@@ -356,11 +356,11 @@ def topk_transform_512_pytorch_vectorized(
|
|||||||
page_size: int,
|
page_size: int,
|
||||||
out_raw_indices: Optional[torch.Tensor] = None,
|
out_raw_indices: Optional[torch.Tensor] = None,
|
||||||
) -> 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
|
All helper tensors (arange, zeros) are cached to avoid device-tensor
|
||||||
creation during HIP/CUDA graph capture."""
|
creation during HIP/CUDA graph capture."""
|
||||||
|
|
||||||
_topk_transform_512_vectorized(
|
_topk_transform_vectorized(
|
||||||
scores,
|
scores,
|
||||||
seq_lens,
|
seq_lens,
|
||||||
page_tables,
|
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,
|
scores: torch.Tensor,
|
||||||
seq_lens: torch.Tensor,
|
seq_lens: torch.Tensor,
|
||||||
page_tables: torch.Tensor,
|
page_tables: torch.Tensor,
|
||||||
@@ -386,7 +386,7 @@ def topk_transform_512_flashinfer_unfused(
|
|||||||
_flashinfer_tie_break_value,
|
_flashinfer_tie_break_value,
|
||||||
)
|
)
|
||||||
|
|
||||||
_topk_transform_512_vectorized(
|
_topk_transform_vectorized(
|
||||||
scores,
|
scores,
|
||||||
seq_lens,
|
seq_lens,
|
||||||
page_tables,
|
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,
|
scores: torch.Tensor,
|
||||||
seq_lens: torch.Tensor,
|
seq_lens: torch.Tensor,
|
||||||
page_tables: torch.Tensor,
|
page_tables: torch.Tensor,
|
||||||
@@ -438,9 +438,9 @@ class C4IndexerBackendMixin:
|
|||||||
self.debug_use_external_c4_sparse_indices: bool = False
|
self.debug_use_external_c4_sparse_indices: bool = False
|
||||||
self.dsa_topk_backend: DSATopKBackend = DSATopKBackend.SGL_KERNEL
|
self.dsa_topk_backend: DSATopKBackend = DSATopKBackend.SGL_KERNEL
|
||||||
self.flashinfer_topk_transform: Callable[..., None] = (
|
self.flashinfer_topk_transform: Callable[..., None] = (
|
||||||
topk_transform_512_flashinfer_fused
|
topk_transform_flashinfer_fused
|
||||||
if envs.SGLANG_DSA_FUSE_TOPK.get()
|
if envs.SGLANG_DSA_FUSE_TOPK.get()
|
||||||
else topk_transform_512_flashinfer_unfused
|
else topk_transform_flashinfer_unfused
|
||||||
)
|
)
|
||||||
|
|
||||||
def _forward_prepare_multi_stream(
|
def _forward_prepare_multi_stream(
|
||||||
@@ -853,7 +853,7 @@ class C4IndexerBackendMixin:
|
|||||||
raw_indices = core_metadata.c4_sparse_raw_indices
|
raw_indices = core_metadata.c4_sparse_raw_indices
|
||||||
|
|
||||||
if self.dsa_topk_backend.is_torch():
|
if self.dsa_topk_backend.is_torch():
|
||||||
topk_transform_512_pytorch_vectorized(
|
topk_transform_pytorch_vectorized(
|
||||||
logits,
|
logits,
|
||||||
c4_seq_lens,
|
c4_seq_lens,
|
||||||
page_table,
|
page_table,
|
||||||
@@ -871,7 +871,7 @@ class C4IndexerBackendMixin:
|
|||||||
raw_indices,
|
raw_indices,
|
||||||
)
|
)
|
||||||
elif self.dsa_topk_backend.should_use_topk_v2() and raw_indices is None:
|
elif self.dsa_topk_backend.should_use_topk_v2() and raw_indices is None:
|
||||||
topk_transform_512_v2(
|
topk_transform_paged_v2(
|
||||||
logits,
|
logits,
|
||||||
c4_seq_lens,
|
c4_seq_lens,
|
||||||
page_table,
|
page_table,
|
||||||
@@ -880,7 +880,7 @@ class C4IndexerBackendMixin:
|
|||||||
indexer_metadata.topk_metadata,
|
indexer_metadata.topk_metadata,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
topk_transform_512(
|
topk_transform_paged(
|
||||||
logits,
|
logits,
|
||||||
c4_seq_lens,
|
c4_seq_lens,
|
||||||
page_table,
|
page_table,
|
||||||
|
|||||||
@@ -3,8 +3,8 @@ import torch
|
|||||||
from sglang.kernels.jit.benchmark import marker
|
from sglang.kernels.jit.benchmark import marker
|
||||||
from sglang.kernels.ops.attention.dsv4.topk import (
|
from sglang.kernels.ops.attention.dsv4.topk import (
|
||||||
plan_topk_v2,
|
plan_topk_v2,
|
||||||
topk_transform_512,
|
topk_transform_paged,
|
||||||
topk_transform_512_v2,
|
topk_transform_paged_v2,
|
||||||
topk_transform_ragged_v2,
|
topk_transform_ragged_v2,
|
||||||
)
|
)
|
||||||
from sglang.test.ci.ci_register import register_cuda_ci
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
@@ -42,10 +42,10 @@ def _build_paged_fn(
|
|||||||
|
|
||||||
def fn(scores, seq_lens, page_table):
|
def fn(scores, seq_lens, page_table):
|
||||||
if provider == "jit_v1":
|
if provider == "jit_v1":
|
||||||
topk_transform_512(scores, seq_lens, page_table, out, N)
|
topk_transform_paged(scores, seq_lens, page_table, out, N)
|
||||||
return out
|
return out
|
||||||
elif provider == "jit_v2":
|
elif provider == "jit_v2":
|
||||||
topk_transform_512_v2(scores, seq_lens, page_table, out, N, metadata)
|
topk_transform_paged_v2(scores, seq_lens, page_table, out, N, metadata)
|
||||||
return out
|
return out
|
||||||
elif provider == "flashinfer":
|
elif provider == "flashinfer":
|
||||||
from flashinfer import top_k_page_table_transform
|
from flashinfer import top_k_page_table_transform
|
||||||
|
|||||||
@@ -31,7 +31,7 @@ import torch
|
|||||||
|
|
||||||
from sglang.kernels.ops.attention.dsv4.topk import (
|
from sglang.kernels.ops.attention.dsv4.topk import (
|
||||||
plan_topk_v2,
|
plan_topk_v2,
|
||||||
topk_transform_512_v2,
|
topk_transform_paged_v2,
|
||||||
topk_transform_ragged_v2,
|
topk_transform_ragged_v2,
|
||||||
)
|
)
|
||||||
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||||
@@ -165,7 +165,7 @@ def _run(scores, seq_lens, page_table, inv_cpu, k):
|
|||||||
batch = scores.shape[0]
|
batch = scores.shape[0]
|
||||||
metadata = _plan(seq_lens)
|
metadata = _plan(seq_lens)
|
||||||
out = torch.full((batch, k), -1, dtype=torch.int32, device=scores.device)
|
out = torch.full((batch, k), -1, dtype=torch.int32, device=scores.device)
|
||||||
topk_transform_512_v2(scores, seq_lens, page_table, out, PAGE_SIZE, metadata)
|
topk_transform_paged_v2(scores, seq_lens, page_table, out, PAGE_SIZE, metadata)
|
||||||
torch.cuda.synchronize()
|
torch.cuda.synchronize()
|
||||||
out_cpu = out.cpu().tolist()
|
out_cpu = out.cpu().tolist()
|
||||||
return [_invert(out_cpu[i], inv_cpu[i]) for i in range(batch)]
|
return [_invert(out_cpu[i], inv_cpu[i]) for i in range(batch)]
|
||||||
@@ -177,7 +177,7 @@ def _run_raw(scores, seq_lens, k):
|
|||||||
batch = scores.shape[0]
|
batch = scores.shape[0]
|
||||||
metadata = _plan(seq_lens)
|
metadata = _plan(seq_lens)
|
||||||
out = torch.full((batch, k), -1, dtype=torch.int32, device=scores.device)
|
out = torch.full((batch, k), -1, dtype=torch.int32, device=scores.device)
|
||||||
topk_transform_512_v2(scores, seq_lens, None, out, PAGE_SIZE, metadata)
|
topk_transform_paged_v2(scores, seq_lens, None, out, PAGE_SIZE, metadata)
|
||||||
torch.cuda.synchronize()
|
torch.cuda.synchronize()
|
||||||
out_cpu = out.cpu().tolist()
|
out_cpu = out.cpu().tolist()
|
||||||
return [[v for v in out_cpu[i] if v != -1] for i in range(batch)]
|
return [[v for v in out_cpu[i] if v != -1] for i in range(batch)]
|
||||||
|
|||||||
Reference in New Issue
Block a user