Tiny rename for spec related fileds. (#18468)
This commit is contained in:
@@ -2136,7 +2136,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
default_extend = getattr(
|
default_extend = getattr(
|
||||||
spec_info, "num_tokens_per_batch", self.speculative_num_steps + 1
|
spec_info, "num_tokens_per_req", self.speculative_num_steps + 1
|
||||||
)
|
)
|
||||||
extend_seq_lens = torch.full(
|
extend_seq_lens = torch.full(
|
||||||
(bs,), default_extend, dtype=torch.int32, device=device
|
(bs,), default_extend, dtype=torch.int32, device=device
|
||||||
@@ -2147,7 +2147,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
metadata.max_seq_len_q = int(max(extend_seq_lens_cpu))
|
metadata.max_seq_len_q = int(max(extend_seq_lens_cpu))
|
||||||
else:
|
else:
|
||||||
metadata.max_seq_len_q = getattr(
|
metadata.max_seq_len_q = getattr(
|
||||||
spec_info, "num_tokens_per_batch", self.speculative_num_steps + 1
|
spec_info, "num_tokens_per_req", self.speculative_num_steps + 1
|
||||||
)
|
)
|
||||||
|
|
||||||
metadata.cu_seqlens_q[1:].copy_(
|
metadata.cu_seqlens_q[1:].copy_(
|
||||||
|
|||||||
@@ -799,7 +799,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
setattr(self, "_original_batch_size", self.batch_size)
|
setattr(self, "_original_batch_size", self.batch_size)
|
||||||
if self.spec_info is not None:
|
if self.spec_info is not None:
|
||||||
bs = self.batch_size = (
|
bs = self.batch_size = (
|
||||||
num_tokens // self.spec_info.num_tokens_per_batch
|
num_tokens // self.spec_info.num_tokens_per_req
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
bs = self.batch_size = num_tokens
|
bs = self.batch_size = num_tokens
|
||||||
@@ -935,7 +935,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
logits_output.next_token_logits = logits_output.next_token_logits[:bs]
|
logits_output.next_token_logits = logits_output.next_token_logits[:bs]
|
||||||
logits_output.hidden_states = logits_output.hidden_states[:bs]
|
logits_output.hidden_states = logits_output.hidden_states[:bs]
|
||||||
elif self.forward_mode.is_draft_extend_v2(): # draft extend_v2
|
elif self.forward_mode.is_draft_extend_v2(): # draft extend_v2
|
||||||
bs = bs * self.spec_info.num_tokens_per_batch
|
bs = bs * self.spec_info.num_tokens_per_req
|
||||||
logits_output.next_token_logits = logits_output.next_token_logits[:bs]
|
logits_output.next_token_logits = logits_output.next_token_logits[:bs]
|
||||||
logits_output.hidden_states = logits_output.hidden_states[:bs]
|
logits_output.hidden_states = logits_output.hidden_states[:bs]
|
||||||
elif self.forward_mode.is_extend() or self.forward_mode.is_idle():
|
elif self.forward_mode.is_extend() or self.forward_mode.is_idle():
|
||||||
|
|||||||
@@ -69,7 +69,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
|||||||
grammar: BaseGrammarObject = None
|
grammar: BaseGrammarObject = None
|
||||||
|
|
||||||
# Shape info for padding
|
# Shape info for padding
|
||||||
num_tokens_per_batch: int = -1
|
num_tokens_per_req: int = -1
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
super().__init__(SpecInputType.EAGLE_VERIFY)
|
super().__init__(SpecInputType.EAGLE_VERIFY)
|
||||||
@@ -634,8 +634,8 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
|
|||||||
kv_indices: torch.Tensor = None
|
kv_indices: torch.Tensor = None
|
||||||
|
|
||||||
# Shape info for padding
|
# Shape info for padding
|
||||||
num_tokens_per_batch: int = -1
|
num_tokens_per_req: int = -1
|
||||||
num_tokens_for_logprob_per_batch: int = -1
|
num_tokens_for_logprob_per_req: int = -1
|
||||||
|
|
||||||
# Inputs for draft extend
|
# Inputs for draft extend
|
||||||
# shape: (b,)
|
# shape: (b,)
|
||||||
@@ -652,7 +652,7 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
|
|||||||
super().__init__(SpecInputType.EAGLE_DRAFT)
|
super().__init__(SpecInputType.EAGLE_DRAFT)
|
||||||
|
|
||||||
def get_spec_adjust_token_coefficient(self) -> Tuple[int, int]:
|
def get_spec_adjust_token_coefficient(self) -> Tuple[int, int]:
|
||||||
return self.num_tokens_per_batch, self.num_tokens_for_logprob_per_batch
|
return self.num_tokens_per_req, self.num_tokens_for_logprob_per_req
|
||||||
|
|
||||||
def prepare_for_extend(self, batch: ScheduleBatch):
|
def prepare_for_extend(self, batch: ScheduleBatch):
|
||||||
|
|
||||||
|
|||||||
@@ -169,8 +169,8 @@ class EagleDraftInputV2Mixin:
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Get a forward batch
|
# Get a forward batch
|
||||||
self.num_tokens_per_batch = topk
|
self.num_tokens_per_req = topk
|
||||||
self.num_tokens_for_logprob_per_batch = topk
|
self.num_tokens_for_logprob_per_req = topk
|
||||||
batch.capture_hidden_mode = CaptureHiddenMode.LAST
|
batch.capture_hidden_mode = CaptureHiddenMode.LAST
|
||||||
self.positions = batch.seq_lens.repeat_interleave(topk, dim=0)
|
self.positions = batch.seq_lens.repeat_interleave(topk, dim=0)
|
||||||
forward_batch = ForwardBatch.init_new(batch, draft_model_runner)
|
forward_batch = ForwardBatch.init_new(batch, draft_model_runner)
|
||||||
|
|||||||
@@ -532,8 +532,8 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
assert isinstance(spec_info, EagleDraftInput)
|
assert isinstance(spec_info, EagleDraftInput)
|
||||||
|
|
||||||
spec_info.capture_hidden_mode = CaptureHiddenMode.LAST
|
spec_info.capture_hidden_mode = CaptureHiddenMode.LAST
|
||||||
spec_info.num_tokens_per_batch = self.topk
|
spec_info.num_tokens_per_req = self.topk
|
||||||
spec_info.num_tokens_for_logprob_per_batch = self.topk
|
spec_info.num_tokens_for_logprob_per_req = self.topk
|
||||||
batch.return_hidden_states = False
|
batch.return_hidden_states = False
|
||||||
|
|
||||||
# Get forward batch
|
# Get forward batch
|
||||||
@@ -683,7 +683,7 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
def verify(self, batch: ScheduleBatch, spec_info: EagleVerifyInput):
|
def verify(self, batch: ScheduleBatch, spec_info: EagleVerifyInput):
|
||||||
seq_lens_pre_verify = batch.seq_lens.clone()
|
seq_lens_pre_verify = batch.seq_lens.clone()
|
||||||
spec_info.prepare_for_verify(batch, self.page_size)
|
spec_info.prepare_for_verify(batch, self.page_size)
|
||||||
spec_info.num_tokens_per_batch = self.speculative_num_steps + 1
|
spec_info.num_tokens_per_req = self.speculative_num_steps + 1
|
||||||
batch.return_hidden_states = False
|
batch.return_hidden_states = False
|
||||||
batch.forward_mode = (
|
batch.forward_mode = (
|
||||||
ForwardMode.TARGET_VERIFY
|
ForwardMode.TARGET_VERIFY
|
||||||
@@ -867,8 +867,8 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
batch.spec_info = EagleDraftInput(
|
batch.spec_info = EagleDraftInput(
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
verified_id=next_token_ids,
|
verified_id=next_token_ids,
|
||||||
num_tokens_per_batch=1,
|
num_tokens_per_req=1,
|
||||||
num_tokens_for_logprob_per_batch=1,
|
num_tokens_for_logprob_per_req=1,
|
||||||
)
|
)
|
||||||
batch.return_hidden_states = False
|
batch.return_hidden_states = False
|
||||||
batch.spec_info.prepare_for_extend(batch)
|
batch.spec_info.prepare_for_extend(batch)
|
||||||
@@ -915,8 +915,8 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
capture_hidden_mode=CaptureHiddenMode.LAST,
|
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||||
)
|
)
|
||||||
|
|
||||||
batch.spec_info.num_tokens_per_batch = self.speculative_num_steps + 1
|
batch.spec_info.num_tokens_per_req = self.speculative_num_steps + 1
|
||||||
batch.spec_info.num_tokens_for_logprob_per_batch = 1
|
batch.spec_info.num_tokens_for_logprob_per_req = 1
|
||||||
batch.spec_info.prepare_extend_after_decode(
|
batch.spec_info.prepare_extend_after_decode(
|
||||||
batch,
|
batch,
|
||||||
self.speculative_num_steps,
|
self.speculative_num_steps,
|
||||||
|
|||||||
@@ -483,9 +483,9 @@ class EagleDraftWorker(BaseDraftWorker):
|
|||||||
hidden_states=target_hidden_states,
|
hidden_states=target_hidden_states,
|
||||||
verified_id=next_token_ids,
|
verified_id=next_token_ids,
|
||||||
new_seq_lens=batch.seq_lens,
|
new_seq_lens=batch.seq_lens,
|
||||||
# draft mode is same with decode mode, only 1 num token per batch
|
# draft mode is same with decode mode, only 1 token per req
|
||||||
num_tokens_per_batch=1,
|
num_tokens_per_req=1,
|
||||||
num_tokens_for_logprob_per_batch=1,
|
num_tokens_for_logprob_per_req=1,
|
||||||
)
|
)
|
||||||
|
|
||||||
batch.spec_info = next_draft_input
|
batch.spec_info = next_draft_input
|
||||||
@@ -508,8 +508,8 @@ class EagleDraftWorker(BaseDraftWorker):
|
|||||||
# Batch 2: Draft extend
|
# Batch 2: Draft extend
|
||||||
draft_input = EagleDraftInput(
|
draft_input = EagleDraftInput(
|
||||||
hidden_states=batch_result.logits_output.hidden_states,
|
hidden_states=batch_result.logits_output.hidden_states,
|
||||||
num_tokens_per_batch=self.speculative_num_steps + 1,
|
num_tokens_per_req=self.speculative_num_steps + 1,
|
||||||
num_tokens_for_logprob_per_batch=self.speculative_num_steps + 1,
|
num_tokens_for_logprob_per_req=self.speculative_num_steps + 1,
|
||||||
)
|
)
|
||||||
select_index = (
|
select_index = (
|
||||||
torch.arange(len(batch.seq_lens), device=self.device)
|
torch.arange(len(batch.seq_lens), device=self.device)
|
||||||
@@ -691,7 +691,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
|
|
||||||
# Parse args
|
# Parse args
|
||||||
verify_input: EagleVerifyInput = batch.spec_info
|
verify_input: EagleVerifyInput = batch.spec_info
|
||||||
verify_input.num_tokens_per_batch = self.speculative_num_steps + 1
|
verify_input.num_tokens_per_req = self.speculative_num_steps + 1
|
||||||
bs = len(batch.seq_lens)
|
bs = len(batch.seq_lens)
|
||||||
|
|
||||||
# Batch 1: Target verify
|
# Batch 1: Target verify
|
||||||
|
|||||||
@@ -497,8 +497,8 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
|
|||||||
|
|
||||||
forward_batch.spec_info.hidden_states = self.hidden_states[:num_tokens]
|
forward_batch.spec_info.hidden_states = self.hidden_states[:num_tokens]
|
||||||
forward_batch.spec_info.accept_length = self.accept_length[:bs]
|
forward_batch.spec_info.accept_length = self.accept_length[:bs]
|
||||||
forward_batch.spec_info.num_tokens_per_batch = self.num_tokens_per_bs
|
forward_batch.spec_info.num_tokens_per_req = self.num_tokens_per_bs
|
||||||
forward_batch.spec_info.num_tokens_for_logprob_per_batch = 1
|
forward_batch.spec_info.num_tokens_for_logprob_per_req = 1
|
||||||
forward_batch.spec_info.positions = self.positions[:num_tokens]
|
forward_batch.spec_info.positions = self.positions[:num_tokens]
|
||||||
forward_batch.spec_info.extend_seq_lens_tensor = self.extend_seq_lens[:bs]
|
forward_batch.spec_info.extend_seq_lens_tensor = self.extend_seq_lens[:bs]
|
||||||
|
|
||||||
|
|||||||
@@ -357,8 +357,8 @@ class MultiLayerEagleWorker(TpModelWorker):
|
|||||||
assert isinstance(spec_info, EagleDraftInput)
|
assert isinstance(spec_info, EagleDraftInput)
|
||||||
|
|
||||||
spec_info.capture_hidden_mode = CaptureHiddenMode.LAST
|
spec_info.capture_hidden_mode = CaptureHiddenMode.LAST
|
||||||
spec_info.num_tokens_per_batch = self.topk
|
spec_info.num_tokens_per_req = self.topk
|
||||||
spec_info.num_tokens_for_logprob_per_batch = self.topk
|
spec_info.num_tokens_for_logprob_per_req = self.topk
|
||||||
batch.return_hidden_states = False
|
batch.return_hidden_states = False
|
||||||
|
|
||||||
# Get forward batch
|
# Get forward batch
|
||||||
@@ -599,8 +599,8 @@ class MultiLayerEagleWorker(TpModelWorker):
|
|||||||
batch.spec_info = EagleDraftInput(
|
batch.spec_info = EagleDraftInput(
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
verified_id=next_token_ids,
|
verified_id=next_token_ids,
|
||||||
num_tokens_per_batch=1,
|
num_tokens_per_req=1,
|
||||||
num_tokens_for_logprob_per_batch=1,
|
num_tokens_for_logprob_per_req=1,
|
||||||
)
|
)
|
||||||
batch.return_hidden_states = False
|
batch.return_hidden_states = False
|
||||||
batch.spec_info.prepare_for_extend(batch)
|
batch.spec_info.prepare_for_extend(batch)
|
||||||
@@ -681,8 +681,8 @@ class MultiLayerEagleWorker(TpModelWorker):
|
|||||||
capture_hidden_mode=CaptureHiddenMode.LAST,
|
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||||
)
|
)
|
||||||
|
|
||||||
batch.spec_info.num_tokens_per_batch = self.speculative_num_steps + 1
|
batch.spec_info.num_tokens_per_req = self.speculative_num_steps + 1
|
||||||
batch.spec_info.num_tokens_for_logprob_per_batch = 1
|
batch.spec_info.num_tokens_for_logprob_per_req = 1
|
||||||
batch.spec_info.prepare_extend_after_decode(
|
batch.spec_info.prepare_extend_after_decode(
|
||||||
batch,
|
batch,
|
||||||
self.speculative_num_steps,
|
self.speculative_num_steps,
|
||||||
|
|||||||
@@ -352,9 +352,9 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker):
|
|||||||
hidden_states=target_hidden_states,
|
hidden_states=target_hidden_states,
|
||||||
verified_id=next_token_ids,
|
verified_id=next_token_ids,
|
||||||
new_seq_lens=batch.seq_lens,
|
new_seq_lens=batch.seq_lens,
|
||||||
# draft mode is same with decode mode, only 1 num token per batch
|
# draft mode is same with decode mode, only 1 token per req
|
||||||
num_tokens_per_batch=1,
|
num_tokens_per_req=1,
|
||||||
num_tokens_for_logprob_per_batch=1,
|
num_tokens_for_logprob_per_req=1,
|
||||||
)
|
)
|
||||||
|
|
||||||
batch.spec_info = next_draft_input
|
batch.spec_info = next_draft_input
|
||||||
@@ -411,8 +411,8 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker):
|
|||||||
# Batch 2: Draft extend
|
# Batch 2: Draft extend
|
||||||
draft_input = EagleDraftInput(
|
draft_input = EagleDraftInput(
|
||||||
hidden_states=batch_result.logits_output.hidden_states,
|
hidden_states=batch_result.logits_output.hidden_states,
|
||||||
num_tokens_per_batch=self.speculative_num_steps + 1,
|
num_tokens_per_req=self.speculative_num_steps + 1,
|
||||||
num_tokens_for_logprob_per_batch=1,
|
num_tokens_for_logprob_per_req=1,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Prepare for draft extend in a separate stream
|
# Prepare for draft extend in a separate stream
|
||||||
|
|||||||
Reference in New Issue
Block a user