Tiny fix mimo model conflicts with main (#15483)
This commit is contained in:
@@ -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:
|
||||||
|
|||||||
Reference in New Issue
Block a user