Adjust wrong mtp meaning introduce by mimo (#15632)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
This commit is contained in:
co-authored by
gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
parent
b736a1525a
commit
3c882db3ad
@@ -275,7 +275,6 @@ class Scheduler(
|
||||
self.spec_algorithm = SpeculativeAlgorithm.from_string(
|
||||
server_args.speculative_algorithm
|
||||
)
|
||||
self.enable_mtp = server_args.enable_mtp
|
||||
self.gpu_id = gpu_id
|
||||
self.page_size = server_args.page_size
|
||||
self.enable_hierarchical_cache = server_args.enable_hierarchical_cache
|
||||
@@ -485,11 +484,13 @@ class Scheduler(
|
||||
draft_worker_kwargs["enable_overlap"] = self.enable_overlap
|
||||
|
||||
# FIXME: refactor the draft worker registration logic
|
||||
if self.enable_mtp:
|
||||
if self.server_args.enable_multi_layer_eagle:
|
||||
if self.enable_overlap:
|
||||
from sglang.srt.speculative.mtp_worker_v2 import MTPWorkerV2
|
||||
from sglang.srt.speculative.multi_layer_eagle_worker_v2 import (
|
||||
MultiLayerEagleWorkerV2,
|
||||
)
|
||||
|
||||
self.draft_worker = MTPWorkerV2(
|
||||
self.draft_worker = MultiLayerEagleWorkerV2(
|
||||
gpu_id=self.gpu_id,
|
||||
tp_rank=self.tp_rank,
|
||||
moe_ep_rank=self.moe_ep_rank,
|
||||
@@ -499,9 +500,11 @@ class Scheduler(
|
||||
dp_rank=self.dp_rank,
|
||||
)
|
||||
else:
|
||||
from sglang.srt.speculative.mtp_worker import MTPWorker
|
||||
from sglang.srt.speculative.multi_layer_eagle_worker import (
|
||||
MultiLayerEagleWorker,
|
||||
)
|
||||
|
||||
self.draft_worker = MTPWorker(
|
||||
self.draft_worker = MultiLayerEagleWorker(
|
||||
gpu_id=self.gpu_id,
|
||||
tp_rank=self.tp_rank,
|
||||
moe_ep_rank=self.moe_ep_rank,
|
||||
@@ -834,7 +837,7 @@ class Scheduler(
|
||||
if self.draft_worker is None or self.spec_algorithm.is_ngram():
|
||||
draft_token_to_kv_pool = None
|
||||
elif self.spec_algorithm.is_eagle() and self.enable_overlap:
|
||||
if self.enable_mtp:
|
||||
if self.server_args.enable_multi_layer_eagle:
|
||||
draft_runner = self.draft_worker.draft_worker.draft_runner_list[0]
|
||||
else:
|
||||
draft_runner = self.draft_worker.draft_worker.draft_runner
|
||||
|
||||
@@ -217,7 +217,7 @@ class TpModelWorker(BaseTpWorker):
|
||||
is_draft_worker: bool = False,
|
||||
req_to_token_pool: Optional[ReqToTokenPool] = None,
|
||||
token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator] = None,
|
||||
is_mtp_worker: bool = False,
|
||||
is_multi_layer_eagle: bool = False,
|
||||
):
|
||||
# Parse args
|
||||
self.tp_size = server_args.tp_size
|
||||
@@ -266,9 +266,9 @@ class TpModelWorker(BaseTpWorker):
|
||||
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_mtp_worker else None,
|
||||
draft_model_idx=0 if is_multi_layer_eagle else None,
|
||||
)
|
||||
if is_mtp_worker:
|
||||
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(
|
||||
|
||||
Reference in New Issue
Block a user