Add extension points on SpeculativeAlgorithm for custom spec v2 (#24999)
This commit is contained in:
@@ -49,7 +49,7 @@ class FutureMap:
|
|||||||
chunked_prefill_size: int,
|
chunked_prefill_size: int,
|
||||||
context_len: int,
|
context_len: int,
|
||||||
device: torch.device,
|
device: torch.device,
|
||||||
spec_algo: Optional[SpeculativeAlgorithm] = None,
|
spec_algo: SpeculativeAlgorithm,
|
||||||
):
|
):
|
||||||
# FIXME: the calculation of future_limit and future_buffer_len maybe too conservative
|
# FIXME: the calculation of future_limit and future_buffer_len maybe too conservative
|
||||||
self.future_ct = 0
|
self.future_ct = 0
|
||||||
|
|||||||
@@ -152,7 +152,6 @@ from sglang.srt.managers.mm_utils import (
|
|||||||
unwrap_shm_features,
|
unwrap_shm_features,
|
||||||
)
|
)
|
||||||
from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors
|
from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors
|
||||||
from sglang.srt.managers.overlap_utils import FutureMap
|
|
||||||
from sglang.srt.managers.prefill_delayer import (
|
from sglang.srt.managers.prefill_delayer import (
|
||||||
PrefillDelayer,
|
PrefillDelayer,
|
||||||
PrefillDelayerSinglePassExecutor,
|
PrefillDelayerSinglePassExecutor,
|
||||||
@@ -1343,12 +1342,11 @@ class Scheduler(
|
|||||||
self.future_map = None
|
self.future_map = None
|
||||||
return
|
return
|
||||||
|
|
||||||
self.future_map = FutureMap(
|
self.future_map = self.spec_algorithm.create_future_map(
|
||||||
self.max_running_requests,
|
self.max_running_requests,
|
||||||
self.chunked_prefill_size,
|
self.chunked_prefill_size,
|
||||||
self.model_config.context_len,
|
self.model_config.context_len,
|
||||||
self.device,
|
self.device,
|
||||||
self.spec_algorithm,
|
|
||||||
)
|
)
|
||||||
self.batch_record_buf = [None] * 2
|
self.batch_record_buf = [None] * 2
|
||||||
self.batch_record_ct = 0
|
self.batch_record_ct = 0
|
||||||
|
|||||||
@@ -36,7 +36,6 @@ from sglang.srt.model_executor.forward_batch_info import (
|
|||||||
PPProxyTensors,
|
PPProxyTensors,
|
||||||
enable_num_token_non_padded,
|
enable_num_token_non_padded,
|
||||||
)
|
)
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
log_info_on_rank0,
|
log_info_on_rank0,
|
||||||
require_attn_tp_gather,
|
require_attn_tp_gather,
|
||||||
@@ -528,7 +527,7 @@ class CPUGraphRunner:
|
|||||||
not self.require_gathered_buffer
|
not self.require_gathered_buffer
|
||||||
), "CPUGraphRunner does not support gathered buffer yet."
|
), "CPUGraphRunner does not support gathered buffer yet."
|
||||||
assert (
|
assert (
|
||||||
model_runner.spec_algorithm == SpeculativeAlgorithm.NONE
|
model_runner.spec_algorithm.is_none()
|
||||||
), "CPUGraphRunner does not support speculative inference yet."
|
), "CPUGraphRunner does not support speculative inference yet."
|
||||||
# TODO add compile support for encoder-decoder models
|
# TODO add compile support for encoder-decoder models
|
||||||
assert (
|
assert (
|
||||||
|
|||||||
@@ -635,8 +635,10 @@ class CudaGraphRunner:
|
|||||||
self.num_tokens_per_bs = 1
|
self.num_tokens_per_bs = 1
|
||||||
if model_runner.spec_algorithm.is_speculative():
|
if model_runner.spec_algorithm.is_speculative():
|
||||||
if self.model_runner.is_draft_worker:
|
if self.model_runner.is_draft_worker:
|
||||||
# DFLASH draft workers reuse this runner for TARGET_VERIFY mode.
|
# Draft workers can use TARGET_VERIFY mode.
|
||||||
if not self.model_runner.spec_algorithm.is_dflash():
|
if (
|
||||||
|
not self.model_runner.spec_algorithm.supports_target_verify_for_draft()
|
||||||
|
):
|
||||||
raise RuntimeError("This should not happen")
|
raise RuntimeError("This should not happen")
|
||||||
self.capture_forward_mode = ForwardMode.TARGET_VERIFY
|
self.capture_forward_mode = ForwardMode.TARGET_VERIFY
|
||||||
self.num_tokens_per_bs = self.speculative_num_draft_tokens
|
self.num_tokens_per_bs = self.speculative_num_draft_tokens
|
||||||
|
|||||||
@@ -4,6 +4,8 @@ from abc import ABC, abstractmethod
|
|||||||
from enum import Enum, IntEnum, auto
|
from enum import Enum, IntEnum, auto
|
||||||
from typing import TYPE_CHECKING, Callable, List, Optional, Tuple, Type, Union
|
from typing import TYPE_CHECKING, Callable, List, Optional, Tuple, Type, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
from sglang.srt.speculative.spec_registry import (
|
from sglang.srt.speculative.spec_registry import (
|
||||||
CustomSpecAlgo,
|
CustomSpecAlgo,
|
||||||
ServerArgsValidator,
|
ServerArgsValidator,
|
||||||
@@ -15,6 +17,7 @@ from sglang.srt.speculative.spec_registry import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.managers.overlap_utils import FutureMap
|
||||||
from sglang.srt.managers.schedule_batch import ModelWorkerBatch
|
from sglang.srt.managers.schedule_batch import ModelWorkerBatch
|
||||||
from sglang.srt.managers.tp_worker import TpModelWorker
|
from sglang.srt.managers.tp_worker import TpModelWorker
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
@@ -109,6 +112,26 @@ class SpeculativeAlgorithm(Enum):
|
|||||||
def is_ngram(self) -> bool:
|
def is_ngram(self) -> bool:
|
||||||
return self == SpeculativeAlgorithm.NGRAM
|
return self == SpeculativeAlgorithm.NGRAM
|
||||||
|
|
||||||
|
def supports_target_verify_for_draft(self) -> bool:
|
||||||
|
return self.is_dflash()
|
||||||
|
|
||||||
|
def create_future_map(
|
||||||
|
self,
|
||||||
|
max_running_requests: int,
|
||||||
|
chunked_prefill_size: int,
|
||||||
|
context_len: int,
|
||||||
|
device: torch.device,
|
||||||
|
) -> FutureMap:
|
||||||
|
from sglang.srt.managers.overlap_utils import FutureMap
|
||||||
|
|
||||||
|
return FutureMap(
|
||||||
|
max_running_requests,
|
||||||
|
chunked_prefill_size,
|
||||||
|
context_len,
|
||||||
|
device,
|
||||||
|
self,
|
||||||
|
)
|
||||||
|
|
||||||
def supports_spec_v2(self) -> bool:
|
def supports_spec_v2(self) -> bool:
|
||||||
return (self.is_eagle() and not self.is_frozen_kv_mtp()) or self.is_standalone()
|
return (self.is_eagle() and not self.is_frozen_kv_mtp()) or self.is_standalone()
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user