[Feature] Support return_hidden_states="last" (#30177)

Co-authored-by: litao.dream <litao.dream@bytedance.com>
This commit is contained in:
Tao Li
2026-08-02 15:09:33 +08:00
committed by GitHub
co-authored by litao.dream
parent 1685d29f21
commit a0b7bcf592
30 changed files with 1096 additions and 268 deletions
@@ -1,10 +1,9 @@
""" """
Usage: Usage:
python hidden_states.py python hidden_states_engine.py
Note that each time you change the `return_hidden_states` parameter, CUDA graphs use the configured maximum hidden-state mode. Requests may select
the cuda graph will be recaptured, which might lead to a performance hit. that mode or a weaker one without triggering mode-dependent recapture.
So avoid getting hidden states and completions alternately.
""" """
import torch import torch
@@ -22,7 +21,7 @@ def main():
# Create an LLM. # Create an LLM.
llm = sgl.Engine( llm = sgl.Engine(
model_path="Alibaba-NLP/gte-Qwen2-1.5B-instruct", model_path="Alibaba-NLP/gte-Qwen2-1.5B-instruct",
enable_return_hidden_states=True, return_hidden_states_mode="last",
) )
sampling_params = { sampling_params = {
@@ -32,15 +31,14 @@ def main():
} }
outputs = llm.generate( outputs = llm.generate(
prompts, sampling_params=sampling_params, return_hidden_states=True prompts, sampling_params=sampling_params, return_hidden_states="last"
) )
llm.shutdown() llm.shutdown()
for prompt, output in zip(prompts, outputs): for prompt, output in zip(prompts, outputs):
for i in range(len(output["meta_info"]["hidden_states"])): hidden_state = torch.tensor(
output["meta_info"]["hidden_states"][i] = torch.tensor( output["meta_info"]["hidden_states"], dtype=torch.bfloat16
output["meta_info"]["hidden_states"][i], dtype=torch.bfloat16
) )
print("===============================") print("===============================")
print( print(
@@ -49,14 +47,8 @@ def main():
f"Prompt_Tokens: {output['meta_info']['prompt_tokens']}\t" f"Prompt_Tokens: {output['meta_info']['prompt_tokens']}\t"
f"Completion_tokens: {output['meta_info']['completion_tokens']}" f"Completion_tokens: {output['meta_info']['completion_tokens']}"
) )
print("Hidden states: ") print("Last hidden state: ")
hidden_states = torch.cat( print(hidden_state)
[
i.unsqueeze(0) if len(i.shape) == 1 else i
for i in output["meta_info"]["hidden_states"]
]
)
print(hidden_states)
print() print()
@@ -3,9 +3,8 @@ Usage:
python hidden_states_server.py python hidden_states_server.py
Note that each time you change the `return_hidden_states` parameter, CUDA graphs use the configured maximum hidden-state mode. Requests may select
the cuda graph will be recaptured, which might lead to a performance hit. that mode or a weaker one without triggering mode-dependent recapture.
So avoid getting hidden states and completions alternately.
""" """
import requests import requests
@@ -23,7 +22,9 @@ else:
def main(): def main():
# Launch the server # Launch the server
server_process, port = launch_server_cmd( 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) wait_for_server(f"http://localhost:{port}", process=server_process)
@@ -43,7 +44,7 @@ def main():
json_data = { json_data = {
"text": prompts, "text": prompts,
"sampling_params": sampling_params, "sampling_params": sampling_params,
"return_hidden_states": True, "return_hidden_states": "last",
} }
response = requests.post( response = requests.post(
@@ -55,9 +56,8 @@ def main():
outputs = response.json() outputs = response.json()
for prompt, output in zip(prompts, outputs): for prompt, output in zip(prompts, outputs):
for i in range(len(output["meta_info"]["hidden_states"])): hidden_state = torch.tensor(
output["meta_info"]["hidden_states"][i] = torch.tensor( output["meta_info"]["hidden_states"], dtype=torch.bfloat16
output["meta_info"]["hidden_states"][i], dtype=torch.bfloat16
) )
print("===============================") print("===============================")
print( print(
@@ -66,14 +66,8 @@ def main():
f"Prompt_Tokens: {output['meta_info']['prompt_tokens']}\t" f"Prompt_Tokens: {output['meta_info']['prompt_tokens']}\t"
f"Completion_tokens: {output['meta_info']['completion_tokens']}" f"Completion_tokens: {output['meta_info']['completion_tokens']}"
) )
print("Hidden states: ") print("Last hidden state: ")
hidden_states = torch.cat( print(hidden_state)
[
i.unsqueeze(0) if len(i.shape) == 1 else i
for i in output["meta_info"]["hidden_states"]
]
)
print(hidden_states)
print() print()
+4 -2
View File
@@ -1,5 +1,5 @@
from abc import ABC, abstractmethod 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 import torch
@@ -23,7 +23,9 @@ class EngineBase(ABC):
token_ids_logprob: Optional[Union[List[List[int]], List[int]]] = None, token_ids_logprob: Optional[Union[List[List[int]], List[int]]] = None,
lora_path: Optional[Union[List[Optional[str]], Optional[str]]] = None, lora_path: Optional[Union[List[Optional[str]], Optional[str]]] = None,
custom_logit_processor: Optional[Union[List[str], 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, stream: Optional[bool] = None,
bootstrap_host: Optional[Union[List[str], str]] = None, bootstrap_host: Optional[Union[List[str], str]] = None,
bootstrap_port: Optional[Union[List[int], int]] = None, bootstrap_port: Optional[Union[List[int], int]] = None,
+7 -2
View File
@@ -76,6 +76,7 @@ from sglang.srt.managers.io_struct import (
ProfileReqType, ProfileReqType,
ReleaseMemoryOccupationReqInput, ReleaseMemoryOccupationReqInput,
ResumeMemoryOccupationReqInput, ResumeMemoryOccupationReqInput,
ReturnHiddenStatesMode,
RpcReqInput, RpcReqInput,
RpcReqOutput, RpcReqOutput,
UnloadLoRAAdapterReqInput, UnloadLoRAAdapterReqInput,
@@ -363,7 +364,9 @@ class Engine(EngineScoreMixin, EngineBase):
lora_path: Optional[List[Optional[str]]] = None, lora_path: Optional[List[Optional[str]]] = None,
custom_logit_processor: Optional[Union[List[str], str]] = None, custom_logit_processor: Optional[Union[List[str], str]] = None,
require_reasoning: bool = False, require_reasoning: bool = False,
return_hidden_states: bool = False, return_hidden_states: Union[
ReturnHiddenStatesMode, List[ReturnHiddenStatesMode]
] = False,
return_routed_experts: bool = False, return_routed_experts: bool = False,
routed_experts_start_len: int = 0, routed_experts_start_len: int = 0,
stream: bool = False, stream: bool = False,
@@ -467,7 +470,9 @@ class Engine(EngineScoreMixin, EngineBase):
lora_path: Optional[List[Optional[str]]] = None, lora_path: Optional[List[Optional[str]]] = None,
custom_logit_processor: Optional[Union[List[str], str]] = None, custom_logit_processor: Optional[Union[List[str], str]] = None,
require_reasoning: bool = False, require_reasoning: bool = False,
return_hidden_states: bool = False, return_hidden_states: Union[
ReturnHiddenStatesMode, List[ReturnHiddenStatesMode]
] = False,
return_routed_experts: bool = False, return_routed_experts: bool = False,
routed_experts_start_len: int = 0, routed_experts_start_len: int = 0,
stream: bool = False, stream: bool = False,
@@ -24,6 +24,7 @@ from typing import (
Any, Any,
Dict, Dict,
List, List,
Literal,
NamedTuple, NamedTuple,
Optional, Optional,
Protocol, Protocol,
@@ -337,7 +338,7 @@ class CompletionRequest(BaseModel):
temperature: float = 1.0 temperature: float = 1.0
top_p: float = 1.0 top_p: float = 1.0
user: Optional[str] = None user: Optional[str] = None
return_hidden_states: bool = False return_hidden_states: Union[bool, Literal["last"]] = False
return_routed_experts: bool = False return_routed_experts: bool = False
routed_experts_start_len: int = 0 routed_experts_start_len: int = 0
return_cached_tokens_details: bool = False return_cached_tokens_details: bool = False
@@ -770,7 +771,7 @@ class ChatCompletionRequest(BaseModel):
default="auto", examples=["none"] default="auto", examples=["none"]
) # noqa ) # noqa
parallel_tool_calls: bool = True parallel_tool_calls: bool = True
return_hidden_states: bool = False return_hidden_states: Union[bool, Literal["last"]] = False
return_routed_experts: bool = False return_routed_experts: bool = False
routed_experts_start_len: int = 0 routed_experts_start_len: int = 0
return_cached_tokens_details: bool = False return_cached_tokens_details: bool = False
@@ -55,6 +55,7 @@ from sglang.srt.entrypoints.openai.usage_processor import UsageProcessor
from sglang.srt.entrypoints.openai.utils import ( from sglang.srt.entrypoints.openai.utils import (
cached_tokens_details_from_dict, cached_tokens_details_from_dict,
process_cached_tokens_details_from_ret, process_cached_tokens_details_from_ret,
process_hidden_states_for_response,
process_hidden_states_from_ret, process_hidden_states_from_ret,
process_routed_experts_from_ret, process_routed_experts_from_ret,
should_include_usage, should_include_usage,
@@ -1600,10 +1601,8 @@ class OpenAIServingChat(OpenAIServingBase):
if request.return_hidden_states and hidden_states: if request.return_hidden_states and hidden_states:
for index, choice_hidden_states in hidden_states.items(): for index, choice_hidden_states in hidden_states.items():
if choice_hidden_states: if choice_hidden_states:
last_token_hidden_states = ( response_hidden_states = process_hidden_states_for_response(
choice_hidden_states[-1] choice_hidden_states, request.return_hidden_states
if len(choice_hidden_states) > 1
else []
) )
hidden_states_chunk = ChatCompletionStreamResponse( hidden_states_chunk = ChatCompletionStreamResponse(
id=content["meta_info"]["id"], id=content["meta_info"]["id"],
@@ -1612,7 +1611,7 @@ class OpenAIServingChat(OpenAIServingBase):
ChatCompletionResponseStreamChoice( ChatCompletionResponseStreamChoice(
index=index, index=index,
delta=DeltaMessage( delta=DeltaMessage(
hidden_states=last_token_hidden_states hidden_states=response_hidden_states
), ),
finish_reason=None, # Hidden states don't need finish_reason finish_reason=None, # Hidden states don't need finish_reason
) )
@@ -22,6 +22,7 @@ from sglang.srt.entrypoints.openai.usage_processor import UsageProcessor
from sglang.srt.entrypoints.openai.utils import ( from sglang.srt.entrypoints.openai.utils import (
cached_tokens_details_from_dict, cached_tokens_details_from_dict,
process_cached_tokens_details_from_ret, process_cached_tokens_details_from_ret,
process_hidden_states_for_response,
process_hidden_states_from_ret, process_hidden_states_from_ret,
process_routed_experts_from_ret, process_routed_experts_from_ret,
should_include_usage, should_include_usage,
@@ -392,10 +393,8 @@ class OpenAIServingCompletion(OpenAIServingBase):
if request.return_hidden_states and hidden_states: if request.return_hidden_states and hidden_states:
for index, choice_hidden_states in hidden_states.items(): for index, choice_hidden_states in hidden_states.items():
if choice_hidden_states: if choice_hidden_states:
last_token_hidden_states = ( response_hidden_states = process_hidden_states_for_response(
choice_hidden_states[-1] choice_hidden_states, request.return_hidden_states
if len(choice_hidden_states) > 1
else []
) )
hidden_states_chunk = CompletionStreamResponse( hidden_states_chunk = CompletionStreamResponse(
id=content["meta_info"]["id"], id=content["meta_info"]["id"],
@@ -405,7 +404,7 @@ class OpenAIServingCompletion(OpenAIServingBase):
CompletionResponseStreamChoice( CompletionResponseStreamChoice(
index=index, index=index,
text="", text="",
hidden_states=last_token_hidden_states, hidden_states=response_hidden_states,
finish_reason=None, finish_reason=None,
) )
], ],
+15 -3
View File
@@ -1,5 +1,5 @@
import logging import logging
from typing import Any, Dict, List, Optional, Union from typing import Any, Dict, List, Literal, Optional, Union
import torch import torch
@@ -71,9 +71,21 @@ def process_hidden_states_from_ret(
return None return None
hidden_states = ret_item["meta_info"].get("hidden_states", None) hidden_states = ret_item["meta_info"].get("hidden_states", None)
if hidden_states is not None: return process_hidden_states_for_response(
hidden_states = hidden_states[-1] if len(hidden_states) > 1 else [] 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
return hidden_states[-1] if len(hidden_states) > 1 else []
def should_include_usage( def should_include_usage(
+26 -3
View File
@@ -53,7 +53,11 @@ from pydantic import PlainValidator
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.lora.lora_registry import LoRARef from sglang.srt.lora.lora_registry import LoRARef
from sglang.srt.managers.embed_types import PositionalEmbeds 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.multimodal.mm_utils import has_valid_data
from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.utils import ImageData, VideoData 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) # Whether to log metrics for this request (e.g. health_generate calls do not log metrics)
log_metrics: bool = True log_metrics: bool = True
# Whether to return hidden states # 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 # Whether to return captured routed experts
return_routed_experts: bool = False return_routed_experts: bool = False
# Absolute start position for returned routings; response covers # Absolute start position for returned routings; response covers
@@ -481,6 +487,7 @@ class GenerateReqInput:
self._normalize_audio_data(num) self._normalize_audio_data(num)
self._normalize_sampling_params(num) self._normalize_sampling_params(num)
self._normalize_logprob_params(num) self._normalize_logprob_params(num)
self._normalize_return_hidden_states(num)
self._normalize_custom_logit_processor(num) self._normalize_custom_logit_processor(num)
self._normalize_extra_key(num) self._normalize_extra_key(num)
self._normalize_bootstrap_params(num) self._normalize_bootstrap_params(num)
@@ -644,6 +651,22 @@ class GenerateReqInput:
"Cannot use list token_ids_logprob with parallel_sample_num > 1" "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): def _normalize_custom_logit_processor(self, num):
"""Normalize custom logit processor for batch processing.""" """Normalize custom logit processor for batch processing."""
if self.custom_logit_processor is None: if self.custom_logit_processor is None:
@@ -838,7 +861,7 @@ class TokenizedGenerateReqInput(BaseReq, kw_only=True):
return_flat_raw_top_logprobs: bool = False return_flat_raw_top_logprobs: bool = False
# Whether to return hidden states # Whether to return hidden states
return_hidden_states: bool = False return_hidden_states: ReturnHiddenStatesMode = False
# Whether to return captured routed experts # Whether to return captured routed experts
return_routed_experts: bool = False return_routed_experts: bool = False
+65 -5
View File
@@ -52,6 +52,7 @@ from typing import (
Any, Any,
Dict, Dict,
List, List,
Literal,
NamedTuple, NamedTuple,
Optional, Optional,
Set, 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.memory_pool import ReqToTokenPool
from sglang.srt.mem_cache.radix_cache import RadixKey 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 ( from sglang.srt.observability.metrics_collector import (
DPCooperationInfo, DPCooperationInfo,
SchedulerMetricsCollector, SchedulerMetricsCollector,
@@ -135,6 +140,47 @@ MM_PAD_SHIFT_VALUE = 1_000_000
logger = logging.getLogger(__name__) 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) @lru_cache(maxsize=1)
def sanity_check_mm_pad_shift_value(vocab_size: int) -> None: def sanity_check_mm_pad_shift_value(vocab_size: int) -> None:
if vocab_size > MM_PAD_SHIFT_VALUE: if vocab_size > MM_PAD_SHIFT_VALUE:
@@ -742,7 +788,7 @@ class Req(ReqDllmMixin):
session: Optional[Session] = None, session: Optional[Session] = None,
custom_logit_processor: Optional[str] = None, custom_logit_processor: Optional[str] = None,
require_reasoning: bool = False, require_reasoning: bool = False,
return_hidden_states: bool = False, return_hidden_states: ReturnHiddenStatesMode = False,
return_routed_experts: bool = False, return_routed_experts: bool = False,
routed_experts_start_len: int = 0, routed_experts_start_len: int = 0,
return_indexer_topk: bool = False, return_indexer_topk: bool = False,
@@ -831,6 +877,9 @@ class Req(ReqDllmMixin):
self.sampling_params = sampling_params self.sampling_params = sampling_params
self.custom_logit_processor = custom_logit_processor self.custom_logit_processor = custom_logit_processor
self.return_hidden_states = return_hidden_states 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) # extra key for classifying the request (e.g. cache_salt)
if lora_id is not None: if lora_id is not None:
@@ -2001,6 +2050,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
# Whether to return hidden states # Whether to return hidden states
return_hidden_states: bool = False return_hidden_states: bool = False
return_hidden_states_mode: CaptureHiddenMode = CaptureHiddenMode.NULL
# Has grammar # Has grammar
has_grammar: bool = False has_grammar: bool = False
@@ -2059,6 +2109,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
): ):
return_logprob = any(req.return_logprob for req in reqs) return_logprob = any(req.return_logprob for req in reqs)
return_hidden_states_mode = get_batch_return_hidden_states_mode(reqs)
batch = cls( batch = cls(
reqs=reqs, reqs=reqs,
req_to_token_pool=req_to_token_pool, req_to_token_pool=req_to_token_pool,
@@ -2070,7 +2122,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
has_grammar=any(req.grammar for req in reqs), has_grammar=any(req.grammar for req in reqs),
device=req_to_token_pool.device, device=req_to_token_pool.device,
spec_algorithm=spec_algorithm, 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), is_prefill_only=all(req.is_prefill_only for req in reqs),
chunked_req=chunked_req, chunked_req=chunked_req,
chunked_req_next_prompt_token=_compute_chunked_req_next_prompt_token( 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() # 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. # keys off reqs, so callers drop the batch before a forward reads them.
self.reqs = [] self.reqs = []
self.return_hidden_states = False
self.return_hidden_states_mode = CaptureHiddenMode.NULL
return return
if len(keep_indices) == len(self.reqs): if len(keep_indices) == len(self.reqs):
@@ -3031,6 +3086,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self.token_ids_logprobs = None self.token_ids_logprobs = None
self.has_grammar = any(req.grammar for req in self.reqs) 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) self.sampling_info.filter_batch(keep_indices, keep_indices_device)
if self.spec_info: if self.spec_info:
@@ -3092,9 +3149,10 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self.return_logprob = self.return_logprob or other.return_logprob self.return_logprob = self.return_logprob or other.return_logprob
self.has_grammar = self.has_grammar or other.has_grammar self.has_grammar = self.has_grammar or other.has_grammar
self.return_hidden_states = ( self.return_hidden_states_mode = max(
self.return_hidden_states or other.return_hidden_states 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 self.is_prefill_only = self.is_prefill_only and other.is_prefill_only
if self.spec_info: if self.spec_info:
@@ -3117,6 +3175,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
out_cache_loc=self.out_cache_loc, out_cache_loc=self.out_cache_loc,
return_logprob=self.return_logprob, return_logprob=self.return_logprob,
has_grammar=self.has_grammar, 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, decoding_reqs=self.decoding_reqs,
spec_algorithm=self.spec_algorithm, spec_algorithm=self.spec_algorithm,
spec_info=self.spec_info, spec_info=self.spec_info,
+5 -2
View File
@@ -3611,16 +3611,19 @@ class Scheduler(
# These 2 values are needed for processing the output, but the values can be # 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 # modified by overlap schedule. So we have to copy them here so that
# we can use the correct values in output processing. # 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 = [ batch_result.extend_input_len_per_req = [
req.extend_range.length if req.extend_range is not None else 0 req.extend_range.length if req.extend_range is not None else 0
for req in batch.reqs 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_result.extend_logprob_start_len_per_req = (
batch.extend_logprob_start_lens batch.extend_logprob_start_lens
) )
else: else:
batch_result.extend_input_len_per_req = None
batch_result.extend_logprob_start_len_per_req = None batch_result.extend_logprob_start_len_per_req = None
ret = batch_result ret = batch_result
@@ -27,6 +27,11 @@ from sglang.srt.mem_cache.common import (
maybe_cache_unfinished_req, maybe_cache_unfinished_req,
release_kv_cache, 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 ( from sglang.srt.runtime_context import (
get_disagg, get_disagg,
get_exec, get_exec,
@@ -220,11 +225,33 @@ class SchedulerBatchResultProcessor:
self._validate_pp_skip_output_comm(batch, result) self._validate_pp_skip_output_comm(batch, result)
hidden_state_offset = 0 hidden_state_offset = 0
prefill_hidden_capture_mode = self._get_prefill_hidden_capture_mode(
batch,
self.server_args,
)
# Check finish conditions # Check finish conditions
logprob_pt = 0 logprob_pt = 0
for i, (req, next_token_id) in enumerate(zip(batch.reqs, next_token_ids)): 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 ( if (
req.finished() and req.inflight_middle_chunks <= 0 req.finished() and req.inflight_middle_chunks <= 0
) or req.is_retracted: ) or req.is_retracted:
@@ -268,16 +295,6 @@ class SchedulerBatchResultProcessor:
if req.return_sampling_mask: if req.return_sampling_mask:
self.add_sampling_mask_return_values(i, req, logits_output) 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: if req.grammar is not None:
self._apply_prefill_grammar( self._apply_prefill_grammar(
req=req, req=req,
@@ -482,20 +499,75 @@ class SchedulerBatchResultProcessor:
req: Req, req: Req,
logits_output: LogitsProcessorOutput, logits_output: LogitsProcessorOutput,
hidden_state_offset: int, hidden_state_offset: int,
capture_hidden_mode: CaptureHiddenMode,
extend_input_len: int,
store: bool = True,
) -> int: ) -> int:
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( req.hidden_states.append(
logits_output.hidden_states[ logits_output.hidden_states[index].cpu().tolist()
hidden_state_offset : (
hidden_state_offset := hidden_state_offset
+ len(req.origin_input_ids)
) )
] else:
.cpu() raise ValueError(
.clone() f"Unexpected hidden states capture mode: {capture_hidden_mode}"
.tolist()
) )
return hidden_state_offset 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( def _apply_prefill_grammar(
self, *, req: Req, next_token_id: int, already_advanced: bool = False self, *, req: Req, next_token_id: int, already_advanced: bool = False
) -> None: ) -> None:
@@ -814,10 +886,11 @@ class SchedulerBatchResultProcessor:
stride = result.speculative_num_draft_tokens or 1 stride = result.speculative_num_draft_tokens or 1
accept_len = len(next_token_id) accept_len = len(next_token_id)
start = i * stride start = i * stride
req.hidden_states.extend( self._append_decode_hidden_states(
logits_output.hidden_states[start : start + accept_len] req=req,
.cpu() hidden_states=logits_output.hidden_states,
.tolist() start=start,
accept_len=accept_len,
) )
if req.grammar is not None: if req.grammar is not None:
@@ -564,7 +564,15 @@ class _GenerationStreamAccumulator:
if self.return_hidden_states: if self.return_hidden_states:
if req.return_hidden_states: if req.return_hidden_states:
# Mirror output_ids_through_stop: spec verify steps can overshoot finished_len. 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 hs = req.hidden_states
if req.finished_len is not None: if req.finished_len is not None:
hs = hs[: req.finished_len] hs = hs[: req.finished_len]
@@ -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.load_snapshot import create_load_snapshot_reader
from sglang.srt.managers.mm_utils import TensorTransportMode, wrap_shm_features 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.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.scheduler_input_blocker import input_blocker_guard_region
from sglang.srt.managers.tokenizer_control_mixin import TokenizerControlMixin from sglang.srt.managers.tokenizer_control_mixin import TokenizerControlMixin
from sglang.srt.managers.tokenizer_manager_score_mixin import TokenizerManagerScoreMixin from sglang.srt.managers.tokenizer_manager_score_mixin import TokenizerManagerScoreMixin
@@ -99,6 +102,9 @@ from sglang.srt.managers.utils import (
compute_num_reserved_tokens, compute_num_reserved_tokens,
is_health_check_generate_req, 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.cpu_monitor import start_cpu_monitor_thread
from sglang.srt.observability.metrics_collector import ( from sglang.srt.observability.metrics_collector import (
STAT_LOGGER_ROLE_TOKENIZER, STAT_LOGGER_ROLE_TOKENIZER,
@@ -1137,13 +1143,23 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# Validate generation-specific fields # Validate generation-specific fields
if isinstance(obj, GenerateReqInput): if isinstance(obj, GenerateReqInput):
self._validate_token_ids_logprob(obj) self._validate_token_ids_logprob(obj)
if ( requested_hidden_mode = get_request_return_hidden_states_mode(
obj.return_hidden_states 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( raise ValueError(
"The server is not configured to return the hidden states. " "The requested return_hidden_states mode exceeds the "
"Please set `--enable-return-hidden-states` to enable this feature." 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 hidden states. "
"Please set `--return-hidden-states-mode last`, "
"`--return-hidden-states-mode full`, or the legacy "
"`--enable-return-hidden-states` flag."
) )
if ( if (
obj.custom_logit_processor obj.custom_logit_processor
@@ -34,6 +34,8 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardMode, ForwardMode,
PPProxyTensors, PPProxyTensors,
enable_num_token_non_padded, 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.forward_context import ForwardContext, forward_context
from sglang.srt.model_executor.runner_utils.capture_mode import model_capture_mode from sglang.srt.model_executor.runner_utils.capture_mode import model_capture_mode
@@ -562,9 +564,12 @@ class CPUGraphRunner:
# Parse args # Parse args
self.model_runner = model_runner self.model_runner = model_runner
self.device = model_runner.device self.device = model_runner.device
self.enable_return_hidden_states = ( self.return_hidden_states_mode = (
model_runner.server_args.enable_return_hidden_states 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) # bs -> compiled fn (text-only / skip_cross_attention=True)
self.graphs = {} self.graphs = {}
# bs -> compiled fn (cross-attention / skip_cross_attention=False, enc-dec only) # 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.pp_size = model_runner.server_args.pp_size
self.capture_forward_mode = ForwardMode.DECODE 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. # Static capture width: CPU graphs are decode-only.
self.captured_req_width = 1 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 ( assert (
not self.model_runner.server_args.enable_lora not self.model_runner.server_args.enable_lora
), "CPUGraphRunner does not support LoRA yet." ), "CPUGraphRunner does not support LoRA yet."
@@ -703,21 +704,7 @@ class CPUGraphRunner:
else forward_batch.batch_size <= self.max_bs else forward_batch.batch_size <= self.max_bs
) )
requested_capture_hidden_mode = max( return is_bs_supported
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
def capture(self) -> None: def capture(self) -> None:
capture_range = ( capture_range = (
@@ -793,9 +780,9 @@ class CPUGraphRunner:
encoder_lens = None encoder_lens = None
spec_info = self.get_spec_info(num_tokens) spec_info = self.get_spec_info(num_tokens)
if self.capture_hidden_mode != CaptureHiddenMode.FULL: self.capture_hidden_mode = get_required_capture_hidden_mode(
self.capture_hidden_mode = ( self.capture_hidden_mode,
spec_info.capture_hidden_mode if spec_info else CaptureHiddenMode.NULL spec_info,
) )
forward_batch = ForwardBatch( forward_batch = ForwardBatch(
@@ -870,43 +857,19 @@ class CPUGraphRunner:
self.captured_forward_batches_cross[bs] = forward_batch self.captured_forward_batches_cross[bs] = forward_batch
return forward, out return forward, out
def recapture_if_needed(self, forward_batch: ForwardBatch): def _validate_capture_hidden_mode(self, forward_batch: ForwardBatch) -> None:
if self.capture_hidden_mode < forward_batch.capture_hidden_mode:
# If the required capture_hidden_mode changes, we need to recapture the graph raise RuntimeError(
"The runtime hidden-state mode exceeds the fixed CPU graph "
# These are the different factors that can influence the capture_hidden_mode f"capture mode ({self.capture_hidden_mode.name})."
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 prepare_replay( def prepare_replay(
self, self,
forward_batch: ForwardBatch, forward_batch: ForwardBatch,
skip: bool = False, 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 graphs = self.graphs_cross if not skip else self.graphs
cfbs = ( cfbs = (
@@ -32,7 +32,7 @@ import warnings
from dataclasses import dataclass from dataclasses import dataclass
from enum import IntEnum, auto from enum import IntEnum, auto
from functools import total_ordering 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 import torch
@@ -231,6 +231,25 @@ def register_attn_tp_sequence_sharded_predicate(
_attn_tp_sequence_sharded_predicate = 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]: 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. """(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 # init_new must not mutate the input ScheduleBatch; per-forward
# overrides go through explicit keyword arguments. # overrides go through explicit keyword arguments.
# capture_hidden_mode=None means no override: derive from # capture_hidden_mode=None means no override: capture the server's
# SB.return_hidden_states / spec_info.capture_hidden_mode. # configured maximum so lower-mode requests can share one graph.
if capture_hidden_mode is None: if capture_hidden_mode is None:
if batch.return_hidden_states: request_capture_hidden_mode = (
capture_hidden_mode = CaptureHiddenMode.FULL CaptureHiddenMode.NULL
elif batch.spec_info is not None: if model_runner.is_draft_worker
capture_hidden_mode = getattr( else max(
batch.spec_info, "capture_hidden_mode", CaptureHiddenMode.NULL batch.return_hidden_states_mode,
get_server_return_hidden_states_mode(model_runner.server_args),
)
)
capture_hidden_mode = get_required_capture_hidden_mode(
request_capture_hidden_mode,
batch.spec_info,
) )
else:
capture_hidden_mode = CaptureHiddenMode.NULL
# extend-mode-only fields are None on decode/idle # extend-mode-only fields are None on decode/idle
if batch.forward_mode.is_decode_or_idle(): if batch.forward_mode.is_decode_or_idle():
@@ -19,6 +19,10 @@ from sglang.srt.model_executor.cuda_graph_config import (
Phase, Phase,
check_cuda_graph_backend, 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.graph_shared_output import GraphSharedOutput
from sglang.srt.model_executor.hook_manager import register_forward_hooks from sglang.srt.model_executor.hook_manager import register_forward_hooks
from sglang.srt.model_executor.model_runner_components.layer_setup import ( 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: if model_runner.is_draft_worker and not force_for_draft_worker:
return None return None
# Skip prefill CG for EAGLE target on tc_piecewise: that backend # Skip prefill CG for EAGLE target on tc_piecewise when the fixed server
# captures CaptureHiddenMode.NULL while runtime requests FULL, so # capture ceiling is below FULL. EAGLE target prefill requests FULL, so a
# the captured graph is dead, and capturing it perturbs FP4 / # NULL or LAST graph is dead; capturing it can perturb FP4/TRTLLM-MoE
# TRTLLM-MoE state and corrupts decode replay (see #28386). BCG # state and corrupt decode replay (see #28386 and #28870). BCG captures
# captures FULL for EAGLE target in PrefillCudaGraphRunner.__init__ # FULL for EAGLE target in PrefillCudaGraphRunner.__init__, so it does not
# (restored from #25795), so it does NOT need this skip. # need this skip.
if ( if (
model_runner.spec_algorithm.is_eagle() model_runner.spec_algorithm.is_eagle()
and not model_runner.is_draft_worker 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) and not check_cuda_graph_backend(Phase.PREFILL, Backend.BREAKABLE)
): ):
logger.info( logger.info(
@@ -38,6 +38,7 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardMode, ForwardMode,
NgramEmbeddingInfo, NgramEmbeddingInfo,
PPProxyTensors, PPProxyTensors,
get_server_return_hidden_states_mode,
) )
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
from sglang.srt.model_executor.runner.flashinfer_autotune import ( from sglang.srt.model_executor.runner.flashinfer_autotune import (
@@ -201,9 +202,12 @@ class BaseRunner(ABC):
self.dp_size = get_parallel().dp_size self.dp_size = get_parallel().dp_size
self.pp_size = model_runner.server_args.pp_size self.pp_size = model_runner.server_args.pp_size
self.enable_pdmux = model_runner.server_args.enable_pdmux self.enable_pdmux = model_runner.server_args.enable_pdmux
self.enable_return_hidden_states = ( self.return_hidden_states_mode = (
model_runner.server_args.enable_return_hidden_states 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_size = get_parallel().attn_tp_size
self.attn_tp_rank = get_parallel().attn_tp_rank self.attn_tp_rank = get_parallel().attn_tp_rank
self.tbo_plugin = TboCudaGraphRunnerPlugin() self.tbo_plugin = TboCudaGraphRunnerPlugin()
@@ -368,7 +372,11 @@ class BaseRunner(ABC):
capture_forward_mode = ForwardMode.DECODE capture_forward_mode = ForwardMode.DECODE
else: else:
capture_forward_mode = ForwardMode.EXTEND 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 num_tokens_per_req = 1
if mr.spec_algorithm.is_speculative(): if mr.spec_algorithm.is_speculative():
if mr.is_draft_worker: if mr.is_draft_worker:
@@ -378,9 +386,6 @@ class BaseRunner(ABC):
capture_forward_mode = ForwardMode.TARGET_VERIFY capture_forward_mode = ForwardMode.TARGET_VERIFY
num_tokens_per_req = mr.decode_num_tokens_per_req() 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 num_tokens = batch_size * num_tokens_per_req
# Caller owns the shape: passes a static buffer >= the dummy shape; no # Caller owns the shape: passes a static buffer >= the dummy shape; no
@@ -62,6 +62,7 @@ from sglang.srt.model_executor.forward_batch_info import (
PPProxyTensors, PPProxyTensors,
compute_local_num_token_non_padded, compute_local_num_token_non_padded,
enable_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.forward_context import ForwardContext, forward_context
from sglang.srt.model_executor.runner.base_cuda_graph_runner import ( from sglang.srt.model_executor.runner.base_cuda_graph_runner import (
@@ -258,7 +259,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
# --- capture mode + tokens-per-bs ------------------------------ # --- capture mode + tokens-per-bs ------------------------------
self.capture_forward_mode = ForwardMode.DECODE self.capture_forward_mode = ForwardMode.DECODE
self.capture_hidden_mode = CaptureHiddenMode.NULL self.capture_hidden_mode = self.return_hidden_states_mode
# Static capture width. # Static capture width.
self.captured_req_width = model_runner.decode_num_tokens_per_req( self.captured_req_width = model_runner.decode_num_tokens_per_req(
num_draft_tokens=self.speculative_num_draft_tokens num_draft_tokens=self.speculative_num_draft_tokens
@@ -304,10 +305,6 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
"feature." "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 # Attention backend
self.max_bs = max(self.capture_bs) self.max_bs = max(self.capture_bs)
self.max_num_token = self.max_bs * self.captured_req_width self.max_num_token = self.max_bs * self.captured_req_width
@@ -557,19 +554,6 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
else True 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 = ( is_tbo_supported = (
forward_batch.can_run_tbo if self.enable_two_batch_overlap else True forward_batch.can_run_tbo if self.enable_two_batch_overlap else True
) )
@@ -587,7 +571,6 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
is_bs_supported is_bs_supported
and is_encoder_lens_supported and is_encoder_lens_supported
and is_tbo_supported and is_tbo_supported
and capture_hidden_mode_matches
and is_ngram_supported and is_ngram_supported
) )
@@ -610,18 +593,8 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
else True 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 = ( capture_hidden_mode_matches = (
requested_capture_hidden_mode == CaptureHiddenMode.NULL forward_batch.capture_hidden_mode <= self.capture_hidden_mode
or requested_capture_hidden_mode == self.capture_hidden_mode
) )
return ( return (
@@ -751,9 +724,9 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
global_dp_buffer_len = None global_dp_buffer_len = None
spec_info = self.get_spec_info(num_tokens) spec_info = self.get_spec_info(num_tokens)
if self.capture_hidden_mode != CaptureHiddenMode.FULL: self.capture_hidden_mode = get_required_capture_hidden_mode(
self.capture_hidden_mode = ( self.capture_hidden_mode,
spec_info.capture_hidden_mode if spec_info else CaptureHiddenMode.NULL spec_info,
) )
if self.model_runner.server_args.enable_lora: if self.model_runner.server_args.enable_lora:
@@ -1023,38 +996,12 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
post_warmup_hook=post_warmup_hook, post_warmup_hook=post_warmup_hook,
) )
def recapture_if_needed(self, forward_batch: ForwardBatch): def _validate_capture_hidden_mode(self, forward_batch: ForwardBatch) -> None:
if self.capture_hidden_mode < forward_batch.capture_hidden_mode:
# If the required capture_hidden_mode changes, we need to recapture the graph raise RuntimeError(
"The runtime hidden-state mode exceeds the fixed CUDA graph "
# These are the different factors that can influence the capture_hidden_mode f"capture mode ({self.capture_hidden_mode.name})."
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 load_batch( def load_batch(
self, self,
@@ -1104,7 +1051,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
return return
buffers = self.buffers buffers = self.buffers
self.recapture_if_needed(forward_batch) self._validate_capture_hidden_mode(forward_batch)
raw_bs = forward_batch.batch_size raw_bs = forward_batch.batch_size
@@ -277,16 +277,12 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
self.prefill_backend_name == Backend.BREAKABLE self.prefill_backend_name == Backend.BREAKABLE
and model_runner.spec_algorithm.is_eagle() 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: if is_breakable_eagle and model_runner.is_draft_worker:
self.capture_hidden_mode = CaptureHiddenMode.LAST 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 self.capture_hidden_mode = CaptureHiddenMode.FULL
else: 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() self.mamba_track_enabled = self._is_mamba_track_enabled()
@@ -1055,7 +1051,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
return False return False
if ( if (
capture_hidden_mode is not None capture_hidden_mode is not None
and capture_hidden_mode != self.capture_hidden_mode and self.capture_hidden_mode < capture_hidden_mode
): ):
return False return False
if return_logprob and not self._uses_eager_prefill_tail(): if return_logprob and not self._uses_eager_prefill_tail():
@@ -1651,9 +1647,17 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
"PPProxyTensors is not supported in PrefillCudaGraphRunner yet." "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( def execute(
self, forward_batch: ForwardBatch, **kwargs self, forward_batch: ForwardBatch, **kwargs
) -> Union[LogitsProcessorOutput, PPProxyTensors, EmbeddingPoolerOutput]: ) -> Union[LogitsProcessorOutput, PPProxyTensors, EmbeddingPoolerOutput]:
self._validate_capture_hidden_mode(forward_batch)
with self.backend.replay_session(): with self.backend.replay_session():
static_forward_batch = self.load_batch(forward_batch, **kwargs) static_forward_batch = self.load_batch(forward_batch, **kwargs)
static_num_tokens = len(static_forward_batch.input_ids) static_num_tokens = len(static_forward_batch.input_ids)
+26 -1
View File
@@ -3305,8 +3305,21 @@ class ServerArgs:
NS("exec.features"), NS("exec.features"),
] = False ] = False
enable_return_hidden_states: A[ 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 ] = 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[ enable_return_routed_experts: A[
bool, bool,
"Enable returning routed experts of each layer with responses.", "Enable returning routed experts of each layer with responses.",
@@ -3408,6 +3421,7 @@ class ServerArgs:
# _handle_model_specific_adjustments never runs. # _handle_model_specific_adjustments never runs.
self._resolved_overrides = [] self._resolved_overrides = []
self._handle_return_hidden_states_mode()
if self.model_path.lower() in ["none", "dummy"]: if self.model_path.lower() in ["none", "dummy"]:
return return
@@ -3584,6 +3598,17 @@ class ServerArgs:
materialize_declarations(self) 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): def _handle_model_capability_adjustments(self):
if parse_connector_type(self.model_path) == ConnectorType.INSTANCE: if parse_connector_type(self.model_path) == ConnectorType.INSTANCE:
return return
+91 -2
View File
@@ -32,7 +32,7 @@ class TestHiddenState(CustomTestCase):
model_path=cls.model_path, model_path=cls.model_path,
random_seed=42, random_seed=42,
skip_tokenizer_init=True, skip_tokenizer_init=True,
enable_return_hidden_states=True, return_hidden_states_mode="full",
mem_fraction_static=0.7, mem_fraction_static=0.7,
) )
@@ -53,8 +53,12 @@ class TestHiddenState(CustomTestCase):
return_hidden_states=True, return_hidden_states=True,
) )
expected_num_hidden_states = self.sampling_params["max_new_tokens"]
for output in outputs: 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"])): for i in range(len(output["meta_info"]["hidden_states"])):
assert isinstance(output["meta_info"]["hidden_states"][i], list) assert isinstance(output["meta_info"]["hidden_states"][i], list)
output["meta_info"]["hidden_states"][i] = torch.tensor( 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): def test_repeatedly_changes_hidden_states(self):
outputs_completion_first_round = self.engine.generate( outputs_completion_first_round = self.engine.generate(
input_ids=self.input_ids, input_ids=self.input_ids,
@@ -100,7 +100,7 @@ class BaseTestOpenAIServerWithHiddenStates(ABC):
) )
for choice in response.choices: 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: if return_hidden_states:
assert choice.hidden_states is not None, "hidden_states was None" assert choice.hidden_states is not None, "hidden_states was None"
@@ -139,7 +139,7 @@ class BaseTestOpenAIServerWithHiddenStates(ABC):
usage = response.usage usage = response.usage
for choice in response.choices: for choice in response.choices:
if hasattr(choice, "hidden_states"): if hasattr(choice, "hidden_states"):
assert return_hidden_states assert bool(return_hidden_states)
assert choice.hidden_states is not None assert choice.hidden_states is not None
hidden_states_list.append(choice.hidden_states) hidden_states_list.append(choice.hidden_states)
@@ -169,7 +169,7 @@ class BaseTestOpenAIServerWithHiddenStates(ABC):
) )
for choice in response.choices: 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: if return_hidden_states:
assert choice.hidden_states is not None, "hidden_states was None" assert choice.hidden_states is not None, "hidden_states was None"
@@ -196,7 +196,7 @@ class BaseTestOpenAIServerWithHiddenStates(ABC):
for response in generator: for response in generator:
for choice in response.choices: for choice in response.choices:
if hasattr(choice.delta, "hidden_states"): if hasattr(choice.delta, "hidden_states"):
assert return_hidden_states assert bool(return_hidden_states)
assert choice.delta.hidden_states is not None assert choice.delta.hidden_states is not None
hidden_states_list.append(choice.delta.hidden_states) hidden_states_list.append(choice.delta.hidden_states)
@@ -227,7 +227,7 @@ class TestOpenAIServerWithHiddenStatesEnabled(
) )
cls.base_url += "/v1" cls.base_url += "/v1"
cls.tokenizer = get_tokenizer(DEFAULT_SMALL_MODEL_NAME_FOR_TEST) 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.use_list_input = [True, False]
cls.parallel_sample_nums = [1, 2] cls.parallel_sample_nums = [1, 2]
@@ -253,7 +253,7 @@ class TestOpenAIServerWithHiddenStatesEnabledAndCUDAGraphDisabled(
) )
cls.base_url += "/v1" cls.base_url += "/v1"
cls.tokenizer = get_tokenizer(DEFAULT_SMALL_MODEL_NAME_FOR_TEST) 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.use_list_input = [True, False]
cls.parallel_sample_nums = [1] cls.parallel_sample_nums = [1]
@@ -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()
@@ -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()
@@ -164,6 +164,48 @@ class TestGenerateReqInputNormalization(CustomTestCase):
# Check text expansion # Check text expansion
self.assertEqual(req.text, expected_text) 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): 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.""" """Test that when some batch items have images and others None, parallel expansion works correctly."""
req = copy.deepcopy(self.base_req) req = copy.deepcopy(self.base_req)
@@ -480,14 +522,14 @@ class TestGenerateReqInputNormalization(CustomTestCase):
logprob_start_len=[10, 5], logprob_start_len=[10, 5],
top_logprobs_num=[5, 3], top_logprobs_num=[5, 3],
token_ids_logprob=[[7, 8, 9], [4, 5, 6]], 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() req.normalize_batch_and_arguments()
self.assertEqual(req.return_logprob, [True, False]) self.assertEqual(req.return_logprob, [True, False])
self.assertEqual(req.logprob_start_len, [10, 5]) self.assertEqual(req.logprob_start_len, [10, 5])
self.assertEqual(req.top_logprobs_num, [5, 3]) self.assertEqual(req.top_logprobs_num, [5, 3])
self.assertEqual(req.token_ids_logprob, [[7, 8, 9], [4, 5, 6]]) 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): def test_custom_logit_processor_normalization(self):
"""Test normalization of custom_logit_processor.""" """Test normalization of custom_logit_processor."""
@@ -559,7 +601,7 @@ class TestGenerateReqInputNormalization(CustomTestCase):
modalities=["image", "image"], modalities=["image", "image"],
lora_path=["path1", "path2"], lora_path=["path1", "path2"],
custom_logit_processor=["processor1", "processor2"], custom_logit_processor=["processor1", "processor2"],
return_hidden_states=True, return_hidden_states=[True, "last"],
) )
req.normalize_batch_and_arguments() req.normalize_batch_and_arguments()
@@ -580,6 +622,7 @@ class TestGenerateReqInputNormalization(CustomTestCase):
self.assertEqual(item0.lora_path, "path1") self.assertEqual(item0.lora_path, "path1")
self.assertEqual(item0.custom_logit_processor, "processor1") self.assertEqual(item0.custom_logit_processor, "processor1")
self.assertEqual(item0.return_hidden_states, True) self.assertEqual(item0.return_hidden_states, True)
self.assertEqual(req[1].return_hidden_states, "last")
def test_getitem_preserves_return_prompt_token_ids(self): def test_getitem_preserves_return_prompt_token_ids(self):
"""Batch subrequests must keep the prompt-token-id return flag.""" """Batch subrequests must keep the prompt-token-id return flag."""
@@ -11,6 +11,7 @@ maybe_stub_sgl_kernel()
from sglang.srt.managers.hisparse_coordinator import HiSparseCoordinator # noqa: E402 from sglang.srt.managers.hisparse_coordinator import HiSparseCoordinator # noqa: E402
from sglang.srt.managers.scheduler import Scheduler # 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") 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, return_logprob=False,
grammar=None, grammar=None,
return_hidden_states=False, return_hidden_states=False,
return_hidden_states_mode=CaptureHiddenMode.NULL,
is_prefill_only=False, is_prefill_only=False,
) )
@@ -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()
@@ -6,8 +6,13 @@ from unittest.mock import patch
import torch 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 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.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 ( from sglang.srt.model_executor.runner.prefill_cuda_graph_runner import (
PrefillCudaGraphRunner, PrefillCudaGraphRunner,
) )
@@ -56,6 +61,29 @@ class _FakeKVIndexKernel:
class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase): 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): def test_prefix_chunk_capacity_is_aggregate_and_can_be_overridden(self):
model_runner = SimpleNamespace( model_runner = SimpleNamespace(
server_args=SimpleNamespace( server_args=SimpleNamespace(
@@ -212,7 +240,7 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
runner = PrefillCudaGraphRunner.__new__(PrefillCudaGraphRunner) runner = PrefillCudaGraphRunner.__new__(PrefillCudaGraphRunner)
runner._capture_req_slots = 4 runner._capture_req_slots = 4
runner.enable_lora = False runner.enable_lora = False
runner.capture_hidden_mode = None runner.capture_hidden_mode = CaptureHiddenMode.NULL
runner.max_num_tokens = 32 runner.max_num_tokens = 32
runner.capture_num_tokens = [4] runner.capture_num_tokens = [4]
runner.backend = SimpleNamespace() runner.backend = SimpleNamespace()
@@ -227,7 +255,7 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
input_embeds=None, input_embeds=None,
replace_embeds=None, replace_embeds=None,
forward_mode=SimpleNamespace(is_target_verify=lambda: False), forward_mode=SimpleNamespace(is_target_verify=lambda: False),
capture_hidden_mode=None, capture_hidden_mode=CaptureHiddenMode.NULL,
global_num_tokens_cpu=None, global_num_tokens_cpu=None,
return_logprob=False, return_logprob=False,
extend_prefix_lens_cpu=[8], extend_prefix_lens_cpu=[8],
@@ -42,6 +42,45 @@ _mock_device.start()
class TestPrepareServerArgs(CustomTestCase): 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): def test_config_nested_dict_args_are_json(self):
with tempfile.NamedTemporaryFile(mode="w", suffix=".yaml", delete=False) as f: with tempfile.NamedTemporaryFile(mode="w", suffix=".yaml", delete=False) as f:
f.write("mm-process-config:\n image:\n resize: 128\n") f.write("mm-process-config:\n image:\n resize: 128\n")