Tiny fix mimo model conflicts with main (#15483)
This commit is contained in:
@@ -494,7 +494,7 @@ class Scheduler(
|
||||
tp_rank=self.tp_rank,
|
||||
moe_ep_rank=self.moe_ep_rank,
|
||||
server_args=self.server_args,
|
||||
nccl_port=self.port_args.nccl_port,
|
||||
nccl_port=self.nccl_port,
|
||||
target_worker=self.tp_worker,
|
||||
dp_rank=self.dp_rank,
|
||||
)
|
||||
@@ -506,7 +506,7 @@ class Scheduler(
|
||||
tp_rank=self.tp_rank,
|
||||
moe_ep_rank=self.moe_ep_rank,
|
||||
server_args=self.server_args,
|
||||
nccl_port=self.port_args.nccl_port,
|
||||
nccl_port=self.nccl_port,
|
||||
target_worker=self.tp_worker,
|
||||
dp_rank=self.dp_rank,
|
||||
)
|
||||
|
||||
@@ -14,7 +14,7 @@
|
||||
|
||||
import contextlib
|
||||
import logging
|
||||
from typing import List, Optional, Tuple
|
||||
from typing import TYPE_CHECKING, List, Optional, Tuple
|
||||
|
||||
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
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.model_executor.model_runner import ModelRunnerOutput
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -376,8 +380,10 @@ class MTPDraftWorker(BaseDraftWorker):
|
||||
topk_p_list = []
|
||||
topk_index_list = []
|
||||
for step in range(self.speculative_num_steps):
|
||||
logits_output, _ = self.draft_runner_list[step].forward(forward_batch)
|
||||
probs = torch.softmax(logits_output.next_token_logits, dim=-1)
|
||||
output: ModelRunnerOutput = self.draft_runner_list[step].forward(
|
||||
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_list.append(topk_p)
|
||||
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_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
|
||||
if forward_batch.extend_seq_lens is not None:
|
||||
|
||||
Reference in New Issue
Block a user