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