diff --git a/python/sglang/srt/layers/attention/dsa/dsa_backend_mtp_precompute.py b/python/sglang/srt/layers/attention/dsa/dsa_backend_mtp_precompute.py index ad6db1997..5d562dd9b 100644 --- a/python/sglang/srt/layers/attention/dsa/dsa_backend_mtp_precompute.py +++ b/python/sglang/srt/layers/attention/dsa/dsa_backend_mtp_precompute.py @@ -15,7 +15,6 @@ from sglang.srt.layers.attention.dsa.utils import compute_dsa_seqlens if TYPE_CHECKING: from sglang.srt.model_executor.forward_batch_info import ForwardMode - from sglang.srt.speculative.spec_info import SpecInput @dataclass @@ -72,7 +71,6 @@ class DeepseekSparseAttnBackendMTPPrecomputeMixin: seq_lens: torch.Tensor, seq_lens_cpu: torch.Tensor, forward_mode: ForwardMode, - spec_info: Optional[SpecInput], ) -> PrecomputedMetadata: """Precompute all shared metadata for multi-step backends. @@ -85,8 +83,7 @@ class DeepseekSparseAttnBackendMTPPrecomputeMixin: req_pool_indices: Request pool indices [bs] seq_lens: Sequence lengths [bs] seq_lens_cpu: Sequence lengths on CPU [bs] - forward_mode: Forward mode (decode/target_verify/draft_extend) - spec_info: Speculative decoding info (for draft_extend mode) + forward_mode: Forward mode (decode/target_verify) Returns: PrecomputedMetadata containing all shared intermediate results @@ -242,84 +239,6 @@ class DeepseekSparseAttnBackendMTPPrecomputeMixin: flashmla_metadata=flashmla_metadata, ) - def _precompute_draft_extend_mode( - self, - bs: int, - req_pool_indices: torch.Tensor, - seq_lens: torch.Tensor, - seq_lens_cpu: torch.Tensor, - spec_info: SpecInput, - ) -> PrecomputedMetadata: - """Precompute metadata for draft extend mode.""" - max_seqlen_k = int(seq_lens_cpu.max().item()) - - # Cache seqlens - cache_seqlens = seq_lens.to(torch.int32) - cu_seqlens_k = compute_cu_seqlens(cache_seqlens) - - # Extend seqlens from spec_info: num_accept_tokens already includes - # the bonus token (drafts + 1). - extend_seq_lens = spec_info.num_accept_tokens[:bs] - extend_seq_lens_cpu = extend_seq_lens.tolist() - - # Page indices (repeated per accept length) - page_indices = self.req_to_token[req_pool_indices, :max_seqlen_k] - page_indices = torch.repeat_interleave( - page_indices, repeats=extend_seq_lens, dim=0 - ).contiguous() - - # Generate expanded seqlens - seqlens_expanded = torch.cat( - [ - torch.arange( - kv_len - qo_len + 1, - kv_len + 1, - dtype=torch.int32, - device=self.device, - ) - for qo_len, kv_len in zip( - extend_seq_lens_cpu, - seq_lens_cpu.tolist(), - strict=True, - ) - ] - ) - - # Compute DSA seqlens - dsa_cache_seqlens = compute_dsa_seqlens(seqlens_expanded, self.dsa_index_topk) - seqlens_expanded_size = seqlens_expanded.shape[0] - - # DSA cumsum - dsa_cu_seqlens_k = compute_cu_seqlens(dsa_cache_seqlens) - - # Transform page table - if self.real_page_size > 1: - real_page_table = self._transform_table_1_to_real(page_indices) - else: - real_page_table = None - - # FlashMLA metadata - flashmla_metadata = None - if self.dsa_decode_impl == "flashmla_kv": - flashmla_metadata = self._compute_flashmla_metadata( - cache_seqlens=dsa_cache_seqlens, - seq_len_q=1, - ) - - return PrecomputedMetadata( - cache_seqlens=cache_seqlens, - cu_seqlens_k=cu_seqlens_k, - page_indices=page_indices, - real_page_table=real_page_table, - seqlens_expanded=seqlens_expanded, - dsa_cache_seqlens=dsa_cache_seqlens, - dsa_cu_seqlens_k=dsa_cu_seqlens_k, - seqlens_expanded_size=seqlens_expanded_size, - max_len=max_seqlen_k, - max_seqlen_k=max_seqlen_k, - flashmla_metadata=flashmla_metadata, - ) - # Backward-compat alias DeepseekSparseAttnBackendMTPPrecomputeMixin = ( diff --git a/python/sglang/srt/layers/attention/dsa_backend.py b/python/sglang/srt/layers/attention/dsa_backend.py index bdfa980ed..cdc7b8ef2 100644 --- a/python/sglang/srt/layers/attention/dsa_backend.py +++ b/python/sglang/srt/layers/attention/dsa_backend.py @@ -2395,7 +2395,6 @@ class DeepseekSparseAttnMultiStepBackend: seq_lens=forward_batch.seq_lens, seq_lens_cpu=forward_batch.seq_lens_cpu, forward_mode=ForwardMode.DECODE, - spec_info=forward_batch.spec_info, ) # Use multi-backend fused copy when we have 3 or more backends diff --git a/python/sglang/srt/observability/req_time_stats.py b/python/sglang/srt/observability/req_time_stats.py index 1003291d0..0c654795d 100644 --- a/python/sglang/srt/observability/req_time_stats.py +++ b/python/sglang/srt/observability/req_time_stats.py @@ -210,11 +210,6 @@ class RequestStage: level=2, ) - SPEC_DRAFT_EXTEND = RequestStageConfig( - "spec_draft_extend", - level=3, - ) - # CPU-side run batch RUN_BATCH_CPU = RequestStageConfig( "run_batch_cpu", @@ -613,7 +608,6 @@ class SchedulerReqTimeStats(ReqTimeStatsBase): # speculative decoding spec_draft_start_time: float = 0.0 spec_verify_start_time: float = 0.0 - spec_draft_extend_start_time: float = 0.0 # other transfer_speed_gb_s: float = 0.0 @@ -679,17 +673,6 @@ class SchedulerReqTimeStats(ReqTimeStatsBase): }, ) - def set_spec_draft_extend_start_time(self, ts=None): - ts = ts or time.perf_counter() - self.spec_draft_extend_start_time = ts - - def set_spec_draft_extend_end_time(self, ts=None): - ts = ts or time.perf_counter() - - if self.trace_ctx.tracing_enable: - stage = RequestStage.SPEC_DRAFT_EXTEND - self.trace_slice(stage, self.spec_draft_extend_start_time, ts) - def set_run_batch_cpu_start_time(self, ts=None, attrs=None): ts = ts or time.perf_counter() self.run_batch_cpu_start_time = ts diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_info.py b/python/sglang/srt/speculative/frozen_kv_mtp_info.py index 1ff0f8c13..614907aa8 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_info.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_info.py @@ -18,7 +18,6 @@ from typing import Dict from sglang.srt.mem_cache.memory_pool import KVCache from sglang.srt.speculative.eagle_info import ( - EagleDraftExtendInput, EagleDraftInput, EagleVerifyInput, ) @@ -53,14 +52,6 @@ class FrozenKVMTPDraftInput(EagleDraftInput): SpecInput.__init__(self, SpecInputType.FROZEN_KV_MTP_DRAFT) -@dataclass -class FrozenKVMTPDraftExtendInput(EagleDraftExtendInput): - """Draft-extend input for Frozen-KV MTP. Tag-only subclass.""" - - def __post_init__(self): - SpecInput.__init__(self, SpecInputType.FROZEN_KV_MTP_DRAFT_EXTEND) - - @dataclass class FrozenKVMTPVerifyInput(EagleVerifyInput): """Verify input for Frozen-KV MTP.""" diff --git a/python/sglang/srt/speculative/spec_info.py b/python/sglang/srt/speculative/spec_info.py index 3fe49c2fb..f3172bc0d 100644 --- a/python/sglang/srt/speculative/spec_info.py +++ b/python/sglang/srt/speculative/spec_info.py @@ -226,7 +226,6 @@ class SpecInputType(IntEnum): EAGLE_DRAFT_EXTEND = auto() EAGLE_VERIFY = auto() FROZEN_KV_MTP_DRAFT = auto() - FROZEN_KV_MTP_DRAFT_EXTEND = auto() FROZEN_KV_MTP_VERIFY = auto() DFLASH_DRAFT = auto() DFLASH_VERIFY = auto() @@ -246,7 +245,6 @@ class SpecInput(ABC): SpecInputType.EAGLE_DRAFT, SpecInputType.EAGLE_DRAFT_EXTEND, SpecInputType.FROZEN_KV_MTP_DRAFT, - SpecInputType.FROZEN_KV_MTP_DRAFT_EXTEND, SpecInputType.DFLASH_DRAFT, }