[Spec][DSA] Add --speculative-dsa-topk-backend (#36313)
This commit is contained in:
@@ -560,9 +560,7 @@ class DeepseekV4AttnBackend(
|
||||
self.enable_deepseek_v4_fp4_indexer: bool = (
|
||||
model_runner.server_args.enable_deepseek_v4_fp4_indexer
|
||||
)
|
||||
self.dsa_topk_backend: DSATopKBackend = DSATopKBackend(
|
||||
model_runner.server_args.dsa_topk_backend
|
||||
)
|
||||
self.dsa_topk_backend: DSATopKBackend = DSATopKBackend.resolve(model_runner)
|
||||
self.dsv4_prefill_backend: str = getattr(
|
||||
model_runner.server_args, "dsv4_prefill_backend", "auto"
|
||||
)
|
||||
|
||||
@@ -1,11 +1,15 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from enum import Enum, IntEnum, auto
|
||||
from typing import Callable, Dict, List, Optional, Tuple
|
||||
from typing import TYPE_CHECKING, Callable, Dict, List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.runtime_context import get_exec, get_spec
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
|
||||
_FLASHINFER_TIE_BREAK_VALUES = {
|
||||
"small": 1,
|
||||
@@ -25,6 +29,17 @@ class DSATopKBackend(Enum):
|
||||
TORCH = "torch"
|
||||
FLASHINFER = "flashinfer"
|
||||
|
||||
@classmethod
|
||||
def resolve(cls, model_runner: ModelRunner) -> DSATopKBackend:
|
||||
"""Resolve the DSA top-k backend for one model runner.
|
||||
|
||||
``--dsa-topk-backend`` selects the target backend, while
|
||||
``--speculative-dsa-topk-backend`` independently selects the draft.
|
||||
"""
|
||||
if model_runner.is_draft_worker:
|
||||
return cls(get_spec().speculative_dsa_topk_backend)
|
||||
return cls(get_exec().kernel.dsa_topk_backend)
|
||||
|
||||
def is_sgl_kernel(self) -> bool:
|
||||
return self == DSATopKBackend.SGL_KERNEL
|
||||
|
||||
|
||||
@@ -340,9 +340,7 @@ class DeepseekSparseAttnBackend(
|
||||
self.supports_mha_one_shot: bool = True
|
||||
self.dsa_prefill_impl: _DSA_IMPL_T = get_exec().kernel.dsa_prefill_backend
|
||||
self.dsa_decode_impl: _DSA_IMPL_T = get_exec().kernel.dsa_decode_backend
|
||||
self.dsa_topk_backend: DSATopKBackend = DSATopKBackend(
|
||||
model_runner.server_args.dsa_topk_backend
|
||||
)
|
||||
self.dsa_topk_backend: DSATopKBackend = DSATopKBackend.resolve(model_runner)
|
||||
if self.num_q_heads <= 64:
|
||||
self.flashmla_kv_num_q_heads = 64
|
||||
elif self.num_q_heads <= 128:
|
||||
|
||||
@@ -1897,7 +1897,7 @@ class ServerArgs:
|
||||
dsa_topk_backend: A[
|
||||
str,
|
||||
Arg(
|
||||
help="DSA indexer top-k backend. Options: 'sgl-kernel', 'torch', 'flashinfer'. The 'torch' backend currently requires SGLANG_DSA_FUSE_TOPK=false.",
|
||||
help="DSA indexer top-k backend for the target model. Options: 'sgl-kernel', 'torch', 'flashinfer'. The 'torch' backend currently requires SGLANG_DSA_FUSE_TOPK=false.",
|
||||
choices=DSA_TOPK_BACKEND_CHOICES,
|
||||
),
|
||||
NS("exec.kernel"),
|
||||
@@ -2264,6 +2264,14 @@ class ServerArgs:
|
||||
),
|
||||
NS("spec"),
|
||||
] = None
|
||||
speculative_dsa_topk_backend: A[
|
||||
str,
|
||||
Arg(
|
||||
help="DSA indexer top-k backend for speculative draft workers. Options: 'sgl-kernel', 'torch', 'flashinfer'. The 'torch' backend currently requires SGLANG_DSA_FUSE_TOPK=false.",
|
||||
choices=DSA_TOPK_BACKEND_CHOICES,
|
||||
),
|
||||
NS("spec"),
|
||||
] = "sgl-kernel"
|
||||
speculative_draft_kv_cache_dtype: A[
|
||||
Optional[str],
|
||||
Arg(
|
||||
|
||||
@@ -298,6 +298,7 @@ class DSAMockModelRunner(ModelRunner):
|
||||
self.prefill_attention_backend_str = case.backend
|
||||
self.decode_attention_backend_str = case.backend
|
||||
self.draft_attention_backend = None
|
||||
self.is_draft_worker = False
|
||||
# For TARGET_VERIFY / DRAFT_EXTEND, the DSA backend uses
|
||||
# `self.speculative_num_draft_tokens` to size `seqlens_expanded`
|
||||
# (`dsa_backend.py:482-486,510-515`). When zero, deep_gemm's
|
||||
|
||||
Reference in New Issue
Block a user