Tiny simplify draft worker init. (#16446)
This commit is contained in:
@@ -481,7 +481,11 @@ class Scheduler(
|
|||||||
nccl_port=self.nccl_port,
|
nccl_port=self.nccl_port,
|
||||||
)
|
)
|
||||||
|
|
||||||
def init_draft_worker(self):
|
def maybe_init_draft_worker(self):
|
||||||
|
if self.spec_algorithm.is_none():
|
||||||
|
self.draft_worker = None
|
||||||
|
return
|
||||||
|
|
||||||
# Launch a draft worker for speculative decoding
|
# Launch a draft worker for speculative decoding
|
||||||
draft_worker_kwargs = dict(
|
draft_worker_kwargs = dict(
|
||||||
server_args=self.server_args,
|
server_args=self.server_args,
|
||||||
@@ -501,50 +505,12 @@ class Scheduler(
|
|||||||
f"Using draft model load_format: '{self.server_args.speculative_draft_load_format}'"
|
f"Using draft model load_format: '{self.server_args.speculative_draft_load_format}'"
|
||||||
)
|
)
|
||||||
|
|
||||||
# FIXME: refactor the draft worker registration logic
|
DraftWorkerClass = self.spec_algorithm.create_worker(self.server_args)
|
||||||
if self.server_args.enable_multi_layer_eagle:
|
self.draft_worker = DraftWorkerClass(**draft_worker_kwargs)
|
||||||
if self.enable_overlap:
|
|
||||||
from sglang.srt.speculative.multi_layer_eagle_worker_v2 import (
|
|
||||||
MultiLayerEagleWorkerV2,
|
|
||||||
)
|
|
||||||
|
|
||||||
self.draft_worker = MultiLayerEagleWorkerV2(
|
|
||||||
gpu_id=self.gpu_id,
|
|
||||||
tp_rank=self.tp_rank,
|
|
||||||
moe_ep_rank=self.moe_ep_rank,
|
|
||||||
server_args=self.server_args,
|
|
||||||
nccl_port=self.nccl_port,
|
|
||||||
target_worker=self.tp_worker,
|
|
||||||
dp_rank=self.dp_rank,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
from sglang.srt.speculative.multi_layer_eagle_worker import (
|
|
||||||
MultiLayerEagleWorker,
|
|
||||||
)
|
|
||||||
|
|
||||||
self.draft_worker = MultiLayerEagleWorker(
|
|
||||||
gpu_id=self.gpu_id,
|
|
||||||
tp_rank=self.tp_rank,
|
|
||||||
moe_ep_rank=self.moe_ep_rank,
|
|
||||||
server_args=self.server_args,
|
|
||||||
nccl_port=self.nccl_port,
|
|
||||||
target_worker=self.tp_worker,
|
|
||||||
dp_rank=self.dp_rank,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
WorkerClass = self.spec_algorithm.create_worker(
|
|
||||||
enable_overlap=self.enable_overlap
|
|
||||||
)
|
|
||||||
|
|
||||||
# FIXME: optimize the init draft worker code path
|
|
||||||
if WorkerClass is not None:
|
|
||||||
self.draft_worker = WorkerClass(**draft_worker_kwargs)
|
|
||||||
else:
|
|
||||||
self.draft_worker = None
|
|
||||||
|
|
||||||
def init_model_worker(self):
|
def init_model_worker(self):
|
||||||
self.init_tp_model_worker()
|
self.init_tp_model_worker()
|
||||||
self.init_draft_worker()
|
self.maybe_init_draft_worker()
|
||||||
|
|
||||||
# Dispatch the model worker
|
# Dispatch the model worker
|
||||||
if self.spec_algorithm.is_none():
|
if self.spec_algorithm.is_none():
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ from typing import TYPE_CHECKING, List, Optional, Tuple, Type, Union
|
|||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
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.speculative.base_spec_worker import BaseSpecWorker
|
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker
|
||||||
from sglang.srt.speculative.ngram_worker import NGRAMWorker
|
from sglang.srt.speculative.ngram_worker import NGRAMWorker
|
||||||
|
|
||||||
@@ -49,12 +50,29 @@ class SpeculativeAlgorithm(Enum):
|
|||||||
return self.is_eagle() or self.is_standalone()
|
return self.is_eagle() or self.is_standalone()
|
||||||
|
|
||||||
def create_worker(
|
def create_worker(
|
||||||
self, enable_overlap: bool = False
|
self, server_args: ServerArgs
|
||||||
) -> Optional[Union[Type[BaseSpecWorker], Type[TpModelWorker], Type[NGRAMWorker]]]:
|
) -> Optional[Union[Type[BaseSpecWorker], Type[TpModelWorker], Type[NGRAMWorker]]]:
|
||||||
if self.is_none():
|
assert (
|
||||||
return None
|
not self.is_none()
|
||||||
|
), "Cannot create worker for NONE speculative algorithm."
|
||||||
|
|
||||||
if self.is_eagle():
|
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:
|
if enable_overlap:
|
||||||
from sglang.srt.speculative.eagle_worker_v2 import EAGLEWorkerV2
|
from sglang.srt.speculative.eagle_worker_v2 import EAGLEWorkerV2
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user