diff --git a/python/sglang/srt/model_executor/cuda_graph_runner.py b/python/sglang/srt/model_executor/cuda_graph_runner.py index e5e6b02bd..55118fa17 100644 --- a/python/sglang/srt/model_executor/cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/cuda_graph_runner.py @@ -1386,6 +1386,12 @@ class CudaGraphRunner: if self.model_runner.is_draft_worker: raise RuntimeError("This should not happen.") else: + + capture_mode = ( + CaptureHiddenMode.NULL + if self.model_runner.spec_algorithm.is_standalone() + else CaptureHiddenMode.FULL + ) spec_info = EagleVerifyInput( draft_token=None, custom_mask=self.buffers.custom_mask, @@ -1397,7 +1403,7 @@ class CudaGraphRunner: spec_steps=self.speculative_num_steps, topk=self.model_runner.server_args.speculative_eagle_topk, draft_token_num=self.speculative_num_draft_tokens, - capture_hidden_mode=CaptureHiddenMode.FULL, + capture_hidden_mode=capture_mode, seq_lens_sum=None, seq_lens_cpu=None, ) diff --git a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py index 5ba7f697a..b0cd9a32d 100644 --- a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py @@ -26,7 +26,10 @@ from sglang.srt.model_executor.forward_batch_info import ( ) from sglang.srt.model_executor.input_buffers import ForwardInputBuffers from sglang.srt.speculative.eagle_info import EagleDraftInput -from sglang.srt.speculative.spec_utils import maybe_detect_nan, maybe_detect_oob +from sglang.srt.speculative.spec_utils import ( + maybe_detect_nan, + maybe_detect_oob, +) from sglang.srt.utils import ( require_attn_tp_gather, require_gathered_buffer, @@ -50,7 +53,7 @@ class EagleDraftInputBuffers(ForwardInputBuffers): extend_seq_lens: torch.Tensor topk_p: torch.Tensor topk_index: torch.Tensor - hidden_states: torch.Tensor + hidden_states: Optional[torch.Tensor] global_num_tokens_gpu: Optional[torch.Tensor] global_num_tokens_for_logprob_gpu: Optional[torch.Tensor] @@ -129,9 +132,14 @@ class EAGLEDraftCudaGraphRunner: extend_seq_lens = torch.ones((self.max_bs,), dtype=torch.int32) topk_p = torch.zeros((self.max_bs, self.topk), dtype=torch.float32) topk_index = torch.zeros((self.max_bs, self.topk), dtype=torch.int64) - hidden_states = torch.zeros( - (self.max_bs, EagleDraftInput.hidden_size_for(self.eagle_worker)), - dtype=EagleDraftInput.dtype_for(self.eagle_worker), + _hidden_size = EagleDraftInput.hidden_size_for(self.eagle_worker) + hidden_states = ( + torch.zeros( + (self.max_bs, _hidden_size), + dtype=EagleDraftInput.dtype_for(self.eagle_worker), + ) + if _hidden_size is not None + else None ) if self.require_gathered_buffer: @@ -253,7 +261,11 @@ class EAGLEDraftCudaGraphRunner: out_cache_loc = buffers.out_cache_loc[: num_tokens * self.speculative_num_steps] positions = buffers.positions[:num_tokens] mrope_positions = buffers.mrope_positions[:, :num_tokens] - hidden_states = buffers.hidden_states[:num_seqs] + hidden_states = ( + buffers.hidden_states[:num_seqs] + if buffers.hidden_states is not None + else None + ) topk_p = buffers.topk_p[:num_seqs] topk_index = buffers.topk_index[:num_seqs] @@ -298,11 +310,16 @@ class EAGLEDraftCudaGraphRunner: global_dp_buffer_len = None global_num_tokens_for_logprob = None + capture_mode = ( + CaptureHiddenMode.NULL + if self.model_runner.spec_algorithm.is_standalone() + else CaptureHiddenMode.LAST + ) spec_info = EagleDraftInput( topk_p=topk_p, topk_index=topk_index, hidden_states=hidden_states, - capture_hidden_mode=CaptureHiddenMode.LAST, + capture_hidden_mode=capture_mode, ) # Forward batch @@ -401,7 +418,8 @@ class EAGLEDraftCudaGraphRunner: buffers.positions.zero_() buffers.topk_p.zero_() buffers.topk_index.zero_() - buffers.hidden_states.zero_() + if buffers.hidden_states is not None: + buffers.hidden_states.zero_() buffers.req_pool_indices.zero_() num_tokens = bs * self.num_tokens_per_bs @@ -425,7 +443,11 @@ class EAGLEDraftCudaGraphRunner: ) buffers.topk_p[:raw_bs].copy_(forward_batch.spec_info.topk_p) buffers.topk_index[:raw_bs].copy_(forward_batch.spec_info.topk_index) - buffers.hidden_states[:raw_bs].copy_(forward_batch.spec_info.hidden_states) + if ( + buffers.hidden_states is not None + and forward_batch.spec_info.hidden_states is not None + ): + buffers.hidden_states[:raw_bs].copy_(forward_batch.spec_info.hidden_states) buffers.req_pool_indices[:raw_bs].copy_(forward_batch.req_pool_indices) # TODO(ch-wan): support num_token_non_padded diff --git a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py index 5a22d3260..0f41562fa 100644 --- a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py @@ -46,7 +46,7 @@ class EagleDraftExtendInputBuffers(ForwardInputBuffers): out_cache_loc: torch.Tensor positions: torch.Tensor mrope_positions: torch.Tensor - hidden_states: torch.Tensor + hidden_states: Optional[torch.Tensor] seq_lens: torch.Tensor seq_lens_cpu: torch.Tensor extend_seq_lens: torch.Tensor @@ -132,12 +132,14 @@ class EAGLEDraftExtendCudaGraphRunner: positions = torch.zeros((self.max_num_token,), dtype=torch.int64) mrope_positions = torch.zeros((3, self.max_num_token), dtype=torch.int64) - hidden_states = torch.zeros( - ( - self.max_num_token, - EagleDraftExtendInput.hidden_size_for(self.eagle_worker), - ), - dtype=EagleDraftExtendInput.dtype_for(self.eagle_worker), + _hidden_size = EagleDraftExtendInput.hidden_size_for(self.eagle_worker) + hidden_states = ( + torch.zeros( + (self.max_num_token, _hidden_size), + dtype=EagleDraftExtendInput.dtype_for(self.eagle_worker), + ) + if _hidden_size is not None + else None ) self.seq_len_fill_value = ( self.model_runner.attn_backend.get_cuda_graph_seq_len_fill_value() @@ -292,7 +294,11 @@ class EAGLEDraftExtendCudaGraphRunner: out_cache_loc = buffers.out_cache_loc[:num_tokens] positions = buffers.positions[:num_tokens] mrope_positions = buffers.mrope_positions[:, :num_tokens] - hidden_states = buffers.hidden_states[:num_tokens] + hidden_states = ( + buffers.hidden_states[:num_tokens] + if buffers.hidden_states is not None + else None + ) num_correct_drafts = buffers.num_correct_drafts[:bs] num_accept_tokens = buffers.num_accept_tokens[:bs] next_token_logits_buffer = buffers.next_token_logits_buffer[ @@ -462,7 +468,9 @@ class EAGLEDraftExtendCudaGraphRunner: buffers.out_cache_loc[:num_tokens].copy_(forward_batch.out_cache_loc) buffers.positions[:num_tokens].copy_(forward_batch.positions) if ( - forward_batch.spec_info.hidden_states.shape[1] + buffers.hidden_states is not None + and forward_batch.spec_info.hidden_states is not None + and forward_batch.spec_info.hidden_states.shape[1] == buffers.hidden_states.shape[1] ): buffers.hidden_states[:num_tokens].copy_( diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py index 81ea6cb2b..62723fd36 100644 --- a/python/sglang/srt/speculative/eagle_info.py +++ b/python/sglang/srt/speculative/eagle_info.py @@ -559,7 +559,11 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): batch.seq_lens_cpu.add_(num_accept_tokens_cpu) draft_extend_input = EagleDraftExtendInput( - hidden_states=batch.spec_info.hidden_states[accept_index], + hidden_states=( + batch.spec_info.hidden_states[accept_index] + if batch.spec_info.hidden_states is not None + else None + ), num_correct_drafts=num_correct_drafts, num_accept_tokens=num_correct_drafts + 1, num_accept_tokens_cpu=num_accept_tokens_list, @@ -627,9 +631,11 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): unfinished_index_device ] draft_extend_input = EagleDraftExtendInput( - hidden_states=batch.spec_info.hidden_states[ - unfinished_accept_index - ], + hidden_states=( + batch.spec_info.hidden_states[unfinished_accept_index] + if batch.spec_info.hidden_states is not None + else None + ), num_accept_tokens_cpu=draft_input_num_accept_tokens_cpu, num_correct_drafts=unfinished_num_correct_drafts, num_accept_tokens=unfinished_num_correct_drafts + 1, @@ -665,7 +671,9 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): topk_p: torch.Tensor = None topk_index: torch.Tensor = None # shape: (b, hidden_size) - one hidden per req, consumed by `draft` forward. - hidden_states: torch.Tensor = None + # None when the spec algorithm's draft doesn't read hidden_states + # (e.g., STANDALONE — vanilla LLM draft). + hidden_states: Optional[torch.Tensor] = None capture_hidden_mode: CaptureHiddenMode = CaptureHiddenMode.FULL # Per-req bonus token (the "+1" target prediction at end of each accept @@ -712,28 +720,37 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): pt += extend_len @classmethod - def hidden_size_for(cls, worker) -> int: + def hidden_size_for(cls, worker) -> Optional[int]: """Decode-phase `hidden_states` width: draft self-chain output (draft model writes its own last hidden back via `capture_for_decode` - and the draft loop).""" + and the draft loop). Returns None when the draft architecture doesn't + consume the field (e.g., STANDALONE).""" + if worker.speculative_algorithm.is_standalone(): + return None return _draft_runner_of(worker).model_config.spec_hidden_size @classmethod - def dtype_for(cls, worker) -> torch.dtype: + def dtype_for(cls, worker) -> Optional[torch.dtype]: + if worker.speculative_algorithm.is_standalone(): + return None return _draft_runner_of(worker).model_config.dtype @classmethod def create_idle_input( cls, device: torch.device, - hidden_size: int, - dtype: torch.dtype, + hidden_size: Optional[int], + dtype: Optional[torch.dtype], topk: int, capture_hidden_mode: CaptureHiddenMode, ): return cls( bonus_tokens=torch.empty((0,), device=device, dtype=torch.int32), - hidden_states=torch.empty((0, hidden_size), device=device, dtype=dtype), + hidden_states=( + torch.empty((0, hidden_size), device=device, dtype=dtype) + if hidden_size is not None + else None + ), topk_p=torch.empty((0, topk), device=device, dtype=torch.float32), topk_index=torch.empty((0, topk), device=device, dtype=torch.int64), capture_hidden_mode=capture_hidden_mode, @@ -758,13 +775,15 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): self.topk_p = self.topk_p[: len(new_indices)] self.topk_index = self.topk_index[: len(new_indices)] - self.hidden_states = self.hidden_states[: len(new_indices)] + if self.hidden_states is not None: + self.hidden_states = self.hidden_states[: len(new_indices)] self.bonus_tokens = self.bonus_tokens[: len(new_indices)] else: # in some cases(e.g draft_extend), we have not filtered the batch by `unfinished_index` self.topk_p = self.topk_p[new_indices] self.topk_index = self.topk_index[new_indices] - self.hidden_states = self.hidden_states[new_indices] + if self.hidden_states is not None: + self.hidden_states = self.hidden_states[new_indices] self.bonus_tokens = self.bonus_tokens[new_indices] def merge_batch(self, spec_info: "EagleDraftInput"): @@ -777,17 +796,21 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin): ) return - if self.hidden_states is None: + # Detect idle stub by `topk_index` length (idle inputs have + # shape[0] == 0 across all fields). Don't use `hidden_states is None`: + # for STANDALONE all non-idle inputs also have None hidden_states. + if len(self.topk_index) == 0: self.hidden_states = spec_info.hidden_states self.bonus_tokens = spec_info.bonus_tokens self.topk_p = spec_info.topk_p self.topk_index = spec_info.topk_index return - if spec_info.hidden_states is None: + if len(spec_info.topk_index) == 0: return - self.hidden_states = torch.cat( - [self.hidden_states, spec_info.hidden_states], axis=0 - ) + if self.hidden_states is not None and spec_info.hidden_states is not None: + self.hidden_states = torch.cat( + [self.hidden_states, spec_info.hidden_states], axis=0 + ) self.bonus_tokens = torch.cat( [self.bonus_tokens, spec_info.bonus_tokens], axis=0 ) @@ -805,8 +828,9 @@ class EagleDraftExtendInput(SpecInput): """ # shape: (total_accepted, hidden_size). Sliced from verify-time hidden_states - # by accept_index; consumed by the draft-extend forward. - hidden_states: torch.Tensor = None + # by accept_index; consumed by the draft-extend forward. None when the spec + # algorithm's draft doesn't read hidden_states (e.g., STANDALONE). + hidden_states: Optional[torch.Tensor] = None # Per-req accept counts. `num_accept_tokens = num_correct_drafts + 1`. # Both kept for cuda-graph buffer indexing and the @@ -845,9 +869,13 @@ class EagleDraftExtendInput(SpecInput): return self.num_tokens_per_req, self.num_tokens_for_logprob_per_req @classmethod - def hidden_size_for(cls, worker) -> int: + def hidden_size_for(cls, worker) -> Optional[int]: """Extend-phase `hidden_states` width: target's `spec_hidden_size`, - widened to `num_aux * target_hidden` for EAGLE-3 aux mode.""" + widened to `num_aux * target_hidden` for EAGLE-3 aux mode. Returns + None when the draft architecture doesn't consume the field + (e.g., STANDALONE).""" + if worker.speculative_algorithm.is_standalone(): + return None target_cfg = worker.target_worker.model_runner.model_config if not ( worker.speculative_algorithm.is_eagle3() @@ -868,19 +896,25 @@ class EagleDraftExtendInput(SpecInput): return target_hidden * num_aux @classmethod - def dtype_for(cls, worker) -> torch.dtype: + def dtype_for(cls, worker) -> Optional[torch.dtype]: + if worker.speculative_algorithm.is_standalone(): + return None return worker.target_worker.model_runner.model_config.dtype @classmethod def create_idle_input( cls, device: torch.device, - hidden_size: int, - dtype: torch.dtype, + hidden_size: Optional[int], + dtype: Optional[torch.dtype], capture_hidden_mode: CaptureHiddenMode = CaptureHiddenMode.LAST, ) -> "EagleDraftExtendInput": return cls( - hidden_states=torch.empty((0, hidden_size), device=device, dtype=dtype), + hidden_states=( + torch.empty((0, hidden_size), device=device, dtype=dtype) + if hidden_size is not None + else None + ), num_correct_drafts=torch.empty((0,), device=device, dtype=torch.int32), num_accept_tokens=torch.empty((0,), device=device, dtype=torch.int32), num_accept_tokens_cpu=[], diff --git a/python/sglang/srt/speculative/eagle_info_v2.py b/python/sglang/srt/speculative/eagle_info_v2.py index fe5aa3214..4f49fed27 100644 --- a/python/sglang/srt/speculative/eagle_info_v2.py +++ b/python/sglang/srt/speculative/eagle_info_v2.py @@ -206,7 +206,12 @@ class EagleDraftInputV2Mixin: # Get a forward batch self.num_tokens_per_req = topk self.num_tokens_for_logprob_per_req = topk - batch.capture_hidden_mode = CaptureHiddenMode.LAST + capture_mode = ( + CaptureHiddenMode.NULL + if draft_model_runner.spec_algorithm.is_standalone() + else CaptureHiddenMode.LAST + ) + batch.capture_hidden_mode = capture_mode self.positions = batch.seq_lens.repeat_interleave(topk, dim=0) forward_batch = ForwardBatch.init_new(batch, draft_model_runner) can_cuda_graph = cuda_graph_runner and cuda_graph_runner.can_run(forward_batch) @@ -231,7 +236,12 @@ class EagleDraftInputV2Mixin: batch.extend_seq_lens = [num_draft_tokens for _ in range(len(batch.seq_lens))] batch.extend_prefix_lens = seq_lens_cpu_.tolist() batch.extend_num_tokens = extend_num_tokens - batch.capture_hidden_mode = CaptureHiddenMode.FULL + capture_mode = ( + CaptureHiddenMode.NULL + if draft_model_runner.spec_algorithm.is_standalone() + else CaptureHiddenMode.FULL + ) + batch.capture_hidden_mode = capture_mode batch.forward_mode = ( ForwardMode.IDLE if batch.forward_mode.is_idle() @@ -297,7 +307,12 @@ class EagleVerifyInputV2Mixin: if batch.forward_mode.is_idle() else ForwardMode.TARGET_VERIFY ) - batch.capture_hidden_mode = CaptureHiddenMode.FULL + capture_mode = ( + CaptureHiddenMode.NULL + if target_worker.model_runner.spec_algorithm.is_standalone() + else CaptureHiddenMode.FULL + ) + batch.capture_hidden_mode = capture_mode verify_forward_batch = ForwardBatch.init_new(batch, target_worker.model_runner) # Run attention backend plan and cuda graph preparation diff --git a/python/sglang/srt/speculative/eagle_worker.py b/python/sglang/srt/speculative/eagle_worker.py index bda1f71ec..f4bd641f4 100644 --- a/python/sglang/srt/speculative/eagle_worker.py +++ b/python/sglang/srt/speculative/eagle_worker.py @@ -576,7 +576,12 @@ class EAGLEWorker(TpModelWorker): # Forward with the target model and get hidden states. # We need the full hidden states to prefill the KV cache of the draft model. model_worker_batch = batch.get_model_worker_batch() - model_worker_batch.capture_hidden_mode = CaptureHiddenMode.FULL + capture_mode = ( + CaptureHiddenMode.NULL + if self.speculative_algorithm.is_standalone() + else CaptureHiddenMode.FULL + ) + model_worker_batch.capture_hidden_mode = capture_mode batch_result = self.target_worker.forward_batch_generation(model_worker_batch) logits_output, next_token_ids = ( batch_result.logits_output, @@ -729,12 +734,17 @@ class EAGLEWorker(TpModelWorker): self.token_to_kv_pool_allocator.restore_state(token_to_kv_pool_state_backup) def _draft_preprocess_idle(self, batch: ScheduleBatch): + capture_mode = ( + CaptureHiddenMode.NULL + if self.speculative_algorithm.is_standalone() + else CaptureHiddenMode.LAST + ) batch.spec_info = EagleDraftInput.create_idle_input( device=self.device, hidden_size=EagleDraftInput.hidden_size_for(self), dtype=EagleDraftInput.dtype_for(self), topk=self.topk, - capture_hidden_mode=CaptureHiddenMode.LAST, + capture_hidden_mode=capture_mode, ) def draft(self, batch: ScheduleBatch): @@ -747,14 +757,19 @@ class EAGLEWorker(TpModelWorker): spec_info = batch.spec_info assert isinstance(spec_info, EagleDraftInput) - spec_info.capture_hidden_mode = CaptureHiddenMode.LAST + draft_capture_mode = ( + CaptureHiddenMode.NULL + if self.speculative_algorithm.is_standalone() + else CaptureHiddenMode.LAST + ) + spec_info.capture_hidden_mode = draft_capture_mode 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 model_worker_batch = batch.get_model_worker_batch() - assert model_worker_batch.capture_hidden_mode == CaptureHiddenMode.LAST + assert model_worker_batch.capture_hidden_mode == draft_capture_mode forward_batch = ForwardBatch.init_new( model_worker_batch, self.draft_model_runner ) @@ -804,6 +819,11 @@ class EAGLEWorker(TpModelWorker): self.speculative_num_draft_tokens, ) + target_capture_mode = ( + CaptureHiddenMode.NULL + if self.speculative_algorithm.is_standalone() + else CaptureHiddenMode.FULL + ) return EagleVerifyInput( draft_token=draft_tokens, custom_mask=tree_mask, @@ -815,7 +835,7 @@ class EAGLEWorker(TpModelWorker): spec_steps=self.speculative_num_steps, topk=self.topk, draft_token_num=self.speculative_num_draft_tokens, - capture_hidden_mode=CaptureHiddenMode.FULL, + capture_hidden_mode=target_capture_mode, seq_lens_sum=forward_batch.seq_lens_sum, seq_lens_cpu=forward_batch.seq_lens_cpu, ) @@ -973,7 +993,10 @@ class EAGLEWorker(TpModelWorker): logits_output.next_token_logits = logits_output.next_token_logits[ res.accept_indices ] - logits_output.hidden_states = logits_output.hidden_states[res.accept_indices] + if logits_output.hidden_states is not None: + logits_output.hidden_states = logits_output.hidden_states[ + res.accept_indices + ] if ( self.target_worker.model_runner.hybrid_gdn_config is not None @@ -1009,7 +1032,7 @@ class EAGLEWorker(TpModelWorker): num_correct_drafts = torch.tensor( res.num_correct_drafts_per_req_cpu, - device=logits_output.hidden_states.device, + device=logits_output.next_token_logits.device, dtype=torch.int64, ) cumulative_num_accept_tokens = torch.cumsum(num_correct_drafts + 1, dim=0) @@ -1099,7 +1122,12 @@ class EAGLEWorker(TpModelWorker): ) batch.return_hidden_states = False batch.spec_info.prepare_for_extend(batch) - batch.spec_info.capture_hidden_mode = CaptureHiddenMode.LAST + capture_mode = ( + CaptureHiddenMode.NULL + if self.speculative_algorithm.is_standalone() + else CaptureHiddenMode.LAST + ) + batch.spec_info.capture_hidden_mode = capture_mode model_worker_batch = batch.get_model_worker_batch( seq_lens_cpu_cache=seq_lens_cpu ) @@ -1128,6 +1156,11 @@ class EAGLEWorker(TpModelWorker): input_is_idle = batch.forward_mode.is_idle() + draft_extend_capture_mode = ( + CaptureHiddenMode.NULL + if self.speculative_algorithm.is_standalone() + else CaptureHiddenMode.LAST + ) if not input_is_idle and draft_extend_input.input_ids.shape[0] == 0: # All reqs finished this verify; swap to an idle ExtendInput. batch = batch.copy() @@ -1136,7 +1169,7 @@ class EAGLEWorker(TpModelWorker): device=self.device, hidden_size=EagleDraftExtendInput.hidden_size_for(self), dtype=EagleDraftExtendInput.dtype_for(self), - capture_hidden_mode=CaptureHiddenMode.LAST, + capture_hidden_mode=draft_extend_capture_mode, ) batch.spec_info = draft_extend_input @@ -1154,8 +1187,12 @@ class EAGLEWorker(TpModelWorker): ) batch.return_hidden_states = False + # Verify-time construction of EagleDraftExtendInput uses the dataclass + # default (LAST); the worker overrides here so get_model_worker_batch() + # propagates the correct mode (NULL for STANDALONE). + draft_extend_input.capture_hidden_mode = draft_extend_capture_mode model_worker_batch = batch.get_model_worker_batch() - assert model_worker_batch.capture_hidden_mode == CaptureHiddenMode.LAST + assert model_worker_batch.capture_hidden_mode == draft_extend_capture_mode forward_batch = ForwardBatch.init_new( model_worker_batch, self.draft_model_runner ) @@ -1200,12 +1237,17 @@ class EAGLEWorker(TpModelWorker): ) # Phase 3: assemble next-iter EagleDraftInput from extend output + next_decode_capture_mode = ( + CaptureHiddenMode.NULL + if self.speculative_algorithm.is_standalone() + else CaptureHiddenMode.LAST + ) next_draft_input = EagleDraftInput( bonus_tokens=draft_extend_input.bonus_tokens, hidden_states=hidden_states, topk_p=topk_p, topk_index=topk_index, - capture_hidden_mode=CaptureHiddenMode.FULL, + capture_hidden_mode=next_decode_capture_mode, ) # Restore batch fields. `seq_lens` etc. were modified by diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 6a7cda9b8..4478ddf6d 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -622,9 +622,10 @@ class EagleDraftWorker(BaseDraftWorker): draft_logits_output.next_token_logits = draft_logits_output.next_token_logits[ select_index ] - draft_logits_output.hidden_states = draft_logits_output.hidden_states[ - select_index - ] + if draft_logits_output.hidden_states is not None: + draft_logits_output.hidden_states = draft_logits_output.hidden_states[ + select_index + ] probs = torch.softmax(draft_logits_output.next_token_logits, dim=-1) ret_topk_p, ret_topk_index = fast_topk(probs, self.topk, dim=-1) ret_hidden_states = draft_logits_output.hidden_states @@ -740,13 +741,23 @@ class EAGLEWorkerV2(BaseSpecWorker): or model_worker_batch.is_extend_in_batch ): # Target prefill - model_worker_batch.capture_hidden_mode = CaptureHiddenMode.FULL + target_capture_mode = ( + CaptureHiddenMode.NULL + if self.speculative_algorithm.is_standalone() + else CaptureHiddenMode.FULL + ) + model_worker_batch.capture_hidden_mode = target_capture_mode batch_output = self.target_worker.forward_batch_generation( model_worker_batch ) # Draft prefill - model_worker_batch.capture_hidden_mode = CaptureHiddenMode.LAST + draft_capture_mode = ( + CaptureHiddenMode.NULL + if self.speculative_algorithm.is_standalone() + else CaptureHiddenMode.LAST + ) + model_worker_batch.capture_hidden_mode = draft_capture_mode with self.draft_worker.draft_tp_context( self.draft_worker.draft_runner.tp_group ), speculative_moe_backend_context(), speculative_moe_a2a_backend_context(): @@ -761,12 +772,17 @@ class EAGLEWorkerV2(BaseSpecWorker): return batch_output else: if model_worker_batch.spec_info is None: + capture_mode = ( + CaptureHiddenMode.NULL + if self.speculative_algorithm.is_standalone() + else CaptureHiddenMode.LAST + ) model_worker_batch.spec_info = EagleDraftInput.create_idle_input( device=self.device, hidden_size=EagleDraftInput.hidden_size_for(self.draft_worker), dtype=EagleDraftInput.dtype_for(self.draft_worker), topk=self.topk, - capture_hidden_mode=CaptureHiddenMode.LAST, + capture_hidden_mode=capture_mode, ) with self.draft_worker.draft_tp_context( self.draft_worker.draft_runner.tp_group 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 e4c1d498a..732bc3e62 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 @@ -359,6 +359,12 @@ class MultiLayerEagleDraftExtendCudaGraphRunner: ) spec_info.positions = None + capture_mode = ( + CaptureHiddenMode.NULL + if self.model_runner.spec_algorithm.is_standalone() + else CaptureHiddenMode.FULL + ) + # Forward batch forward_batch = ForwardBatch( forward_mode=self.forward_mode, @@ -381,7 +387,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner: global_dp_buffer_len=global_dp_buffer_len, spec_algorithm=self.model_runner.spec_algorithm, spec_info=spec_info, - capture_hidden_mode=CaptureHiddenMode.FULL, + capture_hidden_mode=capture_mode, attn_backend=self.eagle_worker.draft_extend_attn_backend_list[self.step], extend_seq_lens=extend_seq_lens, extend_seq_lens_cpu=extend_seq_lens_cpu, diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker.py b/python/sglang/srt/speculative/multi_layer_eagle_worker.py index 2c3c53e73..dc9d2e8bc 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker.py @@ -349,7 +349,12 @@ class MultiLayerEagleWorker(TpModelWorker): # Forward with the target model and get hidden states. # We need the full hidden states to prefill the KV cache of the draft model. model_worker_batch = batch.get_model_worker_batch() - model_worker_batch.capture_hidden_mode = CaptureHiddenMode.FULL + capture_mode = ( + CaptureHiddenMode.NULL + if self.speculative_algorithm.is_standalone() + else CaptureHiddenMode.FULL + ) + model_worker_batch.capture_hidden_mode = capture_mode model_worker_batch.return_hidden_states_before_norm = True batch_result = self.target_worker.forward_batch_generation(model_worker_batch) logits_output, next_token_ids = ( @@ -385,14 +390,19 @@ class MultiLayerEagleWorker(TpModelWorker): spec_info = batch.spec_info assert isinstance(spec_info, EagleDraftInput) - spec_info.capture_hidden_mode = CaptureHiddenMode.LAST + draft_capture_mode = ( + CaptureHiddenMode.NULL + if self.speculative_algorithm.is_standalone() + else CaptureHiddenMode.LAST + ) + spec_info.capture_hidden_mode = draft_capture_mode 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 model_worker_batch = batch.get_model_worker_batch() - assert model_worker_batch.capture_hidden_mode == CaptureHiddenMode.LAST + assert model_worker_batch.capture_hidden_mode == draft_capture_mode forward_batch = ForwardBatch.init_new( model_worker_batch, self.mtp_model_runner(0) ) @@ -470,6 +480,11 @@ class MultiLayerEagleWorker(TpModelWorker): self.speculative_num_draft_tokens, ) + target_capture_mode = ( + CaptureHiddenMode.NULL + if self.speculative_algorithm.is_standalone() + else CaptureHiddenMode.FULL + ) return EagleVerifyInput( draft_token=draft_tokens, custom_mask=tree_mask, @@ -481,7 +496,7 @@ class MultiLayerEagleWorker(TpModelWorker): spec_steps=self.speculative_num_steps, topk=self.topk, draft_token_num=self.server_args.speculative_num_draft_tokens, - capture_hidden_mode=CaptureHiddenMode.FULL, + capture_hidden_mode=target_capture_mode, seq_lens_sum=forward_batch.seq_lens_sum, seq_lens_cpu=forward_batch.seq_lens_cpu, ) @@ -631,7 +646,12 @@ class MultiLayerEagleWorker(TpModelWorker): ) batch.return_hidden_states = False batch.spec_info.prepare_for_extend(batch) - batch.spec_info.capture_hidden_mode = CaptureHiddenMode.LAST + capture_mode = ( + CaptureHiddenMode.NULL + if self.speculative_algorithm.is_standalone() + else CaptureHiddenMode.LAST + ) + batch.spec_info.capture_hidden_mode = capture_mode model_worker_batch = batch.get_model_worker_batch( seq_lens_cpu_cache=seq_lens_cpu ) @@ -681,6 +701,11 @@ class MultiLayerEagleWorker(TpModelWorker): input_is_idle = batch.forward_mode.is_idle() + draft_extend_capture_mode = ( + CaptureHiddenMode.NULL + if self.speculative_algorithm.is_standalone() + else CaptureHiddenMode.LAST + ) if not input_is_idle and draft_extend_input.input_ids.shape[0] == 0: batch = batch.copy() batch.prepare_for_idle() @@ -688,7 +713,7 @@ class MultiLayerEagleWorker(TpModelWorker): device=self.device, hidden_size=EagleDraftExtendInput.hidden_size_for(self), dtype=EagleDraftExtendInput.dtype_for(self), - capture_hidden_mode=CaptureHiddenMode.LAST, + capture_hidden_mode=draft_extend_capture_mode, ) batch.spec_info = draft_extend_input @@ -707,7 +732,7 @@ class MultiLayerEagleWorker(TpModelWorker): batch.return_hidden_states = False model_worker_batch = batch.get_model_worker_batch() - assert model_worker_batch.capture_hidden_mode == CaptureHiddenMode.LAST + assert model_worker_batch.capture_hidden_mode == draft_extend_capture_mode forward_batch = ForwardBatch.init_new( model_worker_batch, self.mtp_model_runner(0) ) @@ -759,12 +784,17 @@ class MultiLayerEagleWorker(TpModelWorker): pt += extend_len # Phase 3: assemble next-iter EagleDraftInput from extend output + next_decode_capture_mode = ( + CaptureHiddenMode.NULL + if self.speculative_algorithm.is_standalone() + else CaptureHiddenMode.LAST + ) next_draft_input = EagleDraftInput( bonus_tokens=draft_extend_input.bonus_tokens, hidden_states=logits_output.hidden_states, topk_p=torch.cat(topk_p_list, dim=1), topk_index=torch.cat(topk_index_list, dim=1), - capture_hidden_mode=CaptureHiddenMode.FULL, + capture_hidden_mode=next_decode_capture_mode, ) # Restore batch fields. `seq_lens` etc. were modified by 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 2befe4082..a47d87426 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -28,7 +28,10 @@ from sglang.srt.managers.io_struct import ( from sglang.srt.managers.schedule_batch import ModelWorkerBatch from sglang.srt.managers.scheduler import GenerationBatchResult from sglang.srt.managers.tp_worker import TpModelWorker -from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode, ForwardBatch +from sglang.srt.model_executor.forward_batch_info import ( + CaptureHiddenMode, + ForwardBatch, +) from sglang.srt.server_args import ServerArgs from sglang.srt.speculative.base_spec_worker import BaseDraftWorker, BaseSpecWorker from sglang.srt.speculative.draft_utils import DraftBackendFactory @@ -658,7 +661,12 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker): or model_worker_batch.is_extend_in_batch ): # Target prefill - model_worker_batch.capture_hidden_mode = CaptureHiddenMode.FULL + target_capture_mode = ( + CaptureHiddenMode.NULL + if self.speculative_algorithm.is_standalone() + else CaptureHiddenMode.FULL + ) + model_worker_batch.capture_hidden_mode = target_capture_mode batch_output = self.target_worker.forward_batch_generation( model_worker_batch ) @@ -678,12 +686,17 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker): return batch_output else: if model_worker_batch.spec_info is None: + capture_mode = ( + CaptureHiddenMode.NULL + if self.speculative_algorithm.is_standalone() + else CaptureHiddenMode.LAST + ) model_worker_batch.spec_info = EagleDraftInput.create_idle_input( device=self.device, hidden_size=EagleDraftInput.hidden_size_for(self.draft_worker), dtype=EagleDraftInput.dtype_for(self.draft_worker), topk=self.topk * self.speculative_num_steps, - capture_hidden_mode=CaptureHiddenMode.LAST, + capture_hidden_mode=capture_mode, ) draft_input: EagleDraftInput = model_worker_batch.spec_info verify_input: EagleVerifyInput = self.draft_worker.draft(model_worker_batch) diff --git a/python/sglang/srt/speculative/spec_utils.py b/python/sglang/srt/speculative/spec_utils.py index fa4f2eded..64cb54cd6 100644 --- a/python/sglang/srt/speculative/spec_utils.py +++ b/python/sglang/srt/speculative/spec_utils.py @@ -56,7 +56,11 @@ def spec_need_hidden_states(server_args: Optional[ServerArgs] = None) -> bool: if server_args is None: server_args = get_global_server_args() - # TODO(lsyin): also skip when 1) step = 1 or 2) standalone draft model + # STANDALONE drafts don't consume `spec_info.hidden_states` (vanilla LLM). + # multi_layer_eagle handles hidden_states internally, not via FutureMap. + # TODO(lsyin): also skip when step == 1. + if server_args.speculative_algorithm == "STANDALONE": + return False return not server_args.enable_multi_layer_eagle @@ -510,7 +514,7 @@ def _select_top_k_tokens_later( topk_index = topk_index.view(-1, topk_sq) input_ids = torch.gather(topk_index, 1, topk_cs_index).flatten() - if hidden_states.shape[0] > 0: + if hidden_states is not None and hidden_states.shape[0] > 0: flat_cs = topk_cs_index.flatten() batch_offsets = torch.arange( 0, hidden_states.shape[0], step=topk, device=flat_cs.device diff --git a/test/registered/spec/test_standalone_speculative_decoding.py b/test/registered/spec/test_standalone_speculative_decoding.py index 240d374f6..dba3a1d7a 100644 --- a/test/registered/spec/test_standalone_speculative_decoding.py +++ b/test/registered/spec/test_standalone_speculative_decoding.py @@ -6,6 +6,7 @@ import requests from sglang.srt.environ import envs from sglang.srt.utils import kill_process_tree from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.kits.radix_cache_server_kit import run_radix_attention_test from sglang.test.run_eval import run_eval from sglang.test.test_utils import ( DEFAULT_DRAFT_MODEL_STANDALONE, @@ -205,6 +206,10 @@ class TestStandaloneV2SpeculativeDecodingTriton( def get_server_args(cls): return DEFAULT_SERVER_ARGS_V2 + ["--attention-backend", "triton"] + def test_radix_attention(self): + run_radix_attention_test(self.base_url) + assert self.process.poll() is None + class TestStandaloneV2SpeculativeDecodingFlashinfer( TestStandaloneV2SpeculativeDecodingBase