spec: route idle hidden_size via EagleDraft{,Extend}Input classmethods (#25013)

This commit is contained in:
Liangsheng Yin
2026-05-11 15:59:51 -07:00
committed by GitHub
parent ce1736fcc6
commit 6c3541a914
4 changed files with 23 additions and 13 deletions
@@ -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,
)