Split model_worker init function (#16160)
This commit is contained in:
@@ -465,7 +465,7 @@ class Scheduler(
|
|||||||
# This must be called after initialize_moe_config
|
# This must be called after initialize_moe_config
|
||||||
self.require_mlp_sync = require_mlp_sync(self.server_args)
|
self.require_mlp_sync = require_mlp_sync(self.server_args)
|
||||||
|
|
||||||
def init_model_worker(self):
|
def init_tp_model_worker(self):
|
||||||
from sglang.srt.managers.tp_worker import TpModelWorker
|
from sglang.srt.managers.tp_worker import TpModelWorker
|
||||||
|
|
||||||
self.tp_worker = TpModelWorker(
|
self.tp_worker = TpModelWorker(
|
||||||
@@ -478,6 +478,7 @@ class Scheduler(
|
|||||||
nccl_port=self.nccl_port,
|
nccl_port=self.nccl_port,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def init_draft_worker(self):
|
||||||
# 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,
|
||||||
@@ -537,6 +538,10 @@ class Scheduler(
|
|||||||
**draft_worker_kwargs
|
**draft_worker_kwargs
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def init_model_worker(self):
|
||||||
|
self.init_tp_model_worker()
|
||||||
|
self.init_draft_worker()
|
||||||
|
|
||||||
# Dispatch the model worker
|
# Dispatch the model worker
|
||||||
if self.spec_algorithm.is_none():
|
if self.spec_algorithm.is_none():
|
||||||
self.model_worker = self.tp_worker
|
self.model_worker = self.tp_worker
|
||||||
|
|||||||
Reference in New Issue
Block a user