spec: route idle hidden_size via EagleDraft{,Extend}Input classmethods (#25013)
This commit is contained in:
@@ -763,8 +763,8 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
if model_worker_batch.spec_info is None:
|
if model_worker_batch.spec_info is None:
|
||||||
model_worker_batch.spec_info = EagleDraftInput.create_idle_input(
|
model_worker_batch.spec_info = EagleDraftInput.create_idle_input(
|
||||||
device=self.device,
|
device=self.device,
|
||||||
hidden_size=self.target_worker.model_config.spec_hidden_size,
|
hidden_size=EagleDraftInput.hidden_size_for(self.draft_worker),
|
||||||
dtype=self.target_worker.model_config.dtype,
|
dtype=EagleDraftInput.dtype_for(self.draft_worker),
|
||||||
topk=self.topk,
|
topk=self.topk,
|
||||||
capture_hidden_mode=CaptureHiddenMode.LAST,
|
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -178,8 +178,11 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
|
|||||||
mrope_positions = torch.zeros((3, self.max_num_token), dtype=torch.int64)
|
mrope_positions = torch.zeros((3, self.max_num_token), dtype=torch.int64)
|
||||||
|
|
||||||
hidden_states = torch.zeros(
|
hidden_states = torch.zeros(
|
||||||
(self.max_num_token, self.model_runner.model_config.hidden_size),
|
(
|
||||||
dtype=self.model_runner.dtype,
|
self.max_num_token,
|
||||||
|
EagleDraftExtendInput.hidden_size_for(self.eagle_worker),
|
||||||
|
),
|
||||||
|
dtype=EagleDraftExtendInput.dtype_for(self.eagle_worker),
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.require_gathered_buffer:
|
if self.require_gathered_buffer:
|
||||||
|
|||||||
@@ -147,6 +147,15 @@ class MultiLayerEagleWorker(TpModelWorker):
|
|||||||
is_multi_layer_eagle=True,
|
is_multi_layer_eagle=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
self.eagle_use_aux_hidden_state = False
|
||||||
|
if self.speculative_algorithm.is_eagle3():
|
||||||
|
eagle_config = getattr(
|
||||||
|
self.model_runner.model_config.hf_config, "eagle_config", {}
|
||||||
|
)
|
||||||
|
self.eagle_use_aux_hidden_state = eagle_config.get(
|
||||||
|
"use_aux_hidden_state", 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()
|
||||||
|
|
||||||
if self.speculative_algorithm.is_eagle3():
|
if self.speculative_algorithm.is_eagle3():
|
||||||
@@ -677,15 +686,10 @@ class MultiLayerEagleWorker(TpModelWorker):
|
|||||||
if not input_is_idle and draft_extend_input.input_ids.shape[0] == 0:
|
if not input_is_idle and draft_extend_input.input_ids.shape[0] == 0:
|
||||||
batch = batch.copy()
|
batch = batch.copy()
|
||||||
batch.prepare_for_idle()
|
batch.prepare_for_idle()
|
||||||
hidden_size = (
|
|
||||||
self.model_config.hidden_size * 3
|
|
||||||
if self.speculative_algorithm.is_eagle3()
|
|
||||||
else self.model_config.hidden_size
|
|
||||||
)
|
|
||||||
draft_extend_input = EagleDraftExtendInput.create_idle_input(
|
draft_extend_input = EagleDraftExtendInput.create_idle_input(
|
||||||
device=self.device,
|
device=self.device,
|
||||||
hidden_size=hidden_size,
|
hidden_size=EagleDraftExtendInput.hidden_size_for(self),
|
||||||
dtype=self.model_config.dtype,
|
dtype=EagleDraftExtendInput.dtype_for(self),
|
||||||
capture_hidden_mode=CaptureHiddenMode.LAST,
|
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||||
)
|
)
|
||||||
batch.spec_info = draft_extend_input
|
batch.spec_info = draft_extend_input
|
||||||
|
|||||||
@@ -138,6 +138,9 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker):
|
|||||||
|
|
||||||
# Alias for better readability
|
# Alias for better readability
|
||||||
self.draft_runner_list: List[ModelRunner] = self.draft_worker.model_runner_list
|
self.draft_runner_list: List[ModelRunner] = self.draft_worker.model_runner_list
|
||||||
|
# Match `EagleDraftWorker.draft_runner` so `_draft_runner_of(self)` works
|
||||||
|
# for the EagleDraftInput shape classmethods.
|
||||||
|
self.draft_runner: ModelRunner = self.draft_runner_list[0]
|
||||||
|
|
||||||
# Chain-style MTP: each step propagates its own output hidden states to the
|
# Chain-style MTP: each step propagates its own output hidden states to the
|
||||||
# next step. Non-chain: each step uses the target model's hidden states.
|
# next step. Non-chain: each step uses the target model's hidden states.
|
||||||
@@ -677,8 +680,8 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
|
|||||||
if model_worker_batch.spec_info is None:
|
if model_worker_batch.spec_info is None:
|
||||||
model_worker_batch.spec_info = EagleDraftInput.create_idle_input(
|
model_worker_batch.spec_info = EagleDraftInput.create_idle_input(
|
||||||
device=self.device,
|
device=self.device,
|
||||||
hidden_size=self.target_worker.model_config.spec_hidden_size,
|
hidden_size=EagleDraftInput.hidden_size_for(self.draft_worker),
|
||||||
dtype=self.target_worker.model_config.dtype,
|
dtype=EagleDraftInput.dtype_for(self.draft_worker),
|
||||||
topk=self.topk * self.speculative_num_steps,
|
topk=self.topk * self.speculative_num_steps,
|
||||||
capture_hidden_mode=CaptureHiddenMode.LAST,
|
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user