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:
|
||||
model_worker_batch.spec_info = EagleDraftInput.create_idle_input(
|
||||
device=self.device,
|
||||
hidden_size=self.target_worker.model_config.spec_hidden_size,
|
||||
dtype=self.target_worker.model_config.dtype,
|
||||
hidden_size=EagleDraftInput.hidden_size_for(self.draft_worker),
|
||||
dtype=EagleDraftInput.dtype_for(self.draft_worker),
|
||||
topk=self.topk,
|
||||
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||
)
|
||||
|
||||
@@ -178,8 +178,11 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
|
||||
mrope_positions = torch.zeros((3, self.max_num_token), dtype=torch.int64)
|
||||
|
||||
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:
|
||||
|
||||
@@ -147,6 +147,15 @@ class MultiLayerEagleWorker(TpModelWorker):
|
||||
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()
|
||||
|
||||
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:
|
||||
batch = batch.copy()
|
||||
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(
|
||||
device=self.device,
|
||||
hidden_size=hidden_size,
|
||||
dtype=self.model_config.dtype,
|
||||
hidden_size=EagleDraftExtendInput.hidden_size_for(self),
|
||||
dtype=EagleDraftExtendInput.dtype_for(self),
|
||||
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||
)
|
||||
batch.spec_info = draft_extend_input
|
||||
|
||||
@@ -138,6 +138,9 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker):
|
||||
|
||||
# Alias for better readability
|
||||
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
|
||||
# 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:
|
||||
model_worker_batch.spec_info = EagleDraftInput.create_idle_input(
|
||||
device=self.device,
|
||||
hidden_size=self.target_worker.model_config.spec_hidden_size,
|
||||
dtype=self.target_worker.model_config.dtype,
|
||||
hidden_size=EagleDraftInput.hidden_size_for(self.draft_worker),
|
||||
dtype=EagleDraftInput.dtype_for(self.draft_worker),
|
||||
topk=self.topk * self.speculative_num_steps,
|
||||
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user