Tiny rename for spec related fileds. (#18468)

This commit is contained in:
Liangsheng Yin
2026-02-09 00:10:39 -08:00
committed by GitHub
parent 107958a489
commit 875ad6cf35
9 changed files with 36 additions and 36 deletions

View File

@@ -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_(

View File

@@ -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():

View File

@@ -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):

View File

@@ -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)

View File

@@ -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,

View File

@@ -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

View File

@@ -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]

View File

@@ -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,

View File

@@ -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