diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py index 2623f187e..1f1875254 100644 --- a/python/sglang/srt/managers/tp_worker.py +++ b/python/sglang/srt/managers/tp_worker.py @@ -20,9 +20,7 @@ from typing import TYPE_CHECKING, Optional import torch -from sglang.srt.configs.model_config import ModelConfig from sglang.srt.distributed import get_pp_group, get_world_group -from sglang.srt.dllm.algorithm.base import DllmAlgorithm from sglang.srt.managers.io_struct import ( DestroyWeightsUpdateGroupReqInput, GetWeightsByNameReqInput, @@ -41,7 +39,6 @@ from sglang.srt.managers.scheduler import GenerationBatchResult from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors -from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.server_args import ServerArgs from sglang.srt.utils import MultiprocessingSerializer, broadcast_pyobj, set_random_seed from sglang.srt.utils.hf_transformers_utils import ( @@ -53,6 +50,7 @@ from sglang.srt.utils.patch_torch import monkey_patch_torch_reductions if TYPE_CHECKING: from sglang.srt.managers.cache_controller import LayerDoneCounter + from sglang.srt.model_executor.model_runner import ModelRunner logger = logging.getLogger(__name__) @@ -64,7 +62,7 @@ class BaseTpWorker(ABC): @property @abstractmethod - def model_runner(self) -> ModelRunner: + def model_runner(self) -> "ModelRunner": pass @property @@ -220,77 +218,32 @@ class TpModelWorker(BaseTpWorker): is_multi_layer_eagle: bool = False, ): # Parse args + self.server_args = server_args self.tp_size = server_args.tp_size + self.ep_size = server_args.ep_size + self.pp_size = server_args.pp_size self.tp_rank = tp_rank self.moe_ep_rank = moe_ep_rank self.pp_rank = pp_rank + self.dp_rank = dp_rank + self.gpu_id = gpu_id + self.nccl_port = nccl_port + self.is_draft_worker = is_draft_worker + self.is_multi_layer_eagle = is_multi_layer_eagle + self.req_to_token_pool = req_to_token_pool + self.token_to_kv_pool_allocator = token_to_kv_pool_allocator # MTP model runners self.model_runner_list = [] - # Init model and tokenizer - self.model_config = ModelConfig.from_server_args( - server_args, - model_path=( - server_args.model_path - if not is_draft_worker - else server_args.speculative_draft_model_path - ), - model_revision=( - server_args.revision - if not is_draft_worker - else server_args.speculative_draft_model_revision - ), - is_draft_model=is_draft_worker, - ) + self._init_model_config() + self._init_model_runner() - # Init DLLM algorithm - if server_args.dllm_algorithm is not None: - self.dllm_algorithm = DllmAlgorithm.from_server_args(server_args) - else: - self.dllm_algorithm = None - - self._model_runner = ModelRunner( - model_config=self.model_config, - mem_fraction_static=server_args.mem_fraction_static, - gpu_id=gpu_id, - tp_rank=tp_rank, - tp_size=server_args.tp_size, - moe_ep_rank=moe_ep_rank, - moe_ep_size=server_args.ep_size, - pp_rank=pp_rank, - pp_size=server_args.pp_size, - nccl_port=nccl_port, - dp_rank=dp_rank, - server_args=server_args, - is_draft_worker=is_draft_worker, - req_to_token_pool=req_to_token_pool, - token_to_kv_pool_allocator=token_to_kv_pool_allocator, - draft_model_idx=0 if is_multi_layer_eagle else None, - ) if is_multi_layer_eagle: - self.model_runner_list.append(self.model_runner) - for i in range(1, server_args.speculative_num_steps): - self.model_runner_list.append( - ModelRunner( - model_config=self.model_config, - mem_fraction_static=server_args.mem_fraction_static, - gpu_id=gpu_id, - tp_rank=tp_rank, - tp_size=server_args.tp_size, - moe_ep_rank=moe_ep_rank, - moe_ep_size=server_args.ep_size, - pp_rank=pp_rank, - pp_size=server_args.pp_size, - nccl_port=nccl_port, - dp_rank=dp_rank, - server_args=server_args, - is_draft_worker=is_draft_worker, - req_to_token_pool=req_to_token_pool, - token_to_kv_pool_allocator=token_to_kv_pool_allocator, - draft_model_idx=i, - ) - ) + self._init_multi_layer_eagle_model_runners() + + self._init_dllm_algorithm() + if server_args.skip_tokenizer_init: self.tokenizer = self.processor = None else: @@ -346,8 +299,82 @@ class TpModelWorker(BaseTpWorker): self.enable_spec = server_args.speculative_algorithm is not None self.hicache_layer_transfer_counter = None + def _init_model_config(self): + from sglang.srt.configs.model_config import ModelConfig + + self.model_config = ModelConfig.from_server_args( + self.server_args, + model_path=( + self.server_args.model_path + if not self.is_draft_worker + else self.server_args.speculative_draft_model_path + ), + model_revision=( + self.server_args.revision + if not self.is_draft_worker + else self.server_args.speculative_draft_model_revision + ), + is_draft_model=self.is_draft_worker, + ) + + def _init_model_runner(self): + from sglang.srt.model_executor.model_runner import ModelRunner + + self._model_runner = ModelRunner( + model_config=self.model_config, + mem_fraction_static=self.server_args.mem_fraction_static, + gpu_id=self.gpu_id, + tp_rank=self.tp_rank, + tp_size=self.tp_size, + moe_ep_rank=self.moe_ep_rank, + moe_ep_size=self.ep_size, + pp_rank=self.pp_rank, + pp_size=self.pp_size, + nccl_port=self.nccl_port, + dp_rank=self.dp_rank, + server_args=self.server_args, + is_draft_worker=self.is_draft_worker, + req_to_token_pool=self.req_to_token_pool, + token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, + draft_model_idx=0 if self.is_multi_layer_eagle else None, + ) + + def _init_multi_layer_eagle_model_runners(self): + from sglang.srt.model_executor.model_runner import ModelRunner + + self.model_runner_list.append(self.model_runner) + for i in range(1, self.server_args.speculative_num_steps): + self.model_runner_list.append( + ModelRunner( + model_config=self.model_config, + mem_fraction_static=self.server_args.mem_fraction_static, + gpu_id=self.gpu_id, + tp_rank=self.tp_rank, + tp_size=self.tp_size, + moe_ep_rank=self.moe_ep_rank, + moe_ep_size=self.ep_size, + pp_rank=self.pp_rank, + pp_size=self.pp_size, + nccl_port=self.nccl_port, + dp_rank=self.dp_rank, + server_args=self.server_args, + is_draft_worker=self.is_draft_worker, + req_to_token_pool=self.req_to_token_pool, + token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, + draft_model_idx=i, + ) + ) + + def _init_dllm_algorithm(self): + from sglang.srt.dllm.algorithm.base import DllmAlgorithm + + if self.server_args.dllm_algorithm is not None: + self.dllm_algorithm = DllmAlgorithm.from_server_args(self.server_args) + else: + self.dllm_algorithm = None + @property - def model_runner(self) -> ModelRunner: + def model_runner(self) -> "ModelRunner": return self._model_runner def register_hicache_layer_transfer_counter(self, counter: LayerDoneCounter):