From f57ec8d6ef6614482ea540dee843fbd57b9caf6c Mon Sep 17 00:00:00 2001 From: Qiaolin Yu Date: Wed, 29 Apr 2026 10:54:50 +0800 Subject: [PATCH] [spec decoding] add extra attribute 'spec_hidden_size' (#23890) --- python/sglang/srt/configs/model_config.py | 4 ++++ python/sglang/srt/managers/scheduler.py | 4 ++-- .../sglang/srt/speculative/eagle_draft_cuda_graph_runner.py | 2 +- .../srt/speculative/eagle_draft_extend_cuda_graph_runner.py | 5 ++++- python/sglang/srt/speculative/eagle_info.py | 4 ++-- python/sglang/srt/speculative/eagle_worker.py | 4 ++-- python/sglang/srt/speculative/eagle_worker_v2.py | 2 +- python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py | 2 +- 8 files changed, 17 insertions(+), 10 deletions(-) diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index e414172b7..d9d0e36d6 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -633,6 +633,10 @@ class ModelConfig: if self.num_key_value_heads is None: self.num_key_value_heads = self.num_attention_heads self.hidden_size = self.hf_text_config.hidden_size + hc_mult = getattr(self.hf_text_config, "hc_mult", 1) + self.spec_hidden_size = ( + self.hidden_size * hc_mult if hc_mult > 1 else self.hidden_size + ) self.num_hidden_layers = self.hf_text_config.num_hidden_layers self.num_attention_layers = self.num_hidden_layers if "LongcatFlashForCausalLM" in self.hf_config.architectures: diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 225b0ee66..1d3ad0b38 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -1089,7 +1089,7 @@ class Scheduler( self.disagg_metadata_buffers = MetadataBuffers( buffer_size, hidden_size=( - model_config.hidden_size + model_config.spec_hidden_size if self.spec_algorithm.is_eagle() else 16 # minimal padding size for RDMA ), @@ -1142,7 +1142,7 @@ class Scheduler( self.disagg_metadata_buffers = MetadataBuffers( buffer_size, hidden_size=( - model_config.hidden_size + model_config.spec_hidden_size if self.spec_algorithm.is_eagle() or self.spec_algorithm.is_standalone() else 16 # minimal padding size for RDMA diff --git a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py index d8439208c..a7e56eee7 100644 --- a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py @@ -128,7 +128,7 @@ class EAGLEDraftCudaGraphRunner: topk_p = torch.zeros((self.max_bs, self.topk), dtype=torch.float32) topk_index = torch.zeros((self.max_bs, self.topk), dtype=torch.int64) hidden_states = torch.zeros( - (self.max_bs, self.model_runner.model_config.hidden_size), + (self.max_bs, self.model_runner.model_config.spec_hidden_size), dtype=self.model_runner.dtype, ) diff --git a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py index 4f813d140..52e6de905 100644 --- a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py @@ -151,7 +151,10 @@ class EAGLEDraftExtendCudaGraphRunner: ) else: hidden_states = torch.zeros( - (self.max_num_token, self.model_runner.model_config.hidden_size), + ( + self.max_num_token, + self.model_runner.model_config.spec_hidden_size, + ), dtype=self.model_runner.dtype, ) self.seq_len_fill_value = ( diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py index f402b9cad..bf91cdeb7 100644 --- a/python/sglang/srt/speculative/eagle_info.py +++ b/python/sglang/srt/speculative/eagle_info.py @@ -242,7 +242,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): return EagleVerifyOutput( draft_input=EagleDraftInput.create_idle_input( device=batch.device, - hidden_size=batch.model_config.hidden_size, + hidden_size=batch.model_config.spec_hidden_size, dtype=batch.model_config.dtype, topk=self.topk, capture_hidden_mode=CaptureHiddenMode.LAST, @@ -624,7 +624,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): else: draft_input = EagleDraftInput.create_idle_input( device=batch.device, - hidden_size=batch.model_config.hidden_size, + hidden_size=batch.model_config.spec_hidden_size, dtype=batch.model_config.dtype, topk=self.topk, capture_hidden_mode=CaptureHiddenMode.LAST, diff --git a/python/sglang/srt/speculative/eagle_worker.py b/python/sglang/srt/speculative/eagle_worker.py index 52ecfb828..c63ec9a72 100644 --- a/python/sglang/srt/speculative/eagle_worker.py +++ b/python/sglang/srt/speculative/eagle_worker.py @@ -708,7 +708,7 @@ class EAGLEWorker(TpModelWorker): def _draft_preprocess_idle(self, batch: ScheduleBatch): batch.spec_info = EagleDraftInput.create_idle_input( device=self.device, - hidden_size=self.model_config.hidden_size, + hidden_size=self.model_config.spec_hidden_size, dtype=self.model_config.dtype, topk=self.topk, capture_hidden_mode=CaptureHiddenMode.LAST, @@ -1113,7 +1113,7 @@ class EAGLEWorker(TpModelWorker): self.model_config.hidden_size * 3 if self.speculative_algorithm.is_eagle3() and self.eagle_use_aux_hidden_state - else self.model_config.hidden_size + else self.model_config.spec_hidden_size ) batch.spec_info = EagleDraftInput.create_idle_input( device=self.device, diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index af883a455..4cc46e2c4 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -720,7 +720,7 @@ 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.hidden_size, + hidden_size=self.target_worker.model_config.spec_hidden_size, dtype=self.target_worker.model_config.dtype, topk=self.topk, capture_hidden_mode=CaptureHiddenMode.LAST, 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 b54dc3363..da2ba93fc 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -676,7 +676,7 @@ 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.hidden_size, + hidden_size=self.target_worker.model_config.spec_hidden_size, dtype=self.target_worker.model_config.dtype, topk=self.topk * self.speculative_num_steps, capture_hidden_mode=CaptureHiddenMode.LAST,