Files
sglang/python/sglang/srt/speculative/spec_info.py
2026-01-09 11:20:29 +08:00

144 lines
4.8 KiB
Python

from __future__ import annotations
from abc import ABC, abstractmethod
from enum import Enum, IntEnum, auto
from typing import TYPE_CHECKING, List, Optional, Tuple, Type, Union
if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import ModelWorkerBatch
from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker
from sglang.srt.speculative.ngram_worker import NGRAMWorker
class SpeculativeAlgorithm(Enum):
"""Enumeration of speculative decoding algorithms."""
EAGLE = auto()
EAGLE3 = auto()
STANDALONE = auto()
NGRAM = auto()
NONE = auto()
@classmethod
def from_string(cls, name: Optional[str]) -> SpeculativeAlgorithm:
if name is None:
return cls.NONE
try:
return cls[name.upper()]
except KeyError:
raise ValueError(f"Unknown speculative algorithm name: {name}")
def is_none(self) -> bool:
return self == SpeculativeAlgorithm.NONE
def is_eagle(self) -> bool:
# NOTE: EAGLE3 is a variant of EAGLE
return self == SpeculativeAlgorithm.EAGLE or self == SpeculativeAlgorithm.EAGLE3
def is_eagle3(self) -> bool:
return self == SpeculativeAlgorithm.EAGLE3
def is_standalone(self) -> bool:
return self == SpeculativeAlgorithm.STANDALONE
def is_ngram(self) -> bool:
return self == SpeculativeAlgorithm.NGRAM
def supports_spec_v2(self) -> bool:
return self.is_eagle() or self.is_standalone()
def create_worker(
self, server_args: ServerArgs
) -> Optional[Union[Type[BaseSpecWorker], Type[TpModelWorker], Type[NGRAMWorker]]]:
assert (
not self.is_none()
), "Cannot create worker for NONE speculative algorithm."
enable_overlap = not server_args.disable_overlap_schedule
if self.is_eagle() and server_args.enable_multi_layer_eagle:
# FIXME: migrate to EagleWorker
if enable_overlap:
from sglang.srt.speculative.multi_layer_eagle_worker_v2 import (
MultiLayerEagleWorkerV2,
)
return MultiLayerEagleWorkerV2
from sglang.srt.speculative.multi_layer_eagle_worker import (
MultiLayerEagleWorker,
)
return MultiLayerEagleWorker
elif self.is_eagle():
if enable_overlap:
from sglang.srt.speculative.eagle_worker_v2 import EAGLEWorkerV2
return EAGLEWorkerV2
from sglang.srt.speculative.eagle_worker import EAGLEWorker
return EAGLEWorker
elif self.is_standalone():
if enable_overlap:
from sglang.srt.speculative.standalone_worker_v2 import (
StandaloneWorkerV2,
)
return StandaloneWorkerV2
from sglang.srt.speculative.standalone_worker import StandaloneWorker
return StandaloneWorker
elif self.is_ngram():
if enable_overlap:
raise ValueError(
f"Speculative algorithm {self.name} does not support overlap worker creation."
)
from sglang.srt.speculative.ngram_worker import NGRAMWorker
return NGRAMWorker
raise ValueError("Unreachable code path in create_worker.")
class SpecInputType(IntEnum):
# NOTE: introduce this to distinguish the SpecInput types of multiple algorithms when asserting in attention backends.
# If all algorithms can share the same datastrucutre of draft_input and verify_input, consider simplify it
EAGLE_DRAFT = auto()
EAGLE_VERIFY = auto()
NGRAM_VERIFY = auto()
class SpecInput(ABC):
def __init__(self, spec_input_type: SpecInputType):
self.spec_input_type = spec_input_type
def is_draft_input(self) -> bool:
# FIXME: remove this function which is only used for assertion
# or use another variable name like `draft_input` to substitute `spec_info`
return self.spec_input_type == SpecInputType.EAGLE_DRAFT
def is_verify_input(self) -> bool:
return self.spec_input_type in {
SpecInputType.EAGLE_VERIFY,
SpecInputType.NGRAM_VERIFY,
}
@abstractmethod
def get_spec_adjust_token_coefficient(self) -> Tuple[int, int]:
pass
def get_spec_adjusted_global_num_tokens(
self, forward_batch: ModelWorkerBatch
) -> Tuple[List[int], List[int]]:
c1, c2 = self.get_spec_adjust_token_coefficient()
global_num_tokens = [x * c1 for x in forward_batch.global_num_tokens]
global_num_tokens_for_logprob = [
x * c2 for x in forward_batch.global_num_tokens_for_logprob
]
return global_num_tokens, global_num_tokens_for_logprob