diff --git a/python/sglang/srt/layers/attention/flashattention_backend.py b/python/sglang/srt/layers/attention/flashattention_backend.py index d03b0d4bd..c0adc510e 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -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_( diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index bcd1fda10..f96651f40 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -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(): diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py index e22eeaee4..7718ba245 100644 --- a/python/sglang/srt/speculative/eagle_info.py +++ b/python/sglang/srt/speculative/eagle_info.py @@ -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): diff --git a/python/sglang/srt/speculative/eagle_info_v2.py b/python/sglang/srt/speculative/eagle_info_v2.py index b542e9615..6b0412a74 100644 --- a/python/sglang/srt/speculative/eagle_info_v2.py +++ b/python/sglang/srt/speculative/eagle_info_v2.py @@ -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) diff --git a/python/sglang/srt/speculative/eagle_worker.py b/python/sglang/srt/speculative/eagle_worker.py index 0086e2aa7..ac689cbf5 100644 --- a/python/sglang/srt/speculative/eagle_worker.py +++ b/python/sglang/srt/speculative/eagle_worker.py @@ -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, diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 1c90a3041..a47c48bd0 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -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 diff --git a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py index 2387ba6a0..94fb58b1d 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py @@ -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] diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker.py b/python/sglang/srt/speculative/multi_layer_eagle_worker.py index 21fc97777..de9b7a855 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker.py @@ -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, diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py index a0432e2cf..dbbfa5bba 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -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