From a0b7bcf59289a6cf916fa5ee44e3cfa865a25f3d Mon Sep 17 00:00:00 2001 From: Tao Li <1435647858@qq.com> Date: Sun, 2 Aug 2026 15:09:33 +0800 Subject: [PATCH] [Feature] Support return_hidden_states="last" (#30177) Co-authored-by: litao.dream --- .../hidden_states/hidden_states_engine.py | 28 +-- .../hidden_states/hidden_states_server.py | 28 +-- python/sglang/srt/entrypoints/EngineBase.py | 6 +- python/sglang/srt/entrypoints/engine.py | 9 +- .../sglang/srt/entrypoints/openai/protocol.py | 5 +- .../srt/entrypoints/openai/serving_chat.py | 9 +- .../entrypoints/openai/serving_completions.py | 9 +- python/sglang/srt/entrypoints/openai/utils.py | 20 +- python/sglang/srt/managers/io_struct.py | 29 ++- python/sglang/srt/managers/schedule_batch.py | 70 +++++- python/sglang/srt/managers/scheduler.py | 7 +- .../batch_result_processor.py | 121 ++++++++-- .../scheduler_components/output_streamer.py | 18 +- .../sglang/srt/managers/tokenizer_manager.py | 28 ++- .../srt/model_executor/cpu_graph_runner.py | 77 ++---- .../srt/model_executor/forward_batch_info.py | 43 +++- .../cuda_graph_setup.py | 19 +- .../srt/model_executor/runner/base_runner.py | 17 +- .../runner/decode_cuda_graph_runner.py | 81 ++----- .../runner/prefill_cuda_graph_runner.py | 18 +- python/sglang/srt/server_args.py | 27 ++- test/registered/core/test_hidden_states.py | 93 +++++++- .../test_openai_server_hidden_states.py | 12 +- ...st_batch_result_processor_hidden_states.py | 225 ++++++++++++++++++ .../managers/test_hidden_state_server_mode.py | 72 ++++++ .../unit/managers/test_io_struct.py | 49 +++- .../test_schedule_batch_req_pool_indices.py | 2 + .../test_hidden_state_graph_recapture.py | 171 +++++++++++++ .../test_prefill_cuda_graph_runner.py | 32 ++- .../unit/server_args/test_server_args.py | 39 +++ 30 files changed, 1096 insertions(+), 268 deletions(-) create mode 100644 test/registered/unit/managers/test_batch_result_processor_hidden_states.py create mode 100644 test/registered/unit/managers/test_hidden_state_server_mode.py create mode 100644 test/registered/unit/model_executor/runner/test_hidden_state_graph_recapture.py diff --git a/examples/runtime/hidden_states/hidden_states_engine.py b/examples/runtime/hidden_states/hidden_states_engine.py index 60ab302ca..f888789bc 100644 --- a/examples/runtime/hidden_states/hidden_states_engine.py +++ b/examples/runtime/hidden_states/hidden_states_engine.py @@ -1,10 +1,9 @@ """ Usage: -python hidden_states.py +python hidden_states_engine.py -Note that each time you change the `return_hidden_states` parameter, -the cuda graph will be recaptured, which might lead to a performance hit. -So avoid getting hidden states and completions alternately. +CUDA graphs use the configured maximum hidden-state mode. Requests may select +that mode or a weaker one without triggering mode-dependent recapture. """ import torch @@ -22,7 +21,7 @@ def main(): # Create an LLM. llm = sgl.Engine( model_path="Alibaba-NLP/gte-Qwen2-1.5B-instruct", - enable_return_hidden_states=True, + return_hidden_states_mode="last", ) sampling_params = { @@ -32,16 +31,15 @@ def main(): } outputs = llm.generate( - prompts, sampling_params=sampling_params, return_hidden_states=True + prompts, sampling_params=sampling_params, return_hidden_states="last" ) llm.shutdown() for prompt, output in zip(prompts, outputs): - for i in range(len(output["meta_info"]["hidden_states"])): - output["meta_info"]["hidden_states"][i] = torch.tensor( - output["meta_info"]["hidden_states"][i], dtype=torch.bfloat16 - ) + hidden_state = torch.tensor( + output["meta_info"]["hidden_states"], dtype=torch.bfloat16 + ) print("===============================") print( f"Prompt: {prompt}\n" @@ -49,14 +47,8 @@ def main(): f"Prompt_Tokens: {output['meta_info']['prompt_tokens']}\t" f"Completion_tokens: {output['meta_info']['completion_tokens']}" ) - print("Hidden states: ") - hidden_states = torch.cat( - [ - i.unsqueeze(0) if len(i.shape) == 1 else i - for i in output["meta_info"]["hidden_states"] - ] - ) - print(hidden_states) + print("Last hidden state: ") + print(hidden_state) print() diff --git a/examples/runtime/hidden_states/hidden_states_server.py b/examples/runtime/hidden_states/hidden_states_server.py index c05646841..74a8c5246 100644 --- a/examples/runtime/hidden_states/hidden_states_server.py +++ b/examples/runtime/hidden_states/hidden_states_server.py @@ -3,9 +3,8 @@ Usage: python hidden_states_server.py -Note that each time you change the `return_hidden_states` parameter, -the cuda graph will be recaptured, which might lead to a performance hit. -So avoid getting hidden states and completions alternately. +CUDA graphs use the configured maximum hidden-state mode. Requests may select +that mode or a weaker one without triggering mode-dependent recapture. """ import requests @@ -23,7 +22,9 @@ else: def main(): # Launch the server server_process, port = launch_server_cmd( - "python -m sglang.launch_server --model-path Alibaba-NLP/gte-Qwen2-1.5B-instruct --enable-return-hidden-states --host 0.0.0.0" + "python -m sglang.launch_server --model-path " + "Alibaba-NLP/gte-Qwen2-1.5B-instruct " + "--return-hidden-states-mode last --host 0.0.0.0" ) wait_for_server(f"http://localhost:{port}", process=server_process) @@ -43,7 +44,7 @@ def main(): json_data = { "text": prompts, "sampling_params": sampling_params, - "return_hidden_states": True, + "return_hidden_states": "last", } response = requests.post( @@ -55,10 +56,9 @@ def main(): outputs = response.json() for prompt, output in zip(prompts, outputs): - for i in range(len(output["meta_info"]["hidden_states"])): - output["meta_info"]["hidden_states"][i] = torch.tensor( - output["meta_info"]["hidden_states"][i], dtype=torch.bfloat16 - ) + hidden_state = torch.tensor( + output["meta_info"]["hidden_states"], dtype=torch.bfloat16 + ) print("===============================") print( f"Prompt: {prompt}\n" @@ -66,14 +66,8 @@ def main(): f"Prompt_Tokens: {output['meta_info']['prompt_tokens']}\t" f"Completion_tokens: {output['meta_info']['completion_tokens']}" ) - print("Hidden states: ") - hidden_states = torch.cat( - [ - i.unsqueeze(0) if len(i.shape) == 1 else i - for i in output["meta_info"]["hidden_states"] - ] - ) - print(hidden_states) + print("Last hidden state: ") + print(hidden_state) print() diff --git a/python/sglang/srt/entrypoints/EngineBase.py b/python/sglang/srt/entrypoints/EngineBase.py index c5d1d18ed..06c4f1d24 100644 --- a/python/sglang/srt/entrypoints/EngineBase.py +++ b/python/sglang/srt/entrypoints/EngineBase.py @@ -1,5 +1,5 @@ from abc import ABC, abstractmethod -from typing import Dict, Iterator, List, Optional, Tuple, Union +from typing import Dict, Iterator, List, Literal, Optional, Tuple, Union import torch @@ -23,7 +23,9 @@ class EngineBase(ABC): token_ids_logprob: Optional[Union[List[List[int]], List[int]]] = None, lora_path: Optional[Union[List[Optional[str]], Optional[str]]] = None, custom_logit_processor: Optional[Union[List[str], str]] = None, - return_hidden_states: Optional[bool] = None, + return_hidden_states: Optional[ + Union[bool, Literal["last"], List[Union[bool, Literal["last"]]]] + ] = None, stream: Optional[bool] = None, bootstrap_host: Optional[Union[List[str], str]] = None, bootstrap_port: Optional[Union[List[int], int]] = None, diff --git a/python/sglang/srt/entrypoints/engine.py b/python/sglang/srt/entrypoints/engine.py index aa7933fa8..c8e785bca 100644 --- a/python/sglang/srt/entrypoints/engine.py +++ b/python/sglang/srt/entrypoints/engine.py @@ -76,6 +76,7 @@ from sglang.srt.managers.io_struct import ( ProfileReqType, ReleaseMemoryOccupationReqInput, ResumeMemoryOccupationReqInput, + ReturnHiddenStatesMode, RpcReqInput, RpcReqOutput, UnloadLoRAAdapterReqInput, @@ -363,7 +364,9 @@ class Engine(EngineScoreMixin, EngineBase): lora_path: Optional[List[Optional[str]]] = None, custom_logit_processor: Optional[Union[List[str], str]] = None, require_reasoning: bool = False, - return_hidden_states: bool = False, + return_hidden_states: Union[ + ReturnHiddenStatesMode, List[ReturnHiddenStatesMode] + ] = False, return_routed_experts: bool = False, routed_experts_start_len: int = 0, stream: bool = False, @@ -467,7 +470,9 @@ class Engine(EngineScoreMixin, EngineBase): lora_path: Optional[List[Optional[str]]] = None, custom_logit_processor: Optional[Union[List[str], str]] = None, require_reasoning: bool = False, - return_hidden_states: bool = False, + return_hidden_states: Union[ + ReturnHiddenStatesMode, List[ReturnHiddenStatesMode] + ] = False, return_routed_experts: bool = False, routed_experts_start_len: int = 0, stream: bool = False, diff --git a/python/sglang/srt/entrypoints/openai/protocol.py b/python/sglang/srt/entrypoints/openai/protocol.py index d035a537a..2352e5531 100644 --- a/python/sglang/srt/entrypoints/openai/protocol.py +++ b/python/sglang/srt/entrypoints/openai/protocol.py @@ -24,6 +24,7 @@ from typing import ( Any, Dict, List, + Literal, NamedTuple, Optional, Protocol, @@ -337,7 +338,7 @@ class CompletionRequest(BaseModel): temperature: float = 1.0 top_p: float = 1.0 user: Optional[str] = None - return_hidden_states: bool = False + return_hidden_states: Union[bool, Literal["last"]] = False return_routed_experts: bool = False routed_experts_start_len: int = 0 return_cached_tokens_details: bool = False @@ -770,7 +771,7 @@ class ChatCompletionRequest(BaseModel): default="auto", examples=["none"] ) # noqa parallel_tool_calls: bool = True - return_hidden_states: bool = False + return_hidden_states: Union[bool, Literal["last"]] = False return_routed_experts: bool = False routed_experts_start_len: int = 0 return_cached_tokens_details: bool = False diff --git a/python/sglang/srt/entrypoints/openai/serving_chat.py b/python/sglang/srt/entrypoints/openai/serving_chat.py index 308f196f9..b12615e60 100644 --- a/python/sglang/srt/entrypoints/openai/serving_chat.py +++ b/python/sglang/srt/entrypoints/openai/serving_chat.py @@ -55,6 +55,7 @@ from sglang.srt.entrypoints.openai.usage_processor import UsageProcessor from sglang.srt.entrypoints.openai.utils import ( cached_tokens_details_from_dict, process_cached_tokens_details_from_ret, + process_hidden_states_for_response, process_hidden_states_from_ret, process_routed_experts_from_ret, should_include_usage, @@ -1600,10 +1601,8 @@ class OpenAIServingChat(OpenAIServingBase): if request.return_hidden_states and hidden_states: for index, choice_hidden_states in hidden_states.items(): if choice_hidden_states: - last_token_hidden_states = ( - choice_hidden_states[-1] - if len(choice_hidden_states) > 1 - else [] + response_hidden_states = process_hidden_states_for_response( + choice_hidden_states, request.return_hidden_states ) hidden_states_chunk = ChatCompletionStreamResponse( id=content["meta_info"]["id"], @@ -1612,7 +1611,7 @@ class OpenAIServingChat(OpenAIServingBase): ChatCompletionResponseStreamChoice( index=index, delta=DeltaMessage( - hidden_states=last_token_hidden_states + hidden_states=response_hidden_states ), finish_reason=None, # Hidden states don't need finish_reason ) diff --git a/python/sglang/srt/entrypoints/openai/serving_completions.py b/python/sglang/srt/entrypoints/openai/serving_completions.py index 8be9bf41a..3867a1c8f 100644 --- a/python/sglang/srt/entrypoints/openai/serving_completions.py +++ b/python/sglang/srt/entrypoints/openai/serving_completions.py @@ -22,6 +22,7 @@ from sglang.srt.entrypoints.openai.usage_processor import UsageProcessor from sglang.srt.entrypoints.openai.utils import ( cached_tokens_details_from_dict, process_cached_tokens_details_from_ret, + process_hidden_states_for_response, process_hidden_states_from_ret, process_routed_experts_from_ret, should_include_usage, @@ -392,10 +393,8 @@ class OpenAIServingCompletion(OpenAIServingBase): if request.return_hidden_states and hidden_states: for index, choice_hidden_states in hidden_states.items(): if choice_hidden_states: - last_token_hidden_states = ( - choice_hidden_states[-1] - if len(choice_hidden_states) > 1 - else [] + response_hidden_states = process_hidden_states_for_response( + choice_hidden_states, request.return_hidden_states ) hidden_states_chunk = CompletionStreamResponse( id=content["meta_info"]["id"], @@ -405,7 +404,7 @@ class OpenAIServingCompletion(OpenAIServingBase): CompletionResponseStreamChoice( index=index, text="", - hidden_states=last_token_hidden_states, + hidden_states=response_hidden_states, finish_reason=None, ) ], diff --git a/python/sglang/srt/entrypoints/openai/utils.py b/python/sglang/srt/entrypoints/openai/utils.py index 7586f62f6..86e6e663b 100644 --- a/python/sglang/srt/entrypoints/openai/utils.py +++ b/python/sglang/srt/entrypoints/openai/utils.py @@ -1,5 +1,5 @@ import logging -from typing import Any, Dict, List, Optional, Union +from typing import Any, Dict, List, Literal, Optional, Union import torch @@ -71,9 +71,21 @@ def process_hidden_states_from_ret( return None hidden_states = ret_item["meta_info"].get("hidden_states", None) - if hidden_states is not None: - hidden_states = hidden_states[-1] if len(hidden_states) > 1 else [] - return hidden_states + return process_hidden_states_for_response( + hidden_states, request.return_hidden_states + ) + + +def process_hidden_states_for_response( + hidden_states: Optional[List], + return_hidden_states: Union[bool, Literal["last"]], +) -> Optional[List]: + """Format scheduler hidden states for OpenAI API responses.""" + if not return_hidden_states or hidden_states is None: + return None + if return_hidden_states == "last": + return hidden_states + return hidden_states[-1] if len(hidden_states) > 1 else [] def should_include_usage( diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index 06282f331..4e7a54b39 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -53,7 +53,11 @@ from pydantic import PlainValidator from sglang.srt.environ import envs from sglang.srt.lora.lora_registry import LoRARef from sglang.srt.managers.embed_types import PositionalEmbeds -from sglang.srt.managers.schedule_batch import Modality +from sglang.srt.managers.schedule_batch import ( + Modality, + ReturnHiddenStatesMode, + get_return_hidden_states_mode, +) from sglang.srt.multimodal.mm_utils import has_valid_data from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.utils import ImageData, VideoData @@ -219,7 +223,9 @@ class GenerateReqInput: # Whether to log metrics for this request (e.g. health_generate calls do not log metrics) log_metrics: bool = True # Whether to return hidden states - return_hidden_states: Union[List[bool], bool] = False + return_hidden_states: Union[ + List[ReturnHiddenStatesMode], ReturnHiddenStatesMode + ] = False # Whether to return captured routed experts return_routed_experts: bool = False # Absolute start position for returned routings; response covers @@ -481,6 +487,7 @@ class GenerateReqInput: self._normalize_audio_data(num) self._normalize_sampling_params(num) self._normalize_logprob_params(num) + self._normalize_return_hidden_states(num) self._normalize_custom_logit_processor(num) self._normalize_extra_key(num) self._normalize_bootstrap_params(num) @@ -644,6 +651,22 @@ class GenerateReqInput: "Cannot use list token_ids_logprob with parallel_sample_num > 1" ) + def _normalize_return_hidden_states(self, num): + """Normalize and validate per-request hidden-state return modes.""" + if isinstance(self.return_hidden_states, list): + if len(self.return_hidden_states) != self.batch_size: + raise ValueError( + "The length of return_hidden_states should be equal to the batch size." + ) + for mode in self.return_hidden_states: + get_return_hidden_states_mode(mode) + self.return_hidden_states = ( + self.return_hidden_states * self.parallel_sample_num + ) + else: + get_return_hidden_states_mode(self.return_hidden_states) + self.return_hidden_states = [self.return_hidden_states] * num + def _normalize_custom_logit_processor(self, num): """Normalize custom logit processor for batch processing.""" if self.custom_logit_processor is None: @@ -838,7 +861,7 @@ class TokenizedGenerateReqInput(BaseReq, kw_only=True): return_flat_raw_top_logprobs: bool = False # Whether to return hidden states - return_hidden_states: bool = False + return_hidden_states: ReturnHiddenStatesMode = False # Whether to return captured routed experts return_routed_experts: bool = False diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index a370c136f..14784944e 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -52,6 +52,7 @@ from typing import ( Any, Dict, List, + Literal, NamedTuple, Optional, Set, @@ -96,7 +97,11 @@ from sglang.srt.mem_cache.common import ( ) from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.mem_cache.radix_cache import RadixKey -from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode +from sglang.srt.model_executor.forward_batch_info import ( + CaptureHiddenMode, + ForwardBatch, + ForwardMode, +) from sglang.srt.observability.metrics_collector import ( DPCooperationInfo, SchedulerMetricsCollector, @@ -135,6 +140,47 @@ MM_PAD_SHIFT_VALUE = 1_000_000 logger = logging.getLogger(__name__) +ReturnHiddenStatesMode = Union[bool, Literal["last"]] + + +def get_return_hidden_states_mode( + return_hidden_states: ReturnHiddenStatesMode, +) -> CaptureHiddenMode: + if return_hidden_states is True: + return CaptureHiddenMode.FULL + if return_hidden_states == "last": + return CaptureHiddenMode.LAST + if return_hidden_states is False: + return CaptureHiddenMode.NULL + raise ValueError( + "return_hidden_states must be a boolean or the string literal 'last'." + ) + + +def get_request_return_hidden_states_mode( + return_hidden_states: Union[List[ReturnHiddenStatesMode], ReturnHiddenStatesMode], +) -> CaptureHiddenMode: + if isinstance(return_hidden_states, list): + return max( + (get_return_hidden_states_mode(mode) for mode in return_hidden_states), + default=CaptureHiddenMode.NULL, + ) + return get_return_hidden_states_mode(return_hidden_states) + + +def get_batch_return_hidden_states_mode(reqs: List[Req]) -> CaptureHiddenMode: + mode = CaptureHiddenMode.NULL + for req in reqs: + mode = max(mode, req.return_hidden_states_mode) + return mode + + +def need_return_hidden_states( + return_hidden_states: Union[List[ReturnHiddenStatesMode], ReturnHiddenStatesMode], +) -> bool: + return get_request_return_hidden_states_mode(return_hidden_states).need_capture() + + @lru_cache(maxsize=1) def sanity_check_mm_pad_shift_value(vocab_size: int) -> None: if vocab_size > MM_PAD_SHIFT_VALUE: @@ -742,7 +788,7 @@ class Req(ReqDllmMixin): session: Optional[Session] = None, custom_logit_processor: Optional[str] = None, require_reasoning: bool = False, - return_hidden_states: bool = False, + return_hidden_states: ReturnHiddenStatesMode = False, return_routed_experts: bool = False, routed_experts_start_len: int = 0, return_indexer_topk: bool = False, @@ -831,6 +877,9 @@ class Req(ReqDllmMixin): self.sampling_params = sampling_params self.custom_logit_processor = custom_logit_processor self.return_hidden_states = return_hidden_states + self.return_hidden_states_mode = get_return_hidden_states_mode( + return_hidden_states + ) # extra key for classifying the request (e.g. cache_salt) if lora_id is not None: @@ -2001,6 +2050,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): # Whether to return hidden states return_hidden_states: bool = False + return_hidden_states_mode: CaptureHiddenMode = CaptureHiddenMode.NULL # Has grammar has_grammar: bool = False @@ -2059,6 +2109,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): ): return_logprob = any(req.return_logprob for req in reqs) + return_hidden_states_mode = get_batch_return_hidden_states_mode(reqs) + batch = cls( reqs=reqs, req_to_token_pool=req_to_token_pool, @@ -2070,7 +2122,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): has_grammar=any(req.grammar for req in reqs), device=req_to_token_pool.device, spec_algorithm=spec_algorithm, - return_hidden_states=any(req.return_hidden_states for req in reqs), + return_hidden_states=return_hidden_states_mode.need_capture(), + return_hidden_states_mode=return_hidden_states_mode, is_prefill_only=all(req.is_prefill_only for req in reqs), chunked_req=chunked_req, chunked_req_next_prompt_token=_compute_chunked_req_next_prompt_token( @@ -2982,6 +3035,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): # Filter out all requests. Stale tensors are left as-is: is_empty() # keys off reqs, so callers drop the batch before a forward reads them. self.reqs = [] + self.return_hidden_states = False + self.return_hidden_states_mode = CaptureHiddenMode.NULL return if len(keep_indices) == len(self.reqs): @@ -3031,6 +3086,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): self.token_ids_logprobs = None self.has_grammar = any(req.grammar for req in self.reqs) + self.return_hidden_states_mode = get_batch_return_hidden_states_mode(self.reqs) + self.return_hidden_states = self.return_hidden_states_mode.need_capture() self.sampling_info.filter_batch(keep_indices, keep_indices_device) if self.spec_info: @@ -3092,9 +3149,10 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): self.return_logprob = self.return_logprob or other.return_logprob self.has_grammar = self.has_grammar or other.has_grammar - self.return_hidden_states = ( - self.return_hidden_states or other.return_hidden_states + self.return_hidden_states_mode = max( + self.return_hidden_states_mode, other.return_hidden_states_mode ) + self.return_hidden_states = self.return_hidden_states_mode.need_capture() self.is_prefill_only = self.is_prefill_only and other.is_prefill_only if self.spec_info: @@ -3117,6 +3175,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): out_cache_loc=self.out_cache_loc, return_logprob=self.return_logprob, has_grammar=self.has_grammar, + return_hidden_states=self.return_hidden_states, + return_hidden_states_mode=self.return_hidden_states_mode, decoding_reqs=self.decoding_reqs, spec_algorithm=self.spec_algorithm, spec_info=self.spec_info, diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 5a0b8264f..f2272447d 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -3611,16 +3611,19 @@ class Scheduler( # These 2 values are needed for processing the output, but the values can be # modified by overlap schedule. So we have to copy them here so that # we can use the correct values in output processing. - if batch.return_logprob: + if batch.return_logprob or batch.return_hidden_states: batch_result.extend_input_len_per_req = [ req.extend_range.length if req.extend_range is not None else 0 for req in batch.reqs ] + else: + batch_result.extend_input_len_per_req = None + + if batch.return_logprob: batch_result.extend_logprob_start_len_per_req = ( batch.extend_logprob_start_lens ) else: - batch_result.extend_input_len_per_req = None batch_result.extend_logprob_start_len_per_req = None ret = batch_result diff --git a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py index f75eab004..dbf83eea9 100644 --- a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py +++ b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py @@ -27,6 +27,11 @@ from sglang.srt.mem_cache.common import ( maybe_cache_unfinished_req, release_kv_cache, ) +from sglang.srt.model_executor.forward_batch_info import ( + CaptureHiddenMode, + get_required_capture_hidden_mode, + get_server_return_hidden_states_mode, +) from sglang.srt.runtime_context import ( get_disagg, get_exec, @@ -220,11 +225,33 @@ class SchedulerBatchResultProcessor: self._validate_pp_skip_output_comm(batch, result) hidden_state_offset = 0 + prefill_hidden_capture_mode = self._get_prefill_hidden_capture_mode( + batch, + self.server_args, + ) # Check finish conditions logprob_pt = 0 for i, (req, next_token_id) in enumerate(zip(batch.reqs, next_token_ids)): + if ( + batch.return_hidden_states + and logits_output.hidden_states is not None + ): + assert extend_input_len_per_req is not None + hidden_state_offset = self._append_prefill_hidden_states( + req=req, + logits_output=logits_output, + hidden_state_offset=hidden_state_offset, + capture_hidden_mode=prefill_hidden_capture_mode, + extend_input_len=extend_input_len_per_req[i], + store=( + not req.finished() + and not req.is_retracted + and req.inflight_middle_chunks <= 0 + ), + ) + if ( req.finished() and req.inflight_middle_chunks <= 0 ) or req.is_retracted: @@ -268,16 +295,6 @@ class SchedulerBatchResultProcessor: if req.return_sampling_mask: self.add_sampling_mask_return_values(i, req, logits_output) - if ( - req.return_hidden_states - and logits_output.hidden_states is not None - ): - hidden_state_offset = self._append_prefill_hidden_states( - req=req, - logits_output=logits_output, - hidden_state_offset=hidden_state_offset, - ) - if req.grammar is not None: self._apply_prefill_grammar( req=req, @@ -482,20 +499,75 @@ class SchedulerBatchResultProcessor: req: Req, logits_output: LogitsProcessorOutput, hidden_state_offset: int, + capture_hidden_mode: CaptureHiddenMode, + extend_input_len: int, + store: bool = True, ) -> int: - req.hidden_states.append( - logits_output.hidden_states[ - hidden_state_offset : ( - hidden_state_offset := hidden_state_offset - + len(req.origin_input_ids) + if capture_hidden_mode.is_full(): + start = hidden_state_offset + hidden_state_offset += extend_input_len + if not store or not req.return_hidden_states: + return hidden_state_offset + + req_hidden_states = logits_output.hidden_states[start:hidden_state_offset] + if req.return_hidden_states is True: + req.hidden_states.append(req_hidden_states.cpu().clone().tolist()) + elif req.return_hidden_states == "last": + req.hidden_states.append(req_hidden_states[-1].cpu().tolist()) + elif capture_hidden_mode.is_last(): + index = hidden_state_offset + hidden_state_offset += 1 + if store and req.return_hidden_states: + req.hidden_states.append( + logits_output.hidden_states[index].cpu().tolist() ) - ] - .cpu() - .clone() - .tolist() - ) + else: + raise ValueError( + f"Unexpected hidden states capture mode: {capture_hidden_mode}" + ) return hidden_state_offset + @staticmethod + def _append_decode_hidden_states( + *, + req: Req, + hidden_states: torch.Tensor, + start: int, + accept_len: int, + ) -> None: + if accept_len <= 0: + return + + if req.return_hidden_states == "last": + valid_accept_len = accept_len + if req.finished_len is not None: + step_start = len(req.output_ids) - accept_len + valid_accept_len = max( + 0, + min(accept_len, req.finished_len - step_start), + ) + if valid_accept_len > 0: + req.hidden_states[:] = [ + hidden_states[start + valid_accept_len - 1].cpu().tolist() + ] + else: + req.hidden_states.extend( + hidden_states[start : start + accept_len].cpu().tolist() + ) + + @staticmethod + def _get_prefill_hidden_capture_mode( + batch: ScheduleBatch, + server_args: ServerArgs, + ) -> CaptureHiddenMode: + return get_required_capture_hidden_mode( + max( + batch.return_hidden_states_mode, + get_server_return_hidden_states_mode(server_args), + ), + batch.spec_info, + ) + def _apply_prefill_grammar( self, *, req: Req, next_token_id: int, already_advanced: bool = False ) -> None: @@ -814,10 +886,11 @@ class SchedulerBatchResultProcessor: stride = result.speculative_num_draft_tokens or 1 accept_len = len(next_token_id) start = i * stride - req.hidden_states.extend( - logits_output.hidden_states[start : start + accept_len] - .cpu() - .tolist() + self._append_decode_hidden_states( + req=req, + hidden_states=logits_output.hidden_states, + start=start, + accept_len=accept_len, ) if req.grammar is not None: diff --git a/python/sglang/srt/managers/scheduler_components/output_streamer.py b/python/sglang/srt/managers/scheduler_components/output_streamer.py index 9f5b2329a..18a42407a 100644 --- a/python/sglang/srt/managers/scheduler_components/output_streamer.py +++ b/python/sglang/srt/managers/scheduler_components/output_streamer.py @@ -564,11 +564,19 @@ class _GenerationStreamAccumulator: if self.return_hidden_states: if req.return_hidden_states: - # Mirror output_ids_through_stop: spec verify steps can overshoot finished_len. - hs = req.hidden_states - if req.finished_len is not None: - hs = hs[: req.finished_len] - self.output_hidden_states.append(hs) + if req.return_hidden_states == "last": + # Collection keeps this list bounded to the final valid + # accepted token, including speculative verify overshoot. + self.output_hidden_states.append( + req.hidden_states[-1] if req.hidden_states else None + ) + else: + # Mirror output_ids_through_stop: spec verify steps can + # overshoot finished_len. + hs = req.hidden_states + if req.finished_len is not None: + hs = hs[: req.finished_len] + self.output_hidden_states.append(hs) else: self.output_hidden_states.append(None) if self.return_routed_experts: diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index bd7fc94ee..1f0091859 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -91,7 +91,10 @@ from sglang.srt.managers.io_struct import ( from sglang.srt.managers.load_snapshot import create_load_snapshot_reader from sglang.srt.managers.mm_utils import TensorTransportMode, wrap_shm_features from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors -from sglang.srt.managers.schedule_batch import MultimodalDataItem +from sglang.srt.managers.schedule_batch import ( + MultimodalDataItem, + get_request_return_hidden_states_mode, +) from sglang.srt.managers.scheduler_input_blocker import input_blocker_guard_region from sglang.srt.managers.tokenizer_control_mixin import TokenizerControlMixin from sglang.srt.managers.tokenizer_manager_score_mixin import TokenizerManagerScoreMixin @@ -99,6 +102,9 @@ from sglang.srt.managers.utils import ( compute_num_reserved_tokens, is_health_check_generate_req, ) +from sglang.srt.model_executor.forward_batch_info import ( + get_server_return_hidden_states_mode, +) from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread from sglang.srt.observability.metrics_collector import ( STAT_LOGGER_ROLE_TOKENIZER, @@ -1137,13 +1143,23 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): # Validate generation-specific fields if isinstance(obj, GenerateReqInput): self._validate_token_ids_logprob(obj) - if ( + requested_hidden_mode = get_request_return_hidden_states_mode( obj.return_hidden_states - and not self.server_args.enable_return_hidden_states - ): + ) + server_hidden_mode = get_server_return_hidden_states_mode(self.server_args) + if requested_hidden_mode > server_hidden_mode: + if server_hidden_mode.need_capture(): + raise ValueError( + "The requested return_hidden_states mode exceeds the " + f"server maximum `{self.server_args.return_hidden_states_mode}`. " + "Please launch with `--return-hidden-states-mode full` " + "to allow return_hidden_states=True." + ) raise ValueError( - "The server is not configured to return the hidden states. " - "Please set `--enable-return-hidden-states` to enable this feature." + "The server is not configured to return hidden states. " + "Please set `--return-hidden-states-mode last`, " + "`--return-hidden-states-mode full`, or the legacy " + "`--enable-return-hidden-states` flag." ) if ( obj.custom_logit_processor diff --git a/python/sglang/srt/model_executor/cpu_graph_runner.py b/python/sglang/srt/model_executor/cpu_graph_runner.py index 2c77f8a2b..65f068706 100644 --- a/python/sglang/srt/model_executor/cpu_graph_runner.py +++ b/python/sglang/srt/model_executor/cpu_graph_runner.py @@ -34,6 +34,8 @@ from sglang.srt.model_executor.forward_batch_info import ( ForwardMode, PPProxyTensors, enable_num_token_non_padded, + get_required_capture_hidden_mode, + get_server_return_hidden_states_mode, ) from sglang.srt.model_executor.forward_context import ForwardContext, forward_context from sglang.srt.model_executor.runner_utils.capture_mode import model_capture_mode @@ -562,9 +564,12 @@ class CPUGraphRunner: # Parse args self.model_runner = model_runner self.device = model_runner.device - self.enable_return_hidden_states = ( - model_runner.server_args.enable_return_hidden_states + self.return_hidden_states_mode = ( + CaptureHiddenMode.NULL + if model_runner.is_draft_worker + else get_server_return_hidden_states_mode(model_runner.server_args) ) + self.enable_return_hidden_states = self.return_hidden_states_mode.need_capture() # bs -> compiled fn (text-only / skip_cross_attention=True) self.graphs = {} # bs -> compiled fn (cross-attention / skip_cross_attention=False, enc-dec only) @@ -589,14 +594,10 @@ class CPUGraphRunner: self.pp_size = model_runner.server_args.pp_size self.capture_forward_mode = ForwardMode.DECODE - self.capture_hidden_mode = CaptureHiddenMode.NULL + self.capture_hidden_mode = self.return_hidden_states_mode # Static capture width: CPU graphs are decode-only. self.captured_req_width = 1 - # If returning hidden states is enabled, set initial capture hidden mode to full to avoid double-capture on startup - if self.enable_return_hidden_states: - self.capture_hidden_mode = CaptureHiddenMode.FULL - assert ( not self.model_runner.server_args.enable_lora ), "CPUGraphRunner does not support LoRA yet." @@ -703,21 +704,7 @@ class CPUGraphRunner: else forward_batch.batch_size <= self.max_bs ) - requested_capture_hidden_mode = max( - forward_batch.capture_hidden_mode, - ( - forward_batch.spec_info.capture_hidden_mode - if getattr(forward_batch.spec_info, "capture_hidden_mode", None) - is not None - else CaptureHiddenMode.NULL - ), - ) - capture_hidden_mode_matches = ( - requested_capture_hidden_mode == CaptureHiddenMode.NULL - or requested_capture_hidden_mode == self.capture_hidden_mode - ) - - return is_bs_supported and capture_hidden_mode_matches + return is_bs_supported def capture(self) -> None: capture_range = ( @@ -793,10 +780,10 @@ class CPUGraphRunner: encoder_lens = None spec_info = self.get_spec_info(num_tokens) - if self.capture_hidden_mode != CaptureHiddenMode.FULL: - self.capture_hidden_mode = ( - spec_info.capture_hidden_mode if spec_info else CaptureHiddenMode.NULL - ) + self.capture_hidden_mode = get_required_capture_hidden_mode( + self.capture_hidden_mode, + spec_info, + ) forward_batch = ForwardBatch( forward_mode=self.capture_forward_mode, @@ -870,43 +857,19 @@ class CPUGraphRunner: self.captured_forward_batches_cross[bs] = forward_batch return forward, out - def recapture_if_needed(self, forward_batch: ForwardBatch): - - # If the required capture_hidden_mode changes, we need to recapture the graph - - # These are the different factors that can influence the capture_hidden_mode - capture_hidden_mode_required_by_forward_batch = ( - forward_batch.capture_hidden_mode - ) - capture_hidden_mode_required_by_spec_info = getattr( - forward_batch.spec_info, "capture_hidden_mode", CaptureHiddenMode.NULL - ) - capture_hidden_mode_required_for_returning_hidden_states = ( - CaptureHiddenMode.FULL - if self.enable_return_hidden_states - else CaptureHiddenMode.NULL - ) - - # Determine the highest capture_hidden_mode required - # (If we have FULL, we can emulate LAST or NULL) - # (If we have LAST, we can emulate NULL) - required_capture_hidden_mode = max( - capture_hidden_mode_required_by_forward_batch, - capture_hidden_mode_required_by_spec_info, - capture_hidden_mode_required_for_returning_hidden_states, - ) - - # If the current hidden mode is no longer aligned with the required hidden mode, we need to set it to what is required and re-capture - if self.capture_hidden_mode != required_capture_hidden_mode: - self.capture_hidden_mode = required_capture_hidden_mode - self.capture() + def _validate_capture_hidden_mode(self, forward_batch: ForwardBatch) -> None: + if self.capture_hidden_mode < forward_batch.capture_hidden_mode: + raise RuntimeError( + "The runtime hidden-state mode exceeds the fixed CPU graph " + f"capture mode ({self.capture_hidden_mode.name})." + ) def prepare_replay( self, forward_batch: ForwardBatch, skip: bool = False, ): - self.recapture_if_needed(forward_batch) + self._validate_capture_hidden_mode(forward_batch) graphs = self.graphs_cross if not skip else self.graphs cfbs = ( diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 9447362b9..a5722c65b 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -32,7 +32,7 @@ import warnings from dataclasses import dataclass from enum import IntEnum, auto from functools import total_ordering -from typing import TYPE_CHECKING, Callable, Dict, List, Optional, Set, Tuple, Union +from typing import TYPE_CHECKING, Any, Callable, Dict, List, Optional, Set, Tuple, Union import torch @@ -231,6 +231,25 @@ def register_attn_tp_sequence_sharded_predicate( _attn_tp_sequence_sharded_predicate = predicate +def get_server_return_hidden_states_mode(server_args: Any) -> CaptureHiddenMode: + mode = getattr(server_args, "return_hidden_states_mode", None) + if mode == "last": + return CaptureHiddenMode.LAST + if mode == "full" or getattr(server_args, "enable_return_hidden_states", False): + return CaptureHiddenMode.FULL + return CaptureHiddenMode.NULL + + +def get_required_capture_hidden_mode( + capture_hidden_mode: CaptureHiddenMode, + spec_info: Optional[SpecInput], +) -> CaptureHiddenMode: + spec_capture_hidden_mode = ( + getattr(spec_info, "capture_hidden_mode", None) or CaptureHiddenMode.NULL + ) + return max(capture_hidden_mode, spec_capture_hidden_mode) + + def _attn_tp_local_shard_bounds(num_tokens_per_dp: int) -> Tuple[int, int]: """(tokens_per_rank, rank_offset) of this attn-TP rank's slice of the sequence. @@ -692,17 +711,21 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): # init_new must not mutate the input ScheduleBatch; per-forward # overrides go through explicit keyword arguments. - # capture_hidden_mode=None means no override: derive from - # SB.return_hidden_states / spec_info.capture_hidden_mode. + # capture_hidden_mode=None means no override: capture the server's + # configured maximum so lower-mode requests can share one graph. if capture_hidden_mode is None: - if batch.return_hidden_states: - capture_hidden_mode = CaptureHiddenMode.FULL - elif batch.spec_info is not None: - capture_hidden_mode = getattr( - batch.spec_info, "capture_hidden_mode", CaptureHiddenMode.NULL + request_capture_hidden_mode = ( + CaptureHiddenMode.NULL + if model_runner.is_draft_worker + else max( + batch.return_hidden_states_mode, + get_server_return_hidden_states_mode(model_runner.server_args), ) - else: - capture_hidden_mode = CaptureHiddenMode.NULL + ) + capture_hidden_mode = get_required_capture_hidden_mode( + request_capture_hidden_mode, + batch.spec_info, + ) # extend-mode-only fields are None on decode/idle if batch.forward_mode.is_decode_or_idle(): diff --git a/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py b/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py index fb026087e..6d7a0160a 100644 --- a/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py +++ b/python/sglang/srt/model_executor/model_runner_components/cuda_graph_setup.py @@ -19,6 +19,10 @@ from sglang.srt.model_executor.cuda_graph_config import ( Phase, check_cuda_graph_backend, ) +from sglang.srt.model_executor.forward_batch_info import ( + CaptureHiddenMode, + get_server_return_hidden_states_mode, +) from sglang.srt.model_executor.graph_shared_output import GraphSharedOutput from sglang.srt.model_executor.hook_manager import register_forward_hooks from sglang.srt.model_executor.model_runner_components.layer_setup import ( @@ -162,16 +166,17 @@ def capture_prefill_graph( if model_runner.is_draft_worker and not force_for_draft_worker: return None - # Skip prefill CG for EAGLE target on tc_piecewise: that backend - # captures CaptureHiddenMode.NULL while runtime requests FULL, so - # the captured graph is dead, and capturing it perturbs FP4 / - # TRTLLM-MoE state and corrupts decode replay (see #28386). BCG - # captures FULL for EAGLE target in PrefillCudaGraphRunner.__init__ - # (restored from #25795), so it does NOT need this skip. + # Skip prefill CG for EAGLE target on tc_piecewise when the fixed server + # capture ceiling is below FULL. EAGLE target prefill requests FULL, so a + # NULL or LAST graph is dead; capturing it can perturb FP4/TRTLLM-MoE + # state and corrupt decode replay (see #28386 and #28870). BCG captures + # FULL for EAGLE target in PrefillCudaGraphRunner.__init__, so it does not + # need this skip. if ( model_runner.spec_algorithm.is_eagle() and not model_runner.is_draft_worker - and not model_runner.server_args.enable_return_hidden_states + and get_server_return_hidden_states_mode(model_runner.server_args) + < CaptureHiddenMode.FULL and not check_cuda_graph_backend(Phase.PREFILL, Backend.BREAKABLE) ): logger.info( diff --git a/python/sglang/srt/model_executor/runner/base_runner.py b/python/sglang/srt/model_executor/runner/base_runner.py index 31dd6c5f3..7664b66d5 100644 --- a/python/sglang/srt/model_executor/runner/base_runner.py +++ b/python/sglang/srt/model_executor/runner/base_runner.py @@ -38,6 +38,7 @@ from sglang.srt.model_executor.forward_batch_info import ( ForwardMode, NgramEmbeddingInfo, PPProxyTensors, + get_server_return_hidden_states_mode, ) from sglang.srt.model_executor.forward_context import ForwardContext, forward_context from sglang.srt.model_executor.runner.flashinfer_autotune import ( @@ -201,9 +202,12 @@ class BaseRunner(ABC): self.dp_size = get_parallel().dp_size self.pp_size = model_runner.server_args.pp_size self.enable_pdmux = model_runner.server_args.enable_pdmux - self.enable_return_hidden_states = ( - model_runner.server_args.enable_return_hidden_states + self.return_hidden_states_mode = ( + CaptureHiddenMode.NULL + if model_runner.is_draft_worker + else get_server_return_hidden_states_mode(model_runner.server_args) ) + self.enable_return_hidden_states = self.return_hidden_states_mode.need_capture() self.attn_tp_size = get_parallel().attn_tp_size self.attn_tp_rank = get_parallel().attn_tp_rank self.tbo_plugin = TboCudaGraphRunnerPlugin() @@ -368,7 +372,11 @@ class BaseRunner(ABC): capture_forward_mode = ForwardMode.DECODE else: capture_forward_mode = ForwardMode.EXTEND - capture_hidden_mode = CaptureHiddenMode.NULL + capture_hidden_mode = ( + CaptureHiddenMode.NULL + if mr.is_draft_worker + else get_server_return_hidden_states_mode(mr.server_args) + ) num_tokens_per_req = 1 if mr.spec_algorithm.is_speculative(): if mr.is_draft_worker: @@ -378,9 +386,6 @@ class BaseRunner(ABC): capture_forward_mode = ForwardMode.TARGET_VERIFY num_tokens_per_req = mr.decode_num_tokens_per_req() - if mr.server_args.enable_return_hidden_states: - capture_hidden_mode = CaptureHiddenMode.FULL - num_tokens = batch_size * num_tokens_per_req # Caller owns the shape: passes a static buffer >= the dummy shape; no diff --git a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py index 6be24d982..bbca3c05d 100644 --- a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py @@ -62,6 +62,7 @@ from sglang.srt.model_executor.forward_batch_info import ( PPProxyTensors, compute_local_num_token_non_padded, enable_num_token_non_padded, + get_required_capture_hidden_mode, ) from sglang.srt.model_executor.forward_context import ForwardContext, forward_context from sglang.srt.model_executor.runner.base_cuda_graph_runner import ( @@ -258,7 +259,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): # --- capture mode + tokens-per-bs ------------------------------ self.capture_forward_mode = ForwardMode.DECODE - self.capture_hidden_mode = CaptureHiddenMode.NULL + self.capture_hidden_mode = self.return_hidden_states_mode # Static capture width. self.captured_req_width = model_runner.decode_num_tokens_per_req( num_draft_tokens=self.speculative_num_draft_tokens @@ -304,10 +305,6 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): "feature." ) - # If returning hidden states is enabled, set initial capture hidden mode to full to avoid double-capture on startup - if self.enable_return_hidden_states: - self.capture_hidden_mode = CaptureHiddenMode.FULL - # Attention backend self.max_bs = max(self.capture_bs) self.max_num_token = self.max_bs * self.captured_req_width @@ -557,19 +554,6 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): else True ) - requested_capture_hidden_mode = max( - forward_batch.capture_hidden_mode, - ( - forward_batch.spec_info.capture_hidden_mode - if getattr(forward_batch.spec_info, "capture_hidden_mode", None) - is not None - else CaptureHiddenMode.NULL - ), - ) - capture_hidden_mode_matches = ( - requested_capture_hidden_mode == CaptureHiddenMode.NULL - or requested_capture_hidden_mode == self.capture_hidden_mode - ) is_tbo_supported = ( forward_batch.can_run_tbo if self.enable_two_batch_overlap else True ) @@ -587,7 +571,6 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): is_bs_supported and is_encoder_lens_supported and is_tbo_supported - and capture_hidden_mode_matches and is_ngram_supported ) @@ -610,18 +593,8 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): else True ) - requested_capture_hidden_mode = max( - forward_batch.capture_hidden_mode, - ( - forward_batch.spec_info.capture_hidden_mode - if getattr(forward_batch.spec_info, "capture_hidden_mode", None) - is not None - else CaptureHiddenMode.NULL - ), - ) capture_hidden_mode_matches = ( - requested_capture_hidden_mode == CaptureHiddenMode.NULL - or requested_capture_hidden_mode == self.capture_hidden_mode + forward_batch.capture_hidden_mode <= self.capture_hidden_mode ) return ( @@ -751,10 +724,10 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): global_dp_buffer_len = None spec_info = self.get_spec_info(num_tokens) - if self.capture_hidden_mode != CaptureHiddenMode.FULL: - self.capture_hidden_mode = ( - spec_info.capture_hidden_mode if spec_info else CaptureHiddenMode.NULL - ) + self.capture_hidden_mode = get_required_capture_hidden_mode( + self.capture_hidden_mode, + spec_info, + ) if self.model_runner.server_args.enable_lora: # It is safe to capture CUDA graph using empty LoRA id, as the LoRA kernels will always be launched whenever @@ -1023,38 +996,12 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): post_warmup_hook=post_warmup_hook, ) - def recapture_if_needed(self, forward_batch: ForwardBatch): - - # If the required capture_hidden_mode changes, we need to recapture the graph - - # These are the different factors that can influence the capture_hidden_mode - capture_hidden_mode_required_by_forward_batch = ( - forward_batch.capture_hidden_mode - ) - capture_hidden_mode_required_by_spec_info = ( - getattr(forward_batch.spec_info, "capture_hidden_mode", None) - or CaptureHiddenMode.NULL - ) - capture_hidden_mode_required_for_returning_hidden_states = ( - CaptureHiddenMode.FULL - if self.enable_return_hidden_states - else CaptureHiddenMode.NULL - ) - - # Determine the highest capture_hidden_mode required - # (If we have FULL, we can emulate LAST or NULL) - # (If we have LAST, we can emulate NULL) - required_capture_hidden_mode = max( - capture_hidden_mode_required_by_forward_batch, - capture_hidden_mode_required_by_spec_info, - capture_hidden_mode_required_for_returning_hidden_states, - ) - - # If the current hidden mode is no longer aligned with the required hidden mode, we need to set it to what is required and re-capture - if self.capture_hidden_mode != required_capture_hidden_mode: - self.capture_hidden_mode = required_capture_hidden_mode - self.backend.cleanup() - self.capture() + def _validate_capture_hidden_mode(self, forward_batch: ForwardBatch) -> None: + if self.capture_hidden_mode < forward_batch.capture_hidden_mode: + raise RuntimeError( + "The runtime hidden-state mode exceeds the fixed CUDA graph " + f"capture mode ({self.capture_hidden_mode.name})." + ) def load_batch( self, @@ -1104,7 +1051,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): return buffers = self.buffers - self.recapture_if_needed(forward_batch) + self._validate_capture_hidden_mode(forward_batch) raw_bs = forward_batch.batch_size diff --git a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py index 4c396c802..ab21cc9dd 100644 --- a/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/prefill_cuda_graph_runner.py @@ -277,16 +277,12 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner): self.prefill_backend_name == Backend.BREAKABLE and model_runner.spec_algorithm.is_eagle() ) - needs_full_hidden_states = ( - model_runner.server_args.enable_return_hidden_states - or model_runner.spec_algorithm.is_dflash_family() - ) if is_breakable_eagle and model_runner.is_draft_worker: self.capture_hidden_mode = CaptureHiddenMode.LAST - elif is_breakable_eagle or needs_full_hidden_states: + elif is_breakable_eagle or model_runner.spec_algorithm.is_dflash_family(): self.capture_hidden_mode = CaptureHiddenMode.FULL else: - self.capture_hidden_mode = CaptureHiddenMode.NULL + self.capture_hidden_mode = self.return_hidden_states_mode self.mamba_track_enabled = self._is_mamba_track_enabled() @@ -1055,7 +1051,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner): return False if ( capture_hidden_mode is not None - and capture_hidden_mode != self.capture_hidden_mode + and self.capture_hidden_mode < capture_hidden_mode ): return False if return_logprob and not self._uses_eager_prefill_tail(): @@ -1651,9 +1647,17 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner): "PPProxyTensors is not supported in PrefillCudaGraphRunner yet." ) + def _validate_capture_hidden_mode(self, forward_batch: ForwardBatch) -> None: + if self.capture_hidden_mode < forward_batch.capture_hidden_mode: + raise RuntimeError( + "The runtime hidden-state mode exceeds the fixed CUDA graph " + f"capture mode ({self.capture_hidden_mode.name})." + ) + def execute( self, forward_batch: ForwardBatch, **kwargs ) -> Union[LogitsProcessorOutput, PPProxyTensors, EmbeddingPoolerOutput]: + self._validate_capture_hidden_mode(forward_batch) with self.backend.replay_session(): static_forward_batch = self.load_batch(forward_batch, **kwargs) static_num_tokens = len(static_forward_batch.input_ids) diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 36cb7c977..86eedf3fd 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -3305,8 +3305,21 @@ class ServerArgs: NS("exec.features"), ] = False enable_return_hidden_states: A[ - bool, "Enable returning hidden states with responses.", NS("exec.features") + bool, + "Enable returning full hidden states with responses. Equivalent to " + "`--return-hidden-states-mode full`.", + NS("exec.features"), ] = False + return_hidden_states_mode: A[ + Optional[str], + Arg( + help="Set the maximum hidden-state return mode supported by the " + "server. `last` allows requests with return_hidden_states=False or " + "`last`; `full` also allows return_hidden_states=True.", + choices=["last", "full"], + ), + NS("exec.features"), + ] = None enable_return_routed_experts: A[ bool, "Enable returning routed experts of each layer with responses.", @@ -3408,6 +3421,7 @@ class ServerArgs: # _handle_model_specific_adjustments never runs. self._resolved_overrides = [] + self._handle_return_hidden_states_mode() if self.model_path.lower() in ["none", "dummy"]: return @@ -3584,6 +3598,17 @@ class ServerArgs: materialize_declarations(self) + def _handle_return_hidden_states_mode(self): + if self.return_hidden_states_mode not in (None, "last", "full"): + raise ValueError( + "return_hidden_states_mode must be one of: None, 'last', or 'full'." + ) + if self.return_hidden_states_mode is None: + if self.enable_return_hidden_states: + self.return_hidden_states_mode = "full" + else: + self.enable_return_hidden_states = True + def _handle_model_capability_adjustments(self): if parse_connector_type(self.model_path) == ConnectorType.INSTANCE: return diff --git a/test/registered/core/test_hidden_states.py b/test/registered/core/test_hidden_states.py index 81dfc9f00..b6bbbb592 100644 --- a/test/registered/core/test_hidden_states.py +++ b/test/registered/core/test_hidden_states.py @@ -32,7 +32,7 @@ class TestHiddenState(CustomTestCase): model_path=cls.model_path, random_seed=42, skip_tokenizer_init=True, - enable_return_hidden_states=True, + return_hidden_states_mode="full", mem_fraction_static=0.7, ) @@ -53,8 +53,12 @@ class TestHiddenState(CustomTestCase): return_hidden_states=True, ) + expected_num_hidden_states = self.sampling_params["max_new_tokens"] for output in outputs: - self.assertEqual(len(output["meta_info"]["hidden_states"]), 8) + self.assertEqual( + len(output["meta_info"]["hidden_states"]), + expected_num_hidden_states, + ) for i in range(len(output["meta_info"]["hidden_states"])): assert isinstance(output["meta_info"]["hidden_states"][i], list) output["meta_info"]["hidden_states"][i] = torch.tensor( @@ -103,6 +107,91 @@ class TestHiddenState(CustomTestCase): ) ) + def test_return_last_hidden_state(self): + outputs = self.engine.generate( + input_ids=self.input_ids, + sampling_params=self.sampling_params, + return_hidden_states="last", + ) + + model = AutoModelForCausalLM.from_pretrained( + self.model_path, torch_dtype=torch.bfloat16, device_map=get_device() + ) + + for input_id, output in zip(self.input_ids, outputs): + sg_hidden_state = torch.tensor( + output["meta_info"]["hidden_states"], dtype=torch.bfloat16 + ).to(get_device()) + self.assertEqual(sg_hidden_state.dim(), 1) + + with torch.inference_mode(): + hf_out = model( + torch.tensor( + [input_id + output["output_ids"][:-1]], device=model.device + ), + output_hidden_states=True, + ) + hf_last_hidden_state = hf_out["hidden_states"][-1][0, -1] + + atol = 0.8 + self.assertTrue( + torch.allclose( + hf_last_hidden_state, + sg_hidden_state, + atol=atol, + rtol=0, + ) + ) + + def test_mixed_return_hidden_states_modes(self): + outputs = self.engine.generate( + input_ids=self.input_ids + [self.input_ids[0]], + sampling_params=self.sampling_params, + return_hidden_states=[False, True, "last"], + ) + + self.assertNotIn("hidden_states", outputs[0]["meta_info"]) + + full_hidden_states = outputs[1]["meta_info"]["hidden_states"] + last_hidden_state = outputs[2]["meta_info"]["hidden_states"] + + self.assertIsInstance(full_hidden_states, list) + self.assertEqual( + len(full_hidden_states), self.sampling_params["max_new_tokens"] + ) + self.assertEqual(torch.tensor(full_hidden_states[0]).dim(), 2) + + last_hidden_state = torch.tensor(last_hidden_state) + self.assertEqual(last_hidden_state.dim(), 1) + + def test_mixed_return_hidden_states_modes_with_warm_cache(self): + # Prime the radix cache so each repeated prompt only extends its + # uncached suffix during the mixed-mode prefill. + self.engine.generate( + input_ids=self.input_ids, + sampling_params={"temperature": 0, "max_new_tokens": 1}, + return_hidden_states=False, + ) + + outputs = self.engine.generate( + input_ids=self.input_ids + [self.input_ids[1]], + sampling_params={"temperature": 0, "max_new_tokens": 1}, + return_hidden_states=[True, True, "last"], + ) + + self.assertEqual( + torch.tensor(outputs[0]["meta_info"]["hidden_states"][0]).dim(), + 2, + ) + self.assertEqual( + torch.tensor(outputs[2]["meta_info"]["hidden_states"]).dim(), + 1, + ) + torch.testing.assert_close( + torch.tensor(outputs[1]["meta_info"]["hidden_states"][0])[-1], + torch.tensor(outputs[2]["meta_info"]["hidden_states"]), + ) + def test_repeatedly_changes_hidden_states(self): outputs_completion_first_round = self.engine.generate( input_ids=self.input_ids, diff --git a/test/registered/openai_server/features/test_openai_server_hidden_states.py b/test/registered/openai_server/features/test_openai_server_hidden_states.py index 5146fc6cf..2028c92fa 100644 --- a/test/registered/openai_server/features/test_openai_server_hidden_states.py +++ b/test/registered/openai_server/features/test_openai_server_hidden_states.py @@ -100,7 +100,7 @@ class BaseTestOpenAIServerWithHiddenStates(ABC): ) for choice in response.choices: - assert hasattr(choice, "hidden_states") == return_hidden_states + assert hasattr(choice, "hidden_states") == bool(return_hidden_states) if return_hidden_states: assert choice.hidden_states is not None, "hidden_states was None" @@ -139,7 +139,7 @@ class BaseTestOpenAIServerWithHiddenStates(ABC): usage = response.usage for choice in response.choices: if hasattr(choice, "hidden_states"): - assert return_hidden_states + assert bool(return_hidden_states) assert choice.hidden_states is not None hidden_states_list.append(choice.hidden_states) @@ -169,7 +169,7 @@ class BaseTestOpenAIServerWithHiddenStates(ABC): ) for choice in response.choices: - assert hasattr(choice, "hidden_states") == return_hidden_states + assert hasattr(choice, "hidden_states") == bool(return_hidden_states) if return_hidden_states: assert choice.hidden_states is not None, "hidden_states was None" @@ -196,7 +196,7 @@ class BaseTestOpenAIServerWithHiddenStates(ABC): for response in generator: for choice in response.choices: if hasattr(choice.delta, "hidden_states"): - assert return_hidden_states + assert bool(return_hidden_states) assert choice.delta.hidden_states is not None hidden_states_list.append(choice.delta.hidden_states) @@ -227,7 +227,7 @@ class TestOpenAIServerWithHiddenStatesEnabled( ) cls.base_url += "/v1" cls.tokenizer = get_tokenizer(DEFAULT_SMALL_MODEL_NAME_FOR_TEST) - cls.return_hidden_states = [False, True] + cls.return_hidden_states = [False, True, "last"] cls.use_list_input = [True, False] cls.parallel_sample_nums = [1, 2] @@ -253,7 +253,7 @@ class TestOpenAIServerWithHiddenStatesEnabledAndCUDAGraphDisabled( ) cls.base_url += "/v1" cls.tokenizer = get_tokenizer(DEFAULT_SMALL_MODEL_NAME_FOR_TEST) - cls.return_hidden_states = [False, True] + cls.return_hidden_states = [False, True, "last"] cls.use_list_input = [True, False] cls.parallel_sample_nums = [1] diff --git a/test/registered/unit/managers/test_batch_result_processor_hidden_states.py b/test/registered/unit/managers/test_batch_result_processor_hidden_states.py new file mode 100644 index 000000000..c3b2809a6 --- /dev/null +++ b/test/registered/unit/managers/test_batch_result_processor_hidden_states.py @@ -0,0 +1,225 @@ +import unittest +from types import SimpleNamespace +from unittest.mock import Mock, patch + +import torch + +from sglang.srt.managers.scheduler_components.batch_result_processor import ( + SchedulerBatchResultProcessor, +) +from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +def _make_processor(server_mode: str = "full") -> SchedulerBatchResultProcessor: + metrics_reporter = Mock() + metrics_reporter.num_generated_tokens = 0 + metrics_reporter.forward_ct_decode = 0 + return SchedulerBatchResultProcessor( + is_generation=True, + disaggregation_mode=None, + enable_overlap=False, + enable_overlap_mlx=False, + server_args=SimpleNamespace( + enable_metrics=False, + enable_hisparse=False, + enable_return_hidden_states=True, + return_hidden_states_mode=server_mode, + ), + model_config=SimpleNamespace(think_end_id=None), + token_to_kv_pool_allocator=Mock(), + tree_cache=None, + hisparse_coordinator=None, + req_to_token_pool=None, + decode_offload_manager=None, + metrics_collector=None, + metrics_reporter=metrics_reporter, + draft_worker=None, + model_worker=Mock(), + logprob_result_processor=None, + output_streamer=Mock(), + abort_request=lambda *args, **kwargs: None, + ) + + +class _PrefillReq: + def __init__(self, *, rid: str, inflight_middle_chunks: int, return_hidden_states): + self.rid = rid + self.inflight_middle_chunks = inflight_middle_chunks + self.return_hidden_states = return_hidden_states + self.hidden_states = [] + self.is_retracted = False + self.output_ids = [] + self.time_stats = Mock() + self.return_logprob = False + self.return_sampling_mask = False + self.grammar = None + self.require_reasoning = False + self.customized_info = None + + def finished(self): + return False + + def update_finish_state(self): + return None + + +class _DecodeReq: + def __init__(self): + self.return_hidden_states = "last" + self.hidden_states = [] + self.output_ids = [] + self.finished_len = None + self.is_retracted = False + self.return_logprob = False + self.return_sampling_mask = False + self.grammar = None + self.time_stats = Mock() + + def finished(self): + return self.finished_len is not None + + def update_finish_state(self, new_accept_len): + if len(self.output_ids) >= 6: + self.finished_len = 5 + + +class TestPrefillHiddenStateOffsets(CustomTestCase): + def test_active_middle_chunk_advances_before_new_last_request(self): + cases = ( + ( + "full", + CaptureHiddenMode.FULL, + torch.tensor([[10.0], [11.0], [20.0], [21.0], [22.0]]), + ), + ( + "last", + CaptureHiddenMode.LAST, + torch.tensor([[11.0], [22.0]]), + ), + ) + + for server_mode, capture_mode, hidden_states in cases: + with self.subTest(server_mode=server_mode): + middle = _PrefillReq( + rid="middle", + inflight_middle_chunks=1, + return_hidden_states=False, + ) + last = _PrefillReq( + rid="last", + inflight_middle_chunks=0, + return_hidden_states="last", + ) + batch = SimpleNamespace( + reqs=[middle, last], + decoding_reqs=[], + return_logprob=False, + return_hidden_states=True, + return_hidden_states_mode=capture_mode, + spec_info=None, + prefill_stats=None, + dp_cooperation_info=None, + ) + result = SimpleNamespace( + copy_done=None, + routed_experts_output=None, + indexer_topk_output=None, + logits_output=SimpleNamespace( + hidden_states=hidden_states, + customized_info=None, + ), + next_token_ids=torch.tensor([0, 1]), + extend_input_len_per_req=[2, 3], + extend_logprob_start_len_per_req=None, + grammar_advanced=False, + can_run_cuda_graph=False, + skipped_output_comm=False, + ) + processor = _make_processor(server_mode) + + with ( + patch( + "sglang.srt.managers.scheduler_components." + "batch_result_processor.maybe_cache_unfinished_req" + ), + patch( + "sglang.srt.managers.scheduler_components." + "batch_result_processor.get_memory", + return_value=SimpleNamespace(enable_hisparse=False), + ), + ): + processor.process_batch_result_prefill(batch, result) + + self.assertEqual(middle.hidden_states, []) + self.assertEqual(last.hidden_states, [[22.0]]) + + +class TestDecodeHiddenStateRetention(CustomTestCase): + def test_last_mode_multi_step_storage_stays_bounded(self): + processor = _make_processor() + req = _DecodeReq() + batch = SimpleNamespace( + reqs=[req], + return_logprob=False, + spec_algorithm=SimpleNamespace(is_none=lambda: False), + batch_size=lambda: 1, + ) + first_step = torch.arange(8, dtype=torch.float32).view(4, 2) + second_step = torch.arange(16, dtype=torch.float32).view(8, 2)[4:] + + def result(hidden_states): + return SimpleNamespace( + copy_done=None, + routed_experts_output=None, + indexer_topk_output=None, + logits_output=SimpleNamespace(hidden_states=hidden_states), + next_token_ids=None, + can_run_cuda_graph=False, + num_correct_drafts=0, + num_block_accept_tokens=0, + num_cap_tokens=0, + speculative_num_draft_tokens=4, + ) + + with ( + patch.object( + SchedulerBatchResultProcessor, + "_normalize_decode_outputs", + side_effect=[ + ([[1, 2, 3]], None), + ([[4, 5, 6]], None), + ], + ), + patch.object( + SchedulerBatchResultProcessor, + "_maybe_update_reasoning_tokens", + ), + patch.object( + SchedulerBatchResultProcessor, + "_handle_finish_state_updated_req", + ), + patch( + "sglang.srt.managers.scheduler_components." + "batch_result_processor.get_observability", + return_value=SimpleNamespace(enable_metrics=False), + ), + ): + processor.process_batch_result_decode(batch, result(first_step)) + + self.assertEqual(req.hidden_states, [first_step[2].tolist()]) + self.assertEqual(len(req.hidden_states), 1) + + # Only the first two accepted tokens are valid because the request + # stops inside this speculative verify step. + processor.process_batch_result_decode(batch, result(second_step)) + + self.assertEqual(req.hidden_states, [second_step[1].tolist()]) + self.assertEqual(len(req.hidden_states), 1) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/managers/test_hidden_state_server_mode.py b/test/registered/unit/managers/test_hidden_state_server_mode.py new file mode 100644 index 000000000..d5730d2b3 --- /dev/null +++ b/test/registered/unit/managers/test_hidden_state_server_mode.py @@ -0,0 +1,72 @@ +import unittest +from types import SimpleNamespace +from unittest.mock import Mock + +from sglang.srt.managers.io_struct import GenerateReqInput +from sglang.srt.managers.tokenizer_manager import TokenizerManager +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=1, suite="base-a-test-cpu") + + +class TestHiddenStateServerMode(CustomTestCase): + @staticmethod + def _make_tokenizer_manager(mode): + manager = TokenizerManager.__new__(TokenizerManager) + manager.context_len = 128 + manager.num_reserved_tokens = 0 + manager.allow_auto_truncate = False + manager.validate_total_tokens = False + manager.is_generation = True + manager.server_args = SimpleNamespace( + enable_return_hidden_states=mode is not None, + return_hidden_states_mode=mode, + enable_custom_logit_processor=False, + ) + manager._validate_token_ids_logprob = Mock() + return manager + + @staticmethod + def _make_request(return_hidden_states): + return GenerateReqInput( + input_ids=[1, 2, 3], + sampling_params={}, + return_hidden_states=return_hidden_states, + ) + + def test_last_server_accepts_false_and_last(self): + manager = self._make_tokenizer_manager("last") + + for mode in (False, "last"): + with self.subTest(mode=mode): + manager._validate_one_request( + self._make_request(mode), + [1, 2, 3], + ) + + def test_last_server_rejects_full(self): + manager = self._make_tokenizer_manager("last") + + with self.assertRaisesRegex( + ValueError, + "server maximum `last`", + ): + manager._validate_one_request( + self._make_request(True), + [1, 2, 3], + ) + + def test_full_server_accepts_all_request_modes(self): + manager = self._make_tokenizer_manager("full") + + for mode in (False, "last", True): + with self.subTest(mode=mode): + manager._validate_one_request( + self._make_request(mode), + [1, 2, 3], + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/managers/test_io_struct.py b/test/registered/unit/managers/test_io_struct.py index f6fe11be7..eb8042855 100644 --- a/test/registered/unit/managers/test_io_struct.py +++ b/test/registered/unit/managers/test_io_struct.py @@ -164,6 +164,48 @@ class TestGenerateReqInputNormalization(CustomTestCase): # Check text expansion self.assertEqual(req.text, expected_text) + def test_return_hidden_states_expands_with_parallel_sampling(self): + req = GenerateReqInput( + text=["Prompt 1", "Prompt 2"], + sampling_params={"n": 2}, + return_hidden_states=[False, "last"], + ) + + req.normalize_batch_and_arguments() + + self.assertEqual( + req.return_hidden_states, + [False, "last", False, "last"], + ) + self.assertEqual( + [req[i].return_hidden_states for i in range(4)], + [False, "last", False, "last"], + ) + + def test_return_hidden_states_batch_length_is_validated(self): + req = GenerateReqInput( + text=["Prompt 1", "Prompt 2"], + return_hidden_states=["last"], + ) + + with self.assertRaisesRegex( + ValueError, + "return_hidden_states should be equal to the batch size", + ): + req.normalize_batch_and_arguments() + + def test_return_hidden_states_batch_modes_are_validated(self): + req = GenerateReqInput( + text=["Prompt 1", "Prompt 2"], + return_hidden_states=[False, "invalid"], + ) + + with self.assertRaisesRegex( + ValueError, + "return_hidden_states must be a boolean or the string literal 'last'", + ): + req.normalize_batch_and_arguments() + def test_mixed_none_and_images_with_parallel_samples(self): """Test that when some batch items have images and others None, parallel expansion works correctly.""" req = copy.deepcopy(self.base_req) @@ -480,14 +522,14 @@ class TestGenerateReqInputNormalization(CustomTestCase): logprob_start_len=[10, 5], top_logprobs_num=[5, 3], token_ids_logprob=[[7, 8, 9], [4, 5, 6]], - return_hidden_states=[False, False, True], + return_hidden_states=[False, True], ) req.normalize_batch_and_arguments() self.assertEqual(req.return_logprob, [True, False]) self.assertEqual(req.logprob_start_len, [10, 5]) self.assertEqual(req.top_logprobs_num, [5, 3]) self.assertEqual(req.token_ids_logprob, [[7, 8, 9], [4, 5, 6]]) - self.assertEqual(req.return_hidden_states, [False, False, True]) + self.assertEqual(req.return_hidden_states, [False, True]) def test_custom_logit_processor_normalization(self): """Test normalization of custom_logit_processor.""" @@ -559,7 +601,7 @@ class TestGenerateReqInputNormalization(CustomTestCase): modalities=["image", "image"], lora_path=["path1", "path2"], custom_logit_processor=["processor1", "processor2"], - return_hidden_states=True, + return_hidden_states=[True, "last"], ) req.normalize_batch_and_arguments() @@ -580,6 +622,7 @@ class TestGenerateReqInputNormalization(CustomTestCase): self.assertEqual(item0.lora_path, "path1") self.assertEqual(item0.custom_logit_processor, "processor1") self.assertEqual(item0.return_hidden_states, True) + self.assertEqual(req[1].return_hidden_states, "last") def test_getitem_preserves_return_prompt_token_ids(self): """Batch subrequests must keep the prompt-token-id return flag.""" diff --git a/test/registered/unit/managers/test_schedule_batch_req_pool_indices.py b/test/registered/unit/managers/test_schedule_batch_req_pool_indices.py index 45316e324..1e129fbfc 100644 --- a/test/registered/unit/managers/test_schedule_batch_req_pool_indices.py +++ b/test/registered/unit/managers/test_schedule_batch_req_pool_indices.py @@ -11,6 +11,7 @@ maybe_stub_sgl_kernel() from sglang.srt.managers.hisparse_coordinator import HiSparseCoordinator # noqa: E402 from sglang.srt.managers.scheduler import Scheduler # noqa: E402 +from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode # noqa: E402 register_cpu_ci(est_time=5, suite="base-a-test-cpu") @@ -23,6 +24,7 @@ def _make_req(req_pool_idx, origin_input_ids, output_ids): return_logprob=False, grammar=None, return_hidden_states=False, + return_hidden_states_mode=CaptureHiddenMode.NULL, is_prefill_only=False, ) diff --git a/test/registered/unit/model_executor/runner/test_hidden_state_graph_recapture.py b/test/registered/unit/model_executor/runner/test_hidden_state_graph_recapture.py new file mode 100644 index 000000000..70cdcde0a --- /dev/null +++ b/test/registered/unit/model_executor/runner/test_hidden_state_graph_recapture.py @@ -0,0 +1,171 @@ +import unittest +from types import SimpleNamespace +from unittest.mock import Mock + +from sglang.srt.model_executor.cpu_graph_runner import CPUGraphRunner +from sglang.srt.model_executor.cuda_graph_config import Backend +from sglang.srt.model_executor.forward_batch_info import ( + CaptureHiddenMode, + ForwardMode, + get_server_return_hidden_states_mode, +) +from sglang.srt.model_executor.runner.decode_cuda_graph_runner import ( + DecodeCudaGraphRunner, +) +from sglang.srt.model_executor.runner.prefill_cuda_graph_runner import ( + PrefillCudaGraphRunner, +) +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=1, suite="base-a-test-cpu") + + +class TestHiddenStateGraphRecapture(CustomTestCase): + def test_server_mode_sets_graph_capture_ceiling(self): + disabled = SimpleNamespace( + enable_return_hidden_states=False, + return_hidden_states_mode=None, + ) + last = SimpleNamespace( + enable_return_hidden_states=True, + return_hidden_states_mode="last", + ) + full = SimpleNamespace( + enable_return_hidden_states=True, + return_hidden_states_mode="full", + ) + + self.assertEqual( + get_server_return_hidden_states_mode(disabled), + CaptureHiddenMode.NULL, + ) + self.assertEqual( + get_server_return_hidden_states_mode(last), + CaptureHiddenMode.LAST, + ) + self.assertEqual( + get_server_return_hidden_states_mode(full), + CaptureHiddenMode.FULL, + ) + + @staticmethod + def _make_runner(runner_cls, capture_hidden_mode): + runner = runner_cls.__new__(runner_cls) + runner.capture_hidden_mode = capture_hidden_mode + runner.backend = Mock() + runner.capture = Mock() + return runner + + @staticmethod + def _make_forward_batch(capture_hidden_mode): + return SimpleNamespace( + capture_hidden_mode=capture_hidden_mode, + spec_info=None, + ) + + @staticmethod + def _make_prefill_runner_for_can_run(capture_hidden_mode): + runner = PrefillCudaGraphRunner.__new__(PrefillCudaGraphRunner) + runner._is_full_backend = False + runner.prefill_backend_name = Backend.BREAKABLE + runner.has_mha_companion_layers = False + runner.enable_lora = False + runner._capture_chunked_prefix = False + runner.capture_hidden_mode = capture_hidden_mode + runner.capture_num_tokens = [4] + runner.max_num_tokens = 4 + return runner + + @staticmethod + def _make_prefill_forward_batch(capture_hidden_mode, spec_capture_hidden_mode): + return SimpleNamespace( + batch_size=1, + input_embeds=None, + replace_embeds=None, + forward_mode=ForwardMode.EXTEND, + capture_hidden_mode=capture_hidden_mode, + spec_info=SimpleNamespace(capture_hidden_mode=spec_capture_hidden_mode), + global_num_tokens_cpu=None, + return_logprob=False, + extend_prefix_lens_cpu=None, + input_ids=list(range(4)), + ) + + def test_stronger_graph_is_reused_for_weaker_modes(self): + runner = self._make_runner(DecodeCudaGraphRunner, CaptureHiddenMode.FULL) + + for required_mode in ( + CaptureHiddenMode.FULL, + CaptureHiddenMode.NULL, + CaptureHiddenMode.LAST, + CaptureHiddenMode.FULL, + CaptureHiddenMode.NULL, + ): + with self.subTest(required_mode=required_mode): + runner._validate_capture_hidden_mode( + self._make_forward_batch(required_mode) + ) + + self.assertEqual(runner.capture_hidden_mode, CaptureHiddenMode.FULL) + runner.backend.cleanup.assert_not_called() + runner.capture.assert_not_called() + + def test_graph_does_not_recapture_above_fixed_server_mode(self): + for runner_cls in ( + DecodeCudaGraphRunner, + PrefillCudaGraphRunner, + CPUGraphRunner, + ): + runner = self._make_runner(runner_cls, CaptureHiddenMode.NULL) + + with self.subTest(runner_cls=runner_cls), self.assertRaisesRegex( + RuntimeError, + "exceeds the fixed (CUDA|CPU) graph capture mode", + ): + runner._validate_capture_hidden_mode( + self._make_forward_batch(CaptureHiddenMode.LAST) + ) + + self.assertEqual(runner.capture_hidden_mode, CaptureHiddenMode.NULL) + runner.backend.cleanup.assert_not_called() + runner.capture.assert_not_called() + + def test_spec_worker_override_is_the_effective_runtime_mode(self): + runner = self._make_prefill_runner_for_can_run(CaptureHiddenMode.LAST) + forward_batch = self._make_prefill_forward_batch( + CaptureHiddenMode.LAST, + CaptureHiddenMode.FULL, + ) + + self.assertTrue(runner.can_run_graph(forward_batch)) + for runner_cls in ( + DecodeCudaGraphRunner, + PrefillCudaGraphRunner, + CPUGraphRunner, + ): + graph_runner = self._make_runner(runner_cls, CaptureHiddenMode.LAST) + with self.subTest(runner_cls=runner_cls): + graph_runner._validate_capture_hidden_mode(forward_batch) + + def test_prefill_graph_falls_back_for_stronger_effective_mode(self): + runner = self._make_prefill_runner_for_can_run(CaptureHiddenMode.LAST) + forward_batch = self._make_prefill_forward_batch( + CaptureHiddenMode.FULL, + CaptureHiddenMode.LAST, + ) + + self.assertFalse(runner.can_run_graph(forward_batch)) + + def test_prefill_graph_accepts_weaker_spec_mode(self): + runner = self._make_prefill_runner_for_can_run(CaptureHiddenMode.FULL) + forward_batch = self._make_prefill_forward_batch( + CaptureHiddenMode.NULL, + CaptureHiddenMode.LAST, + ) + + self.assertTrue(runner.can_run_graph(forward_batch)) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/model_executor/test_prefill_cuda_graph_runner.py b/test/registered/unit/model_executor/test_prefill_cuda_graph_runner.py index fc814fc2e..c76986050 100644 --- a/test/registered/unit/model_executor/test_prefill_cuda_graph_runner.py +++ b/test/registered/unit/model_executor/test_prefill_cuda_graph_runner.py @@ -6,8 +6,13 @@ from unittest.mock import patch import torch +import sglang.srt.model_executor.model_runner_components.cuda_graph_setup as graph_setup import sglang.srt.model_executor.runner.prefill_cuda_graph_runner as runner_module from sglang.srt.model_executor.cuda_graph_config import Backend +from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode +from sglang.srt.model_executor.model_runner_components.cuda_graph_setup import ( + capture_prefill_graph, +) from sglang.srt.model_executor.runner.prefill_cuda_graph_runner import ( PrefillCudaGraphRunner, ) @@ -56,6 +61,29 @@ class _FakeKVIndexKernel: class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase): + def test_eagle_target_tc_piecewise_skips_last_mode_capture(self): + eager_runner = object() + model_runner = SimpleNamespace( + is_draft_worker=False, + spec_algorithm=SimpleNamespace(is_eagle=lambda: True), + server_args=SimpleNamespace( + enable_return_hidden_states=True, + return_hidden_states_mode="last", + ), + ) + + with patch.object( + graph_setup, + "check_cuda_graph_backend", + return_value=False, + ): + runner = capture_prefill_graph( + model_runner=model_runner, + eager_runner=eager_runner, + ) + + self.assertIs(runner, eager_runner) + def test_prefix_chunk_capacity_is_aggregate_and_can_be_overridden(self): model_runner = SimpleNamespace( server_args=SimpleNamespace( @@ -212,7 +240,7 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase): runner = PrefillCudaGraphRunner.__new__(PrefillCudaGraphRunner) runner._capture_req_slots = 4 runner.enable_lora = False - runner.capture_hidden_mode = None + runner.capture_hidden_mode = CaptureHiddenMode.NULL runner.max_num_tokens = 32 runner.capture_num_tokens = [4] runner.backend = SimpleNamespace() @@ -227,7 +255,7 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase): input_embeds=None, replace_embeds=None, forward_mode=SimpleNamespace(is_target_verify=lambda: False), - capture_hidden_mode=None, + capture_hidden_mode=CaptureHiddenMode.NULL, global_num_tokens_cpu=None, return_logprob=False, extend_prefix_lens_cpu=[8], diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index 2ea9ee132..a7102bc25 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -42,6 +42,45 @@ _mock_device.start() class TestPrepareServerArgs(CustomTestCase): + def test_return_hidden_states_mode_configuration(self): + disabled = ServerArgs(model_path="dummy") + self.assertFalse(disabled.enable_return_hidden_states) + self.assertIsNone(disabled.return_hidden_states_mode) + + last = ServerArgs( + model_path="dummy", + return_hidden_states_mode="last", + ) + self.assertTrue(last.enable_return_hidden_states) + self.assertEqual(last.return_hidden_states_mode, "last") + + legacy_full = ServerArgs( + model_path="dummy", + enable_return_hidden_states=True, + ) + self.assertTrue(legacy_full.enable_return_hidden_states) + self.assertEqual(legacy_full.return_hidden_states_mode, "full") + + parsed_last = prepare_server_args( + [ + "--model-path", + "dummy", + "--return-hidden-states-mode", + "last", + ] + ) + self.assertTrue(parsed_last.enable_return_hidden_states) + self.assertEqual(parsed_last.return_hidden_states_mode, "last") + + with self.assertRaisesRegex( + ValueError, + "return_hidden_states_mode must be one of", + ): + ServerArgs( + model_path="dummy", + return_hidden_states_mode="lst", + ) + def test_config_nested_dict_args_are_json(self): with tempfile.NamedTemporaryFile(mode="w", suffix=".yaml", delete=False) as f: f.write("mm-process-config:\n image:\n resize: 128\n")