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]
parent b736a1525a
commit 3c882db3ad
10 changed files with 66 additions and 58 deletions
+3 -3
View File
@@ -100,7 +100,7 @@ class ModelConfig:
model_impl: Union[str, ModelImpl] = ModelImpl.AUTO, model_impl: Union[str, ModelImpl] = ModelImpl.AUTO,
sampling_defaults: str = "openai", sampling_defaults: str = "openai",
quantize_and_serve: bool = False, quantize_and_serve: bool = False,
is_mtp: bool = False, is_multi_layer_eagle: bool = False,
encoder_only: bool = False, encoder_only: bool = False,
language_only: bool = False, language_only: bool = False,
) -> None: ) -> None:
@@ -112,7 +112,7 @@ class ModelConfig:
self.model_impl = model_impl self.model_impl = model_impl
self.sampling_defaults = sampling_defaults self.sampling_defaults = sampling_defaults
self.quantize_and_serve = quantize_and_serve 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 # Validate quantize_and_serve configuration
self._validate_quantize_and_serve_config() self._validate_quantize_and_serve_config()
@@ -252,7 +252,7 @@ class ModelConfig:
sampling_defaults=server_args.sampling_defaults, sampling_defaults=server_args.sampling_defaults,
quantize_and_serve=server_args.quantize_and_serve, quantize_and_serve=server_args.quantize_and_serve,
override_config_file=server_args.decrypted_config_file, 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, language_only=server_args.language_only,
encoder_only=server_args.encoder_only, encoder_only=server_args.encoder_only,
is_draft_model=is_draft_model, is_draft_model=is_draft_model,
@@ -133,7 +133,7 @@ class ScheduleBatchDisaggregationDecodeMixin:
# Simulate the eagle run. # Simulate the eagle run.
if self.spec_algorithm.is_eagle(): if self.spec_algorithm.is_eagle():
num_states = server_args.speculative_eagle_topk 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 num_states *= server_args.speculative_num_steps
topk_p = torch.stack( topk_p = torch.stack(
[ [
+10 -7
View File
@@ -275,7 +275,6 @@ class Scheduler(
self.spec_algorithm = SpeculativeAlgorithm.from_string( self.spec_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm server_args.speculative_algorithm
) )
self.enable_mtp = server_args.enable_mtp
self.gpu_id = gpu_id self.gpu_id = gpu_id
self.page_size = server_args.page_size self.page_size = server_args.page_size
self.enable_hierarchical_cache = server_args.enable_hierarchical_cache self.enable_hierarchical_cache = server_args.enable_hierarchical_cache
@@ -485,11 +484,13 @@ class Scheduler(
draft_worker_kwargs["enable_overlap"] = self.enable_overlap draft_worker_kwargs["enable_overlap"] = self.enable_overlap
# FIXME: refactor the draft worker registration logic # FIXME: refactor the draft worker registration logic
if self.enable_mtp: if self.server_args.enable_multi_layer_eagle:
if self.enable_overlap: 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, gpu_id=self.gpu_id,
tp_rank=self.tp_rank, tp_rank=self.tp_rank,
moe_ep_rank=self.moe_ep_rank, moe_ep_rank=self.moe_ep_rank,
@@ -499,9 +500,11 @@ class Scheduler(
dp_rank=self.dp_rank, dp_rank=self.dp_rank,
) )
else: 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, gpu_id=self.gpu_id,
tp_rank=self.tp_rank, tp_rank=self.tp_rank,
moe_ep_rank=self.moe_ep_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(): if self.draft_worker is None or self.spec_algorithm.is_ngram():
draft_token_to_kv_pool = None draft_token_to_kv_pool = None
elif self.spec_algorithm.is_eagle() and self.enable_overlap: 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] draft_runner = self.draft_worker.draft_worker.draft_runner_list[0]
else: else:
draft_runner = self.draft_worker.draft_worker.draft_runner 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, is_draft_worker: bool = False,
req_to_token_pool: Optional[ReqToTokenPool] = None, req_to_token_pool: Optional[ReqToTokenPool] = None,
token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator] = None, token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator] = None,
is_mtp_worker: bool = False, is_multi_layer_eagle: bool = False,
): ):
# Parse args # Parse args
self.tp_size = server_args.tp_size self.tp_size = server_args.tp_size
@@ -266,9 +266,9 @@ class TpModelWorker(BaseTpWorker):
is_draft_worker=is_draft_worker, is_draft_worker=is_draft_worker,
req_to_token_pool=req_to_token_pool, req_to_token_pool=req_to_token_pool,
token_to_kv_pool_allocator=token_to_kv_pool_allocator, 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) self.model_runner_list.append(self.model_runner)
for i in range(1, server_args.speculative_num_steps): for i in range(1, server_args.speculative_num_steps):
self.model_runner_list.append( self.model_runner_list.append(
+6 -7
View File
@@ -439,10 +439,7 @@ class ServerArgs:
speculative_ngram_match_type: Literal["BFS", "PROB"] = "BFS" speculative_ngram_match_type: Literal["BFS", "PROB"] = "BFS"
speculative_ngram_branch_length: int = 18 speculative_ngram_branch_length: int = 18
speculative_ngram_capacity: int = 10 * 1000 * 1000 speculative_ngram_capacity: int = 10 * 1000 * 1000
enable_multi_layer_eagle: bool = False
# For Multi-Layer MTP
# FIXME: rename -> enable_multi_layer_mtp
enable_mtp: bool = False
# Expert parallelism # Expert parallelism
ep_size: int = 1 ep_size: int = 1
@@ -1189,6 +1186,8 @@ class ServerArgs:
self.disable_hybrid_swa_memory = True self.disable_hybrid_swa_memory = True
elif "MiMoV2FlashForCausalLM" in model_arch: 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 self.swa_full_tokens_ratio = 1.0
logger.warning( logger.warning(
"Reset swa_full_tokens_ratio to 1.0 for MiMoV2FlashForCausalLM model" "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.", help="The cache capacity for ngram speculative decoding.",
) )
# Speculative decoding (MTP) # Multi-layer Eagle speculative decoding
parser.add_argument( parser.add_argument(
"--enable-mtp", "--enable-multi-layer-eagle",
action="store_true", action="store_true",
help="Enable multi-layer MTP speculative decoding.", help="Enable multi-layer Eagle speculative decoding.",
) )
# Expert parallelism # Expert parallelism
@@ -40,7 +40,7 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardMode, ForwardMode,
) )
from sglang.srt.speculative.eagle_info import EagleDraftInput 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.speculative.spec_utils import fast_topk
from sglang.srt.utils import ( from sglang.srt.utils import (
get_available_gpu_memory, get_available_gpu_memory,
@@ -51,18 +51,20 @@ from sglang.srt.utils import (
) )
if TYPE_CHECKING: 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__) logger = logging.getLogger(__name__)
class MTPDraftExtendCudaGraphRunner: class MultiLayerEagleDraftExtendCudaGraphRunner:
def __init__(self, mtp_worker: MTPDraftWorker, step: int): def __init__(self, eagle_worker: MultiLayerEagleDraftWorker, step: int):
# Parse args # Parse args
self.step = step self.step = step
self.mtp_worker = mtp_worker self.eagle_worker = eagle_worker
self.model_runner = model_runner = mtp_worker.mtp_model_runner(self.step) self.model_runner = model_runner = eagle_worker.mtp_model_runner(self.step)
self.forward_mode = ForwardMode.DRAFT_EXTEND_V2 self.forward_mode = ForwardMode.DRAFT_EXTEND_V2
self.graphs = {} self.graphs = {}
@@ -93,10 +95,10 @@ class MTPDraftExtendCudaGraphRunner:
self.max_bs = max(self.capture_bs) self.max_bs = max(self.capture_bs)
self.max_num_token = self.max_bs * self.num_tokens_per_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.eagle_worker.draft_extend_attn_backend_list[
self.max_bs, self.max_num_token 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.seq_len_fill_value = self.eagle_worker.draft_extend_attn_backend_list[
self.step self.step
].get_cuda_graph_seq_len_fill_value() ].get_cuda_graph_seq_len_fill_value()
@@ -331,7 +333,7 @@ class MTPDraftExtendCudaGraphRunner:
spec_algorithm=self.model_runner.spec_algorithm, spec_algorithm=self.model_runner.spec_algorithm,
spec_info=spec_info, spec_info=spec_info,
capture_hidden_mode=CaptureHiddenMode.FULL, 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=extend_seq_lens,
extend_seq_lens_cpu=extend_seq_lens_cpu, extend_seq_lens_cpu=extend_seq_lens_cpu,
padded_static_len=self.padded_static_len, padded_static_len=self.padded_static_len,
@@ -352,7 +354,7 @@ class MTPDraftExtendCudaGraphRunner:
num_tokens = bs * self.num_tokens_per_bs num_tokens = bs * self.num_tokens_per_bs
forward_batch = self.get_forward_batch(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 self.step
].init_forward_metadata_capture_cuda_graph( ].init_forward_metadata_capture_cuda_graph(
bs=bs, bs=bs,
@@ -420,7 +422,7 @@ class MTPDraftExtendCudaGraphRunner:
self.step, self.step,
forward_batch.req_pool_indices, forward_batch.req_pool_indices,
forward_batch.req_to_token_pool.req_to_token, 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.next_cuda_graph_runner.swa_out_cache_loc.copy_(
self.model_runner.token_to_kv_pool.translate_loc_from_full_to_swa( 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.positions = self.positions[:num_tokens]
forward_batch.spec_info.extend_seq_lens_tensor = self.extend_seq_lens[:bs] 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 self.step
].init_forward_metadata_replay_cuda_graph( ].init_forward_metadata_replay_cuda_graph(
bs=bs, bs=bs,
@@ -540,13 +542,15 @@ class MTPDraftExtendCudaGraphRunner:
return out return out
class MTPMultiStepDraftExtendCudaGraphRunner: class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
def __init__(self, mtp_worker: MTPDraftWorker): def __init__(self, eagle_worker: MultiLayerEagleDraftWorker):
self.mtp_worker = mtp_worker self.eagle_worker = eagle_worker
self.device = mtp_worker.device self.device = eagle_worker.device
self.gpu_id = mtp_worker.gpu_id self.gpu_id = eagle_worker.gpu_id
self.speculative_num_steps = mtp_worker.speculative_num_steps self.speculative_num_steps = eagle_worker.speculative_num_steps
self.draft_extend_attn_backend_list = mtp_worker.draft_extend_attn_backend_list self.draft_extend_attn_backend_list = (
eagle_worker.draft_extend_attn_backend_list
)
self.runners = [] self.runners = []
self.cuda_graph_buffers = {} self.cuda_graph_buffers = {}
@@ -557,7 +561,7 @@ class MTPMultiStepDraftExtendCudaGraphRunner:
self._init_and_capture() self._init_and_capture()
def _init_and_capture(self): 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 self.runners = [None] * self.speculative_num_steps
return return
@@ -567,7 +571,9 @@ class MTPMultiStepDraftExtendCudaGraphRunner:
# 1. Capture loop # 1. Capture loop
for step in range(self.speculative_num_steps): for step in range(self.speculative_num_steps):
if self.draft_extend_attn_backend_list[step]: 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.runners.append(runner)
self.seq_len_fill_value = runner.seq_len_fill_value self.seq_len_fill_value = runner.seq_len_fill_value
@@ -48,8 +48,8 @@ from sglang.srt.speculative.eagle_utils import (
organize_draft_results, organize_draft_results,
) )
from sglang.srt.speculative.eagle_worker import get_last_loc_large_page_size_top_k_1 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 ( from sglang.srt.speculative.multi_layer_eagle_draft_extend_cuda_graph_runner import (
MTPDraftExtendCudaGraphRunner, MultiLayerEagleDraftExtendCudaGraphRunner,
) )
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.speculative.spec_utils import ( 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") SGLANG_RETURN_ORIGINAL_LOGPROB = get_bool_env_var("SGLANG_RETURN_ORIGINAL_LOGPROB")
class MTPWorker(TpModelWorker): class MultiLayerEagleWorker(TpModelWorker):
def __init__( def __init__(
self, self,
@@ -152,7 +152,7 @@ class MTPWorker(TpModelWorker):
is_draft_worker=True, is_draft_worker=True,
req_to_token_pool=self.req_to_token_pool, req_to_token_pool=self.req_to_token_pool,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, 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() 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" 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( 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) after_mem = get_available_gpu_memory(self.device, self.gpu_id)
logger.info( logger.info(
@@ -33,10 +33,10 @@ from sglang.srt.speculative.eagle_info_v2 import (
fill_new_verified_id, fill_new_verified_id,
) )
from sglang.srt.speculative.eagle_utils import TreeMaskMode, build_tree_kernel_efficient from sglang.srt.speculative.eagle_utils import TreeMaskMode, build_tree_kernel_efficient
from sglang.srt.speculative.mtp_draft_extend_cuda_graph_runner import ( from sglang.srt.speculative.multi_layer_eagle_draft_extend_cuda_graph_runner import (
MTPMultiStepDraftExtendCudaGraphRunner, MultiLayerEagleMultiStepDraftExtendCudaGraphRunner,
) )
from sglang.srt.speculative.mtp_utils import ( from sglang.srt.speculative.multi_layer_eagle_utils import (
assign_hidden_states_pool_triton, assign_hidden_states_pool_triton,
rotate_input_ids_triton, rotate_input_ids_triton,
) )
@@ -66,7 +66,7 @@ def _get_plan_stream(
return None, contextlib.nullcontext() return None, contextlib.nullcontext()
class MTPDraftWorker(BaseDraftWorker): class MultiLayerEagleDraftWorker(BaseDraftWorker):
def __init__( def __init__(
self, self,
server_args: ServerArgs, server_args: ServerArgs,
@@ -125,7 +125,7 @@ class MTPDraftWorker(BaseDraftWorker):
is_draft_worker=True, is_draft_worker=True,
req_to_token_pool=self.req_to_token_pool, req_to_token_pool=self.req_to_token_pool,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator, token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
is_mtp_worker=True, is_multi_layer_eagle=True,
) )
# Alias for better readability # Alias for better readability
@@ -200,7 +200,7 @@ class MTPDraftWorker(BaseDraftWorker):
return return
self.cuda_graph_runner_for_draft_extend = ( self.cuda_graph_runner_for_draft_extend = (
MTPMultiStepDraftExtendCudaGraphRunner(self) MultiLayerEagleMultiStepDraftExtendCudaGraphRunner(self)
) )
def reset_cuda_graph_buffers(self, forward_batch, batch_result): 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__( def __init__(
self, self,
server_args: ServerArgs, 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. # 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 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 server_args, gpu_id, tp_rank, dp_rank, moe_ep_rank, nccl_port, target_worker
) )
+1 -1
View File
@@ -55,7 +55,7 @@ def spec_need_hidden_states(server_args: Optional[ServerArgs] = None) -> bool:
server_args = get_global_server_args() server_args = get_global_server_args()
# TODO(lsyin): also skip when 1) step = 1 or 2) standalone draft model # 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 @triton.jit