Tiny fix mimo model conflicts with main (#15483)

This commit is contained in:
Liangsheng Yin
2025-12-19 23:20:59 +08:00
committed by GitHub
parent 241ae17b25
commit 933cef16cc
2 changed files with 11 additions and 6 deletions
+2 -2
View File
@@ -494,7 +494,7 @@ class Scheduler(
tp_rank=self.tp_rank, tp_rank=self.tp_rank,
moe_ep_rank=self.moe_ep_rank, moe_ep_rank=self.moe_ep_rank,
server_args=self.server_args, server_args=self.server_args,
nccl_port=self.port_args.nccl_port, nccl_port=self.nccl_port,
target_worker=self.tp_worker, target_worker=self.tp_worker,
dp_rank=self.dp_rank, dp_rank=self.dp_rank,
) )
@@ -506,7 +506,7 @@ class Scheduler(
tp_rank=self.tp_rank, tp_rank=self.tp_rank,
moe_ep_rank=self.moe_ep_rank, moe_ep_rank=self.moe_ep_rank,
server_args=self.server_args, server_args=self.server_args,
nccl_port=self.port_args.nccl_port, nccl_port=self.nccl_port,
target_worker=self.tp_worker, target_worker=self.tp_worker,
dp_rank=self.dp_rank, dp_rank=self.dp_rank,
) )
@@ -14,7 +14,7 @@
import contextlib import contextlib
import logging import logging
from typing import List, Optional, Tuple from typing import TYPE_CHECKING, List, Optional, Tuple
import torch import torch
@@ -48,6 +48,10 @@ from sglang.srt.speculative.spec_utils import (
) )
from sglang.srt.utils.common import empty_context, fast_topk, next_power_of_2 from sglang.srt.utils.common import empty_context, fast_topk, next_power_of_2
if TYPE_CHECKING:
from sglang.srt.model_executor.model_runner import ModelRunnerOutput
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -376,8 +380,10 @@ class MTPDraftWorker(BaseDraftWorker):
topk_p_list = [] topk_p_list = []
topk_index_list = [] topk_index_list = []
for step in range(self.speculative_num_steps): for step in range(self.speculative_num_steps):
logits_output, _ = self.draft_runner_list[step].forward(forward_batch) output: ModelRunnerOutput = self.draft_runner_list[step].forward(
probs = torch.softmax(logits_output.next_token_logits, dim=-1) forward_batch
)
probs = torch.softmax(output.logits_output.next_token_logits, dim=-1)
topk_p, topk_index = fast_topk(probs, self.topk, dim=-1) topk_p, topk_index = fast_topk(probs, self.topk, dim=-1)
topk_p_list.append(topk_p) topk_p_list.append(topk_p)
topk_index_list.append(topk_index) topk_index_list.append(topk_index)
@@ -390,7 +396,6 @@ class MTPDraftWorker(BaseDraftWorker):
) )
next_draft_input.topk_p = torch.cat(topk_p_list, dim=1) next_draft_input.topk_p = torch.cat(topk_p_list, dim=1)
next_draft_input.topk_index = torch.cat(topk_index_list, dim=1) next_draft_input.topk_index = torch.cat(topk_index_list, dim=1)
# next_draft_input.hidden_states = logits_output.hidden_states
# Update req_to_hidden_states_pool for KV Cache reversion # Update req_to_hidden_states_pool for KV Cache reversion
if forward_batch.extend_seq_lens is not None: if forward_batch.extend_seq_lens is not None: