[spec decoding] add extra attribute 'spec_hidden_size' (#23890)

This commit is contained in:
Qiaolin Yu
2026-04-28 19:54:50 -07:00
committed by GitHub
parent 2a771a40ac
commit f57ec8d6ef
8 changed files with 17 additions and 10 deletions
@@ -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:
+2 -2
View File
@@ -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
@@ -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,
)
@@ -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 = (
+2 -2
View File
@@ -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,
@@ -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,
@@ -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,
@@ -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,