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:
co-authored by
gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
parent
b736a1525a
commit
3c882db3ad
+29
-23
@@ -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
|
||||
+5
-5
@@ -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(
|
||||
+8
-8
@@ -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
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user