[Spec] Clear dead DRAFT_EXTEND objects left after EAGLE v1 removal (#28133)
This commit is contained in:
@@ -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,
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user