From 3c882db3adf506b058c10c3c9321723fd4d2b635 Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Tue, 23 Dec 2025 02:06:46 +0800 Subject: [PATCH] Adjust wrong `mtp` meaning introduce by mimo (#15632) Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> --- python/sglang/srt/configs/model_config.py | 6 +-- .../decode_schedule_batch_mixin.py | 2 +- python/sglang/srt/managers/scheduler.py | 17 +++--- python/sglang/srt/managers/tp_worker.py | 6 +-- python/sglang/srt/server_args.py | 13 +++-- ...r_eagle_draft_extend_cuda_graph_runner.py} | 52 +++++++++++-------- ...tp_utils.py => multi_layer_eagle_utils.py} | 0 ..._worker.py => multi_layer_eagle_worker.py} | 10 ++-- ...r_v2.py => multi_layer_eagle_worker_v2.py} | 16 +++--- python/sglang/srt/speculative/spec_utils.py | 2 +- 10 files changed, 66 insertions(+), 58 deletions(-) rename python/sglang/srt/speculative/{mtp_draft_extend_cuda_graph_runner.py => multi_layer_eagle_draft_extend_cuda_graph_runner.py} (94%) rename python/sglang/srt/speculative/{mtp_utils.py => multi_layer_eagle_utils.py} (100%) rename python/sglang/srt/speculative/{mtp_worker.py => multi_layer_eagle_worker.py} (99%) rename python/sglang/srt/speculative/{mtp_worker_v2.py => multi_layer_eagle_worker_v2.py} (98%) diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index 74cf0b8b1..d25f6f952 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -100,7 +100,7 @@ class ModelConfig: model_impl: Union[str, ModelImpl] = ModelImpl.AUTO, sampling_defaults: str = "openai", quantize_and_serve: bool = False, - is_mtp: bool = False, + is_multi_layer_eagle: bool = False, encoder_only: bool = False, language_only: bool = False, ) -> None: @@ -112,7 +112,7 @@ class ModelConfig: self.model_impl = model_impl self.sampling_defaults = sampling_defaults self.quantize_and_serve = quantize_and_serve - self.is_mtp = is_mtp + self.is_multi_layer_eagle = is_multi_layer_eagle # Validate quantize_and_serve configuration self._validate_quantize_and_serve_config() @@ -252,7 +252,7 @@ class ModelConfig: sampling_defaults=server_args.sampling_defaults, quantize_and_serve=server_args.quantize_and_serve, override_config_file=server_args.decrypted_config_file, - is_mtp=server_args.enable_mtp, + is_multi_layer_eagle=server_args.enable_multi_layer_eagle, language_only=server_args.language_only, encoder_only=server_args.encoder_only, is_draft_model=is_draft_model, diff --git a/python/sglang/srt/disaggregation/decode_schedule_batch_mixin.py b/python/sglang/srt/disaggregation/decode_schedule_batch_mixin.py index f4dea976d..81bdb722a 100644 --- a/python/sglang/srt/disaggregation/decode_schedule_batch_mixin.py +++ b/python/sglang/srt/disaggregation/decode_schedule_batch_mixin.py @@ -133,7 +133,7 @@ class ScheduleBatchDisaggregationDecodeMixin: # Simulate the eagle run. if self.spec_algorithm.is_eagle(): num_states = server_args.speculative_eagle_topk - if server_args.enable_mtp: + if server_args.enable_multi_layer_eagle: num_states *= server_args.speculative_num_steps topk_p = torch.stack( [ diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 189a32bb5..ae44104cc 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -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 diff --git a/python/sglang/srt/managers/tp_worker.py b/python/sglang/srt/managers/tp_worker.py index 1b8615e30..2623f187e 100644 --- a/python/sglang/srt/managers/tp_worker.py +++ b/python/sglang/srt/managers/tp_worker.py @@ -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( diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index fc9d2af52..d3033d5e1 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -439,10 +439,7 @@ class ServerArgs: speculative_ngram_match_type: Literal["BFS", "PROB"] = "BFS" speculative_ngram_branch_length: int = 18 speculative_ngram_capacity: int = 10 * 1000 * 1000 - - # For Multi-Layer MTP - # FIXME: rename -> enable_multi_layer_mtp - enable_mtp: bool = False + enable_multi_layer_eagle: bool = False # Expert parallelism ep_size: int = 1 @@ -1189,6 +1186,8 @@ class ServerArgs: self.disable_hybrid_swa_memory = True elif "MiMoV2FlashForCausalLM" in model_arch: + self.enable_multi_layer_eagle = True + logger.info("Enable multi-layer eagle for MiMoV2FlashForCausalLM model") self.swa_full_tokens_ratio = 1.0 logger.warning( "Reset swa_full_tokens_ratio to 1.0 for MiMoV2FlashForCausalLM model" @@ -3478,11 +3477,11 @@ class ServerArgs: help="The cache capacity for ngram speculative decoding.", ) - # Speculative decoding (MTP) + # Multi-layer Eagle speculative decoding parser.add_argument( - "--enable-mtp", + "--enable-multi-layer-eagle", action="store_true", - help="Enable multi-layer MTP speculative decoding.", + help="Enable multi-layer Eagle speculative decoding.", ) # Expert parallelism diff --git a/python/sglang/srt/speculative/mtp_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py similarity index 94% rename from python/sglang/srt/speculative/mtp_draft_extend_cuda_graph_runner.py rename to python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py index f94e51ae4..2387ba6a0 100644 --- a/python/sglang/srt/speculative/mtp_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py @@ -40,7 +40,7 @@ from sglang.srt.model_executor.forward_batch_info import ( ForwardMode, ) from sglang.srt.speculative.eagle_info import EagleDraftInput -from sglang.srt.speculative.mtp_utils import assign_new_state_triton +from sglang.srt.speculative.multi_layer_eagle_utils import assign_new_state_triton from sglang.srt.speculative.spec_utils import fast_topk from sglang.srt.utils import ( get_available_gpu_memory, @@ -51,18 +51,20 @@ from sglang.srt.utils import ( ) if TYPE_CHECKING: - from sglang.srt.speculative.mtp_worker_v2 import MTPDraftWorker + from sglang.srt.speculative.multi_layer_eagle_worker_v2 import ( + MultiLayerEagleDraftWorker, + ) logger = logging.getLogger(__name__) -class MTPDraftExtendCudaGraphRunner: - def __init__(self, mtp_worker: MTPDraftWorker, step: int): +class MultiLayerEagleDraftExtendCudaGraphRunner: + def __init__(self, eagle_worker: MultiLayerEagleDraftWorker, step: int): # Parse args self.step = step - self.mtp_worker = mtp_worker - self.model_runner = model_runner = mtp_worker.mtp_model_runner(self.step) + self.eagle_worker = eagle_worker + self.model_runner = model_runner = eagle_worker.mtp_model_runner(self.step) self.forward_mode = ForwardMode.DRAFT_EXTEND_V2 self.graphs = {} @@ -93,10 +95,10 @@ class MTPDraftExtendCudaGraphRunner: self.max_bs = max(self.capture_bs) self.max_num_token = self.max_bs * self.num_tokens_per_bs - self.mtp_worker.draft_extend_attn_backend_list[self.step].init_cuda_graph_state( - self.max_bs, self.max_num_token - ) - self.seq_len_fill_value = self.mtp_worker.draft_extend_attn_backend_list[ + self.eagle_worker.draft_extend_attn_backend_list[ + self.step + ].init_cuda_graph_state(self.max_bs, self.max_num_token) + self.seq_len_fill_value = self.eagle_worker.draft_extend_attn_backend_list[ self.step ].get_cuda_graph_seq_len_fill_value() @@ -331,7 +333,7 @@ class MTPDraftExtendCudaGraphRunner: spec_algorithm=self.model_runner.spec_algorithm, spec_info=spec_info, capture_hidden_mode=CaptureHiddenMode.FULL, - attn_backend=self.mtp_worker.draft_extend_attn_backend_list[self.step], + attn_backend=self.eagle_worker.draft_extend_attn_backend_list[self.step], extend_seq_lens=extend_seq_lens, extend_seq_lens_cpu=extend_seq_lens_cpu, padded_static_len=self.padded_static_len, @@ -352,7 +354,7 @@ class MTPDraftExtendCudaGraphRunner: num_tokens = bs * self.num_tokens_per_bs forward_batch = self.get_forward_batch(bs) - self.mtp_worker.draft_extend_attn_backend_list[ + self.eagle_worker.draft_extend_attn_backend_list[ self.step ].init_forward_metadata_capture_cuda_graph( bs=bs, @@ -420,7 +422,7 @@ class MTPDraftExtendCudaGraphRunner: self.step, forward_batch.req_pool_indices, forward_batch.req_to_token_pool.req_to_token, - self.mtp_worker.req_to_hidden_states_pool, + self.eagle_worker.req_to_hidden_states_pool, ) self.next_cuda_graph_runner.swa_out_cache_loc.copy_( self.model_runner.token_to_kv_pool.translate_loc_from_full_to_swa( @@ -500,7 +502,7 @@ class MTPDraftExtendCudaGraphRunner: forward_batch.spec_info.positions = self.positions[:num_tokens] forward_batch.spec_info.extend_seq_lens_tensor = self.extend_seq_lens[:bs] - self.mtp_worker.draft_extend_attn_backend_list[ + self.eagle_worker.draft_extend_attn_backend_list[ self.step ].init_forward_metadata_replay_cuda_graph( bs=bs, @@ -540,13 +542,15 @@ class MTPDraftExtendCudaGraphRunner: return out -class MTPMultiStepDraftExtendCudaGraphRunner: - def __init__(self, mtp_worker: MTPDraftWorker): - self.mtp_worker = mtp_worker - self.device = mtp_worker.device - self.gpu_id = mtp_worker.gpu_id - self.speculative_num_steps = mtp_worker.speculative_num_steps - self.draft_extend_attn_backend_list = mtp_worker.draft_extend_attn_backend_list +class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner: + def __init__(self, eagle_worker: MultiLayerEagleDraftWorker): + self.eagle_worker = eagle_worker + self.device = eagle_worker.device + self.gpu_id = eagle_worker.gpu_id + self.speculative_num_steps = eagle_worker.speculative_num_steps + self.draft_extend_attn_backend_list = ( + eagle_worker.draft_extend_attn_backend_list + ) self.runners = [] self.cuda_graph_buffers = {} @@ -557,7 +561,7 @@ class MTPMultiStepDraftExtendCudaGraphRunner: self._init_and_capture() def _init_and_capture(self): - if self.mtp_worker.server_args.disable_cuda_graph: + if self.eagle_worker.server_args.disable_cuda_graph: self.runners = [None] * self.speculative_num_steps return @@ -567,7 +571,9 @@ class MTPMultiStepDraftExtendCudaGraphRunner: # 1. Capture loop for step in range(self.speculative_num_steps): if self.draft_extend_attn_backend_list[step]: - runner = MTPDraftExtendCudaGraphRunner(self.mtp_worker, step) + runner = MultiLayerEagleDraftExtendCudaGraphRunner( + self.eagle_worker, step + ) self.runners.append(runner) self.seq_len_fill_value = runner.seq_len_fill_value diff --git a/python/sglang/srt/speculative/mtp_utils.py b/python/sglang/srt/speculative/multi_layer_eagle_utils.py similarity index 100% rename from python/sglang/srt/speculative/mtp_utils.py rename to python/sglang/srt/speculative/multi_layer_eagle_utils.py diff --git a/python/sglang/srt/speculative/mtp_worker.py b/python/sglang/srt/speculative/multi_layer_eagle_worker.py similarity index 99% rename from python/sglang/srt/speculative/mtp_worker.py rename to python/sglang/srt/speculative/multi_layer_eagle_worker.py index 24cd20a98..455a8e289 100644 --- a/python/sglang/srt/speculative/mtp_worker.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker.py @@ -48,8 +48,8 @@ from sglang.srt.speculative.eagle_utils import ( organize_draft_results, ) from sglang.srt.speculative.eagle_worker import get_last_loc_large_page_size_top_k_1 -from sglang.srt.speculative.mtp_draft_extend_cuda_graph_runner import ( - MTPDraftExtendCudaGraphRunner, +from sglang.srt.speculative.multi_layer_eagle_draft_extend_cuda_graph_runner import ( + MultiLayerEagleDraftExtendCudaGraphRunner, ) from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_utils import ( @@ -80,7 +80,7 @@ logger = logging.getLogger(__name__) SGLANG_RETURN_ORIGINAL_LOGPROB = get_bool_env_var("SGLANG_RETURN_ORIGINAL_LOGPROB") -class MTPWorker(TpModelWorker): +class MultiLayerEagleWorker(TpModelWorker): def __init__( self, @@ -152,7 +152,7 @@ class MTPWorker(TpModelWorker): is_draft_worker=True, req_to_token_pool=self.req_to_token_pool, token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, - is_mtp_worker=True, + is_multi_layer_eagle=True, ) embed, head = self.target_worker.model_runner.model.get_embed_and_head() @@ -235,7 +235,7 @@ class MTPWorker(TpModelWorker): f"Capture draft extend cuda graph begin. This can take up to several minutes. avail mem={before_mem:.2f} GB" ) self.cuda_graph_runner_for_draft_extend_list.append( - MTPDraftExtendCudaGraphRunner(self, step) + MultiLayerEagleDraftExtendCudaGraphRunner(self, step) ) after_mem = get_available_gpu_memory(self.device, self.gpu_id) logger.info( diff --git a/python/sglang/srt/speculative/mtp_worker_v2.py b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py similarity index 98% rename from python/sglang/srt/speculative/mtp_worker_v2.py rename to python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py index f12780c59..425f78ac5 100644 --- a/python/sglang/srt/speculative/mtp_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -33,10 +33,10 @@ from sglang.srt.speculative.eagle_info_v2 import ( fill_new_verified_id, ) from sglang.srt.speculative.eagle_utils import TreeMaskMode, build_tree_kernel_efficient -from sglang.srt.speculative.mtp_draft_extend_cuda_graph_runner import ( - MTPMultiStepDraftExtendCudaGraphRunner, +from sglang.srt.speculative.multi_layer_eagle_draft_extend_cuda_graph_runner import ( + MultiLayerEagleMultiStepDraftExtendCudaGraphRunner, ) -from sglang.srt.speculative.mtp_utils import ( +from sglang.srt.speculative.multi_layer_eagle_utils import ( assign_hidden_states_pool_triton, rotate_input_ids_triton, ) @@ -66,7 +66,7 @@ def _get_plan_stream( return None, contextlib.nullcontext() -class MTPDraftWorker(BaseDraftWorker): +class MultiLayerEagleDraftWorker(BaseDraftWorker): def __init__( self, server_args: ServerArgs, @@ -125,7 +125,7 @@ class MTPDraftWorker(BaseDraftWorker): is_draft_worker=True, req_to_token_pool=self.req_to_token_pool, token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, - is_mtp_worker=True, + is_multi_layer_eagle=True, ) # Alias for better readability @@ -200,7 +200,7 @@ class MTPDraftWorker(BaseDraftWorker): return self.cuda_graph_runner_for_draft_extend = ( - MTPMultiStepDraftExtendCudaGraphRunner(self) + MultiLayerEagleMultiStepDraftExtendCudaGraphRunner(self) ) def reset_cuda_graph_buffers(self, forward_batch, batch_result): @@ -528,7 +528,7 @@ class MTPDraftWorker(BaseDraftWorker): ) -class MTPWorkerV2(BaseSpecWorker): +class MultiLayerEagleWorkerV2(BaseSpecWorker): def __init__( self, server_args: ServerArgs, @@ -560,7 +560,7 @@ class MTPWorkerV2(BaseSpecWorker): # Override the context length of the draft model to be the same as the target model. server_args.context_length = target_worker.model_runner.model_config.context_len - self._draft_worker = MTPDraftWorker( + self._draft_worker = MultiLayerEagleDraftWorker( server_args, gpu_id, tp_rank, dp_rank, moe_ep_rank, nccl_port, target_worker ) diff --git a/python/sglang/srt/speculative/spec_utils.py b/python/sglang/srt/speculative/spec_utils.py index cf2569b19..2e8057e8f 100644 --- a/python/sglang/srt/speculative/spec_utils.py +++ b/python/sglang/srt/speculative/spec_utils.py @@ -55,7 +55,7 @@ def spec_need_hidden_states(server_args: Optional[ServerArgs] = None) -> bool: server_args = get_global_server_args() # TODO(lsyin): also skip when 1) step = 1 or 2) standalone draft model - return not server_args.enable_mtp + return not server_args.enable_multi_layer_eagle @triton.jit