[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:
|
if self.num_key_value_heads is None:
|
||||||
self.num_key_value_heads = self.num_attention_heads
|
self.num_key_value_heads = self.num_attention_heads
|
||||||
self.hidden_size = self.hf_text_config.hidden_size
|
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_hidden_layers = self.hf_text_config.num_hidden_layers
|
||||||
self.num_attention_layers = self.num_hidden_layers
|
self.num_attention_layers = self.num_hidden_layers
|
||||||
if "LongcatFlashForCausalLM" in self.hf_config.architectures:
|
if "LongcatFlashForCausalLM" in self.hf_config.architectures:
|
||||||
|
|||||||
@@ -1089,7 +1089,7 @@ class Scheduler(
|
|||||||
self.disagg_metadata_buffers = MetadataBuffers(
|
self.disagg_metadata_buffers = MetadataBuffers(
|
||||||
buffer_size,
|
buffer_size,
|
||||||
hidden_size=(
|
hidden_size=(
|
||||||
model_config.hidden_size
|
model_config.spec_hidden_size
|
||||||
if self.spec_algorithm.is_eagle()
|
if self.spec_algorithm.is_eagle()
|
||||||
else 16 # minimal padding size for RDMA
|
else 16 # minimal padding size for RDMA
|
||||||
),
|
),
|
||||||
@@ -1142,7 +1142,7 @@ class Scheduler(
|
|||||||
self.disagg_metadata_buffers = MetadataBuffers(
|
self.disagg_metadata_buffers = MetadataBuffers(
|
||||||
buffer_size,
|
buffer_size,
|
||||||
hidden_size=(
|
hidden_size=(
|
||||||
model_config.hidden_size
|
model_config.spec_hidden_size
|
||||||
if self.spec_algorithm.is_eagle()
|
if self.spec_algorithm.is_eagle()
|
||||||
or self.spec_algorithm.is_standalone()
|
or self.spec_algorithm.is_standalone()
|
||||||
else 16 # minimal padding size for RDMA
|
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_p = torch.zeros((self.max_bs, self.topk), dtype=torch.float32)
|
||||||
topk_index = torch.zeros((self.max_bs, self.topk), dtype=torch.int64)
|
topk_index = torch.zeros((self.max_bs, self.topk), dtype=torch.int64)
|
||||||
hidden_states = torch.zeros(
|
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,
|
dtype=self.model_runner.dtype,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -151,7 +151,10 @@ class EAGLEDraftExtendCudaGraphRunner:
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
hidden_states = torch.zeros(
|
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,
|
dtype=self.model_runner.dtype,
|
||||||
)
|
)
|
||||||
self.seq_len_fill_value = (
|
self.seq_len_fill_value = (
|
||||||
|
|||||||
@@ -242,7 +242,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
|||||||
return EagleVerifyOutput(
|
return EagleVerifyOutput(
|
||||||
draft_input=EagleDraftInput.create_idle_input(
|
draft_input=EagleDraftInput.create_idle_input(
|
||||||
device=batch.device,
|
device=batch.device,
|
||||||
hidden_size=batch.model_config.hidden_size,
|
hidden_size=batch.model_config.spec_hidden_size,
|
||||||
dtype=batch.model_config.dtype,
|
dtype=batch.model_config.dtype,
|
||||||
topk=self.topk,
|
topk=self.topk,
|
||||||
capture_hidden_mode=CaptureHiddenMode.LAST,
|
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||||
@@ -624,7 +624,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
|||||||
else:
|
else:
|
||||||
draft_input = EagleDraftInput.create_idle_input(
|
draft_input = EagleDraftInput.create_idle_input(
|
||||||
device=batch.device,
|
device=batch.device,
|
||||||
hidden_size=batch.model_config.hidden_size,
|
hidden_size=batch.model_config.spec_hidden_size,
|
||||||
dtype=batch.model_config.dtype,
|
dtype=batch.model_config.dtype,
|
||||||
topk=self.topk,
|
topk=self.topk,
|
||||||
capture_hidden_mode=CaptureHiddenMode.LAST,
|
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||||
|
|||||||
@@ -708,7 +708,7 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
def _draft_preprocess_idle(self, batch: ScheduleBatch):
|
def _draft_preprocess_idle(self, batch: ScheduleBatch):
|
||||||
batch.spec_info = EagleDraftInput.create_idle_input(
|
batch.spec_info = EagleDraftInput.create_idle_input(
|
||||||
device=self.device,
|
device=self.device,
|
||||||
hidden_size=self.model_config.hidden_size,
|
hidden_size=self.model_config.spec_hidden_size,
|
||||||
dtype=self.model_config.dtype,
|
dtype=self.model_config.dtype,
|
||||||
topk=self.topk,
|
topk=self.topk,
|
||||||
capture_hidden_mode=CaptureHiddenMode.LAST,
|
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||||
@@ -1113,7 +1113,7 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
self.model_config.hidden_size * 3
|
self.model_config.hidden_size * 3
|
||||||
if self.speculative_algorithm.is_eagle3()
|
if self.speculative_algorithm.is_eagle3()
|
||||||
and self.eagle_use_aux_hidden_state
|
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(
|
batch.spec_info = EagleDraftInput.create_idle_input(
|
||||||
device=self.device,
|
device=self.device,
|
||||||
|
|||||||
@@ -720,7 +720,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
if model_worker_batch.spec_info is None:
|
if model_worker_batch.spec_info is None:
|
||||||
model_worker_batch.spec_info = EagleDraftInput.create_idle_input(
|
model_worker_batch.spec_info = EagleDraftInput.create_idle_input(
|
||||||
device=self.device,
|
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,
|
dtype=self.target_worker.model_config.dtype,
|
||||||
topk=self.topk,
|
topk=self.topk,
|
||||||
capture_hidden_mode=CaptureHiddenMode.LAST,
|
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||||
|
|||||||
@@ -676,7 +676,7 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
|
|||||||
if model_worker_batch.spec_info is None:
|
if model_worker_batch.spec_info is None:
|
||||||
model_worker_batch.spec_info = EagleDraftInput.create_idle_input(
|
model_worker_batch.spec_info = EagleDraftInput.create_idle_input(
|
||||||
device=self.device,
|
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,
|
dtype=self.target_worker.model_config.dtype,
|
||||||
topk=self.topk * self.speculative_num_steps,
|
topk=self.topk * self.speculative_num_steps,
|
||||||
capture_hidden_mode=CaptureHiddenMode.LAST,
|
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||||
|
|||||||
Reference in New Issue
Block a user