[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:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||||
from sglang.srt.speculative.spec_info import SpecInput
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -72,7 +71,6 @@ class DeepseekSparseAttnBackendMTPPrecomputeMixin:
|
|||||||
seq_lens: torch.Tensor,
|
seq_lens: torch.Tensor,
|
||||||
seq_lens_cpu: torch.Tensor,
|
seq_lens_cpu: torch.Tensor,
|
||||||
forward_mode: ForwardMode,
|
forward_mode: ForwardMode,
|
||||||
spec_info: Optional[SpecInput],
|
|
||||||
) -> PrecomputedMetadata:
|
) -> PrecomputedMetadata:
|
||||||
"""Precompute all shared metadata for multi-step backends.
|
"""Precompute all shared metadata for multi-step backends.
|
||||||
|
|
||||||
@@ -85,8 +83,7 @@ class DeepseekSparseAttnBackendMTPPrecomputeMixin:
|
|||||||
req_pool_indices: Request pool indices [bs]
|
req_pool_indices: Request pool indices [bs]
|
||||||
seq_lens: Sequence lengths [bs]
|
seq_lens: Sequence lengths [bs]
|
||||||
seq_lens_cpu: Sequence lengths on CPU [bs]
|
seq_lens_cpu: Sequence lengths on CPU [bs]
|
||||||
forward_mode: Forward mode (decode/target_verify/draft_extend)
|
forward_mode: Forward mode (decode/target_verify)
|
||||||
spec_info: Speculative decoding info (for draft_extend mode)
|
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
PrecomputedMetadata containing all shared intermediate results
|
PrecomputedMetadata containing all shared intermediate results
|
||||||
@@ -242,84 +239,6 @@ class DeepseekSparseAttnBackendMTPPrecomputeMixin:
|
|||||||
flashmla_metadata=flashmla_metadata,
|
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
|
# Backward-compat alias
|
||||||
DeepseekSparseAttnBackendMTPPrecomputeMixin = (
|
DeepseekSparseAttnBackendMTPPrecomputeMixin = (
|
||||||
|
|||||||
@@ -2395,7 +2395,6 @@ class DeepseekSparseAttnMultiStepBackend:
|
|||||||
seq_lens=forward_batch.seq_lens,
|
seq_lens=forward_batch.seq_lens,
|
||||||
seq_lens_cpu=forward_batch.seq_lens_cpu,
|
seq_lens_cpu=forward_batch.seq_lens_cpu,
|
||||||
forward_mode=ForwardMode.DECODE,
|
forward_mode=ForwardMode.DECODE,
|
||||||
spec_info=forward_batch.spec_info,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# Use multi-backend fused copy when we have 3 or more backends
|
# Use multi-backend fused copy when we have 3 or more backends
|
||||||
|
|||||||
@@ -210,11 +210,6 @@ class RequestStage:
|
|||||||
level=2,
|
level=2,
|
||||||
)
|
)
|
||||||
|
|
||||||
SPEC_DRAFT_EXTEND = RequestStageConfig(
|
|
||||||
"spec_draft_extend",
|
|
||||||
level=3,
|
|
||||||
)
|
|
||||||
|
|
||||||
# CPU-side run batch
|
# CPU-side run batch
|
||||||
RUN_BATCH_CPU = RequestStageConfig(
|
RUN_BATCH_CPU = RequestStageConfig(
|
||||||
"run_batch_cpu",
|
"run_batch_cpu",
|
||||||
@@ -613,7 +608,6 @@ class SchedulerReqTimeStats(ReqTimeStatsBase):
|
|||||||
# speculative decoding
|
# speculative decoding
|
||||||
spec_draft_start_time: float = 0.0
|
spec_draft_start_time: float = 0.0
|
||||||
spec_verify_start_time: float = 0.0
|
spec_verify_start_time: float = 0.0
|
||||||
spec_draft_extend_start_time: float = 0.0
|
|
||||||
|
|
||||||
# other
|
# other
|
||||||
transfer_speed_gb_s: float = 0.0
|
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):
|
def set_run_batch_cpu_start_time(self, ts=None, attrs=None):
|
||||||
ts = ts or time.perf_counter()
|
ts = ts or time.perf_counter()
|
||||||
self.run_batch_cpu_start_time = ts
|
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.mem_cache.memory_pool import KVCache
|
||||||
from sglang.srt.speculative.eagle_info import (
|
from sglang.srt.speculative.eagle_info import (
|
||||||
EagleDraftExtendInput,
|
|
||||||
EagleDraftInput,
|
EagleDraftInput,
|
||||||
EagleVerifyInput,
|
EagleVerifyInput,
|
||||||
)
|
)
|
||||||
@@ -53,14 +52,6 @@ class FrozenKVMTPDraftInput(EagleDraftInput):
|
|||||||
SpecInput.__init__(self, SpecInputType.FROZEN_KV_MTP_DRAFT)
|
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
|
@dataclass
|
||||||
class FrozenKVMTPVerifyInput(EagleVerifyInput):
|
class FrozenKVMTPVerifyInput(EagleVerifyInput):
|
||||||
"""Verify input for Frozen-KV MTP."""
|
"""Verify input for Frozen-KV MTP."""
|
||||||
|
|||||||
@@ -226,7 +226,6 @@ class SpecInputType(IntEnum):
|
|||||||
EAGLE_DRAFT_EXTEND = auto()
|
EAGLE_DRAFT_EXTEND = auto()
|
||||||
EAGLE_VERIFY = auto()
|
EAGLE_VERIFY = auto()
|
||||||
FROZEN_KV_MTP_DRAFT = auto()
|
FROZEN_KV_MTP_DRAFT = auto()
|
||||||
FROZEN_KV_MTP_DRAFT_EXTEND = auto()
|
|
||||||
FROZEN_KV_MTP_VERIFY = auto()
|
FROZEN_KV_MTP_VERIFY = auto()
|
||||||
DFLASH_DRAFT = auto()
|
DFLASH_DRAFT = auto()
|
||||||
DFLASH_VERIFY = auto()
|
DFLASH_VERIFY = auto()
|
||||||
@@ -246,7 +245,6 @@ class SpecInput(ABC):
|
|||||||
SpecInputType.EAGLE_DRAFT,
|
SpecInputType.EAGLE_DRAFT,
|
||||||
SpecInputType.EAGLE_DRAFT_EXTEND,
|
SpecInputType.EAGLE_DRAFT_EXTEND,
|
||||||
SpecInputType.FROZEN_KV_MTP_DRAFT,
|
SpecInputType.FROZEN_KV_MTP_DRAFT,
|
||||||
SpecInputType.FROZEN_KV_MTP_DRAFT_EXTEND,
|
|
||||||
SpecInputType.DFLASH_DRAFT,
|
SpecInputType.DFLASH_DRAFT,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user