[spec decoding] add extra attribute 'spec_hidden_size' (#23890)
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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 = (
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user