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:
Liangsheng Yin
2025-12-23 02:06:46 +08:00
committed by GitHub
co-authored by gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
parent b736a1525a
commit 3c882db3ad
10 changed files with 66 additions and 58 deletions
+10 -7
View File
@@ -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
+3 -3
View File
@@ -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(