diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 27c9a14ed..e59986906 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -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, ) diff --git a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py index 18050f9e0..608bee11f 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py @@ -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: diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker.py b/python/sglang/srt/speculative/multi_layer_eagle_worker.py index c268168ba..871c96d22 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker.py @@ -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 diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py index 3d819717f..2befe4082 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -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, )