[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
+4 -2
View File
@@ -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,
+7 -2
View File
@@ -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,
@@ -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
@@ -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
)
@@ -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,
)
],
+16 -4
View File
@@ -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(
+26 -3
View File
@@ -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
+65 -5
View File
@@ -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,
+5 -2
View File
@@ -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
@@ -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:
@@ -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:
@@ -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
@@ -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 = (
@@ -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():
@@ -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(
@@ -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
@@ -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
@@ -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)
+26 -1
View File
@@ -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