[Spec] Clear dead DRAFT_EXTEND objects left after EAGLE v1 removal (#28133)

This commit is contained in:
Cheng Wan
2026-06-13 13:04:51 -07:00
committed by GitHub
parent bde6bccf39
commit 27ba13358e
5 changed files with 1 additions and 111 deletions
@@ -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 = (
@@ -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
@@ -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
@@ -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."""
@@ -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,
}