[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:
python hidden_states.py
python hidden_states_engine.py
Note that each time you change the `return_hidden_states` parameter,
the cuda graph will be recaptured, which might lead to a performance hit.
So avoid getting hidden states and completions alternately.
CUDA graphs use the configured maximum hidden-state mode. Requests may select
that mode or a weaker one without triggering mode-dependent recapture.
"""
import torch
@@ -22,7 +21,7 @@ def main():
# Create an LLM.
llm = sgl.Engine(
model_path="Alibaba-NLP/gte-Qwen2-1.5B-instruct",
enable_return_hidden_states=True,
return_hidden_states_mode="last",
)
sampling_params = {
@@ -32,16 +31,15 @@ def main():
}
outputs = llm.generate(
prompts, sampling_params=sampling_params, return_hidden_states=True
prompts, sampling_params=sampling_params, return_hidden_states="last"
)
llm.shutdown()
for prompt, output in zip(prompts, outputs):
for i in range(len(output["meta_info"]["hidden_states"])):
output["meta_info"]["hidden_states"][i] = torch.tensor(
output["meta_info"]["hidden_states"][i], dtype=torch.bfloat16
)
hidden_state = torch.tensor(
output["meta_info"]["hidden_states"], dtype=torch.bfloat16
)
print("===============================")
print(
f"Prompt: {prompt}\n"
@@ -49,14 +47,8 @@ def main():
f"Prompt_Tokens: {output['meta_info']['prompt_tokens']}\t"
f"Completion_tokens: {output['meta_info']['completion_tokens']}"
)
print("Hidden states: ")
hidden_states = torch.cat(
[
i.unsqueeze(0) if len(i.shape) == 1 else i
for i in output["meta_info"]["hidden_states"]
]
)
print(hidden_states)
print("Last hidden state: ")
print(hidden_state)
print()
@@ -3,9 +3,8 @@ Usage:
python hidden_states_server.py
Note that each time you change the `return_hidden_states` parameter,
the cuda graph will be recaptured, which might lead to a performance hit.
So avoid getting hidden states and completions alternately.
CUDA graphs use the configured maximum hidden-state mode. Requests may select
that mode or a weaker one without triggering mode-dependent recapture.
"""
import requests
@@ -23,7 +22,9 @@ else:
def main():
# Launch the server
server_process, port = launch_server_cmd(
"python -m sglang.launch_server --model-path Alibaba-NLP/gte-Qwen2-1.5B-instruct --enable-return-hidden-states --host 0.0.0.0"
"python -m sglang.launch_server --model-path "
"Alibaba-NLP/gte-Qwen2-1.5B-instruct "
"--return-hidden-states-mode last --host 0.0.0.0"
)
wait_for_server(f"http://localhost:{port}", process=server_process)
@@ -43,7 +44,7 @@ def main():
json_data = {
"text": prompts,
"sampling_params": sampling_params,
"return_hidden_states": True,
"return_hidden_states": "last",
}
response = requests.post(
@@ -55,10 +56,9 @@ def main():
outputs = response.json()
for prompt, output in zip(prompts, outputs):
for i in range(len(output["meta_info"]["hidden_states"])):
output["meta_info"]["hidden_states"][i] = torch.tensor(
output["meta_info"]["hidden_states"][i], dtype=torch.bfloat16
)
hidden_state = torch.tensor(
output["meta_info"]["hidden_states"], dtype=torch.bfloat16
)
print("===============================")
print(
f"Prompt: {prompt}\n"
@@ -66,14 +66,8 @@ def main():
f"Prompt_Tokens: {output['meta_info']['prompt_tokens']}\t"
f"Completion_tokens: {output['meta_info']['completion_tokens']}"
)
print("Hidden states: ")
hidden_states = torch.cat(
[
i.unsqueeze(0) if len(i.shape) == 1 else i
for i in output["meta_info"]["hidden_states"]
]
)
print(hidden_states)
print("Last hidden state: ")
print(hidden_state)
print()
+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
+91 -2
View File
@@ -32,7 +32,7 @@ class TestHiddenState(CustomTestCase):
model_path=cls.model_path,
random_seed=42,
skip_tokenizer_init=True,
enable_return_hidden_states=True,
return_hidden_states_mode="full",
mem_fraction_static=0.7,
)
@@ -53,8 +53,12 @@ class TestHiddenState(CustomTestCase):
return_hidden_states=True,
)
expected_num_hidden_states = self.sampling_params["max_new_tokens"]
for output in outputs:
self.assertEqual(len(output["meta_info"]["hidden_states"]), 8)
self.assertEqual(
len(output["meta_info"]["hidden_states"]),
expected_num_hidden_states,
)
for i in range(len(output["meta_info"]["hidden_states"])):
assert isinstance(output["meta_info"]["hidden_states"][i], list)
output["meta_info"]["hidden_states"][i] = torch.tensor(
@@ -103,6 +107,91 @@ class TestHiddenState(CustomTestCase):
)
)
def test_return_last_hidden_state(self):
outputs = self.engine.generate(
input_ids=self.input_ids,
sampling_params=self.sampling_params,
return_hidden_states="last",
)
model = AutoModelForCausalLM.from_pretrained(
self.model_path, torch_dtype=torch.bfloat16, device_map=get_device()
)
for input_id, output in zip(self.input_ids, outputs):
sg_hidden_state = torch.tensor(
output["meta_info"]["hidden_states"], dtype=torch.bfloat16
).to(get_device())
self.assertEqual(sg_hidden_state.dim(), 1)
with torch.inference_mode():
hf_out = model(
torch.tensor(
[input_id + output["output_ids"][:-1]], device=model.device
),
output_hidden_states=True,
)
hf_last_hidden_state = hf_out["hidden_states"][-1][0, -1]
atol = 0.8
self.assertTrue(
torch.allclose(
hf_last_hidden_state,
sg_hidden_state,
atol=atol,
rtol=0,
)
)
def test_mixed_return_hidden_states_modes(self):
outputs = self.engine.generate(
input_ids=self.input_ids + [self.input_ids[0]],
sampling_params=self.sampling_params,
return_hidden_states=[False, True, "last"],
)
self.assertNotIn("hidden_states", outputs[0]["meta_info"])
full_hidden_states = outputs[1]["meta_info"]["hidden_states"]
last_hidden_state = outputs[2]["meta_info"]["hidden_states"]
self.assertIsInstance(full_hidden_states, list)
self.assertEqual(
len(full_hidden_states), self.sampling_params["max_new_tokens"]
)
self.assertEqual(torch.tensor(full_hidden_states[0]).dim(), 2)
last_hidden_state = torch.tensor(last_hidden_state)
self.assertEqual(last_hidden_state.dim(), 1)
def test_mixed_return_hidden_states_modes_with_warm_cache(self):
# Prime the radix cache so each repeated prompt only extends its
# uncached suffix during the mixed-mode prefill.
self.engine.generate(
input_ids=self.input_ids,
sampling_params={"temperature": 0, "max_new_tokens": 1},
return_hidden_states=False,
)
outputs = self.engine.generate(
input_ids=self.input_ids + [self.input_ids[1]],
sampling_params={"temperature": 0, "max_new_tokens": 1},
return_hidden_states=[True, True, "last"],
)
self.assertEqual(
torch.tensor(outputs[0]["meta_info"]["hidden_states"][0]).dim(),
2,
)
self.assertEqual(
torch.tensor(outputs[2]["meta_info"]["hidden_states"]).dim(),
1,
)
torch.testing.assert_close(
torch.tensor(outputs[1]["meta_info"]["hidden_states"][0])[-1],
torch.tensor(outputs[2]["meta_info"]["hidden_states"]),
)
def test_repeatedly_changes_hidden_states(self):
outputs_completion_first_round = self.engine.generate(
input_ids=self.input_ids,
@@ -100,7 +100,7 @@ class BaseTestOpenAIServerWithHiddenStates(ABC):
)
for choice in response.choices:
assert hasattr(choice, "hidden_states") == return_hidden_states
assert hasattr(choice, "hidden_states") == bool(return_hidden_states)
if return_hidden_states:
assert choice.hidden_states is not None, "hidden_states was None"
@@ -139,7 +139,7 @@ class BaseTestOpenAIServerWithHiddenStates(ABC):
usage = response.usage
for choice in response.choices:
if hasattr(choice, "hidden_states"):
assert return_hidden_states
assert bool(return_hidden_states)
assert choice.hidden_states is not None
hidden_states_list.append(choice.hidden_states)
@@ -169,7 +169,7 @@ class BaseTestOpenAIServerWithHiddenStates(ABC):
)
for choice in response.choices:
assert hasattr(choice, "hidden_states") == return_hidden_states
assert hasattr(choice, "hidden_states") == bool(return_hidden_states)
if return_hidden_states:
assert choice.hidden_states is not None, "hidden_states was None"
@@ -196,7 +196,7 @@ class BaseTestOpenAIServerWithHiddenStates(ABC):
for response in generator:
for choice in response.choices:
if hasattr(choice.delta, "hidden_states"):
assert return_hidden_states
assert bool(return_hidden_states)
assert choice.delta.hidden_states is not None
hidden_states_list.append(choice.delta.hidden_states)
@@ -227,7 +227,7 @@ class TestOpenAIServerWithHiddenStatesEnabled(
)
cls.base_url += "/v1"
cls.tokenizer = get_tokenizer(DEFAULT_SMALL_MODEL_NAME_FOR_TEST)
cls.return_hidden_states = [False, True]
cls.return_hidden_states = [False, True, "last"]
cls.use_list_input = [True, False]
cls.parallel_sample_nums = [1, 2]
@@ -253,7 +253,7 @@ class TestOpenAIServerWithHiddenStatesEnabledAndCUDAGraphDisabled(
)
cls.base_url += "/v1"
cls.tokenizer = get_tokenizer(DEFAULT_SMALL_MODEL_NAME_FOR_TEST)
cls.return_hidden_states = [False, True]
cls.return_hidden_states = [False, True, "last"]
cls.use_list_input = [True, False]
cls.parallel_sample_nums = [1]
@@ -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
self.assertEqual(req.text, expected_text)
def test_return_hidden_states_expands_with_parallel_sampling(self):
req = GenerateReqInput(
text=["Prompt 1", "Prompt 2"],
sampling_params={"n": 2},
return_hidden_states=[False, "last"],
)
req.normalize_batch_and_arguments()
self.assertEqual(
req.return_hidden_states,
[False, "last", False, "last"],
)
self.assertEqual(
[req[i].return_hidden_states for i in range(4)],
[False, "last", False, "last"],
)
def test_return_hidden_states_batch_length_is_validated(self):
req = GenerateReqInput(
text=["Prompt 1", "Prompt 2"],
return_hidden_states=["last"],
)
with self.assertRaisesRegex(
ValueError,
"return_hidden_states should be equal to the batch size",
):
req.normalize_batch_and_arguments()
def test_return_hidden_states_batch_modes_are_validated(self):
req = GenerateReqInput(
text=["Prompt 1", "Prompt 2"],
return_hidden_states=[False, "invalid"],
)
with self.assertRaisesRegex(
ValueError,
"return_hidden_states must be a boolean or the string literal 'last'",
):
req.normalize_batch_and_arguments()
def test_mixed_none_and_images_with_parallel_samples(self):
"""Test that when some batch items have images and others None, parallel expansion works correctly."""
req = copy.deepcopy(self.base_req)
@@ -480,14 +522,14 @@ class TestGenerateReqInputNormalization(CustomTestCase):
logprob_start_len=[10, 5],
top_logprobs_num=[5, 3],
token_ids_logprob=[[7, 8, 9], [4, 5, 6]],
return_hidden_states=[False, False, True],
return_hidden_states=[False, True],
)
req.normalize_batch_and_arguments()
self.assertEqual(req.return_logprob, [True, False])
self.assertEqual(req.logprob_start_len, [10, 5])
self.assertEqual(req.top_logprobs_num, [5, 3])
self.assertEqual(req.token_ids_logprob, [[7, 8, 9], [4, 5, 6]])
self.assertEqual(req.return_hidden_states, [False, False, True])
self.assertEqual(req.return_hidden_states, [False, True])
def test_custom_logit_processor_normalization(self):
"""Test normalization of custom_logit_processor."""
@@ -559,7 +601,7 @@ class TestGenerateReqInputNormalization(CustomTestCase):
modalities=["image", "image"],
lora_path=["path1", "path2"],
custom_logit_processor=["processor1", "processor2"],
return_hidden_states=True,
return_hidden_states=[True, "last"],
)
req.normalize_batch_and_arguments()
@@ -580,6 +622,7 @@ class TestGenerateReqInputNormalization(CustomTestCase):
self.assertEqual(item0.lora_path, "path1")
self.assertEqual(item0.custom_logit_processor, "processor1")
self.assertEqual(item0.return_hidden_states, True)
self.assertEqual(req[1].return_hidden_states, "last")
def test_getitem_preserves_return_prompt_token_ids(self):
"""Batch subrequests must keep the prompt-token-id return flag."""
@@ -11,6 +11,7 @@ maybe_stub_sgl_kernel()
from sglang.srt.managers.hisparse_coordinator import HiSparseCoordinator # noqa: E402
from sglang.srt.managers.scheduler import Scheduler # noqa: E402
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode # noqa: E402
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
@@ -23,6 +24,7 @@ def _make_req(req_pool_idx, origin_input_ids, output_ids):
return_logprob=False,
grammar=None,
return_hidden_states=False,
return_hidden_states_mode=CaptureHiddenMode.NULL,
is_prefill_only=False,
)
@@ -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 sglang.srt.model_executor.model_runner_components.cuda_graph_setup as graph_setup
import sglang.srt.model_executor.runner.prefill_cuda_graph_runner as runner_module
from sglang.srt.model_executor.cuda_graph_config import Backend
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode
from sglang.srt.model_executor.model_runner_components.cuda_graph_setup import (
capture_prefill_graph,
)
from sglang.srt.model_executor.runner.prefill_cuda_graph_runner import (
PrefillCudaGraphRunner,
)
@@ -56,6 +61,29 @@ class _FakeKVIndexKernel:
class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
def test_eagle_target_tc_piecewise_skips_last_mode_capture(self):
eager_runner = object()
model_runner = SimpleNamespace(
is_draft_worker=False,
spec_algorithm=SimpleNamespace(is_eagle=lambda: True),
server_args=SimpleNamespace(
enable_return_hidden_states=True,
return_hidden_states_mode="last",
),
)
with patch.object(
graph_setup,
"check_cuda_graph_backend",
return_value=False,
):
runner = capture_prefill_graph(
model_runner=model_runner,
eager_runner=eager_runner,
)
self.assertIs(runner, eager_runner)
def test_prefix_chunk_capacity_is_aggregate_and_can_be_overridden(self):
model_runner = SimpleNamespace(
server_args=SimpleNamespace(
@@ -212,7 +240,7 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
runner = PrefillCudaGraphRunner.__new__(PrefillCudaGraphRunner)
runner._capture_req_slots = 4
runner.enable_lora = False
runner.capture_hidden_mode = None
runner.capture_hidden_mode = CaptureHiddenMode.NULL
runner.max_num_tokens = 32
runner.capture_num_tokens = [4]
runner.backend = SimpleNamespace()
@@ -227,7 +255,7 @@ class TestPrefillCudaGraphRunnerChunkedPrefix(CustomTestCase):
input_embeds=None,
replace_embeds=None,
forward_mode=SimpleNamespace(is_target_verify=lambda: False),
capture_hidden_mode=None,
capture_hidden_mode=CaptureHiddenMode.NULL,
global_num_tokens_cpu=None,
return_logprob=False,
extend_prefix_lens_cpu=[8],
@@ -42,6 +42,45 @@ _mock_device.start()
class TestPrepareServerArgs(CustomTestCase):
def test_return_hidden_states_mode_configuration(self):
disabled = ServerArgs(model_path="dummy")
self.assertFalse(disabled.enable_return_hidden_states)
self.assertIsNone(disabled.return_hidden_states_mode)
last = ServerArgs(
model_path="dummy",
return_hidden_states_mode="last",
)
self.assertTrue(last.enable_return_hidden_states)
self.assertEqual(last.return_hidden_states_mode, "last")
legacy_full = ServerArgs(
model_path="dummy",
enable_return_hidden_states=True,
)
self.assertTrue(legacy_full.enable_return_hidden_states)
self.assertEqual(legacy_full.return_hidden_states_mode, "full")
parsed_last = prepare_server_args(
[
"--model-path",
"dummy",
"--return-hidden-states-mode",
"last",
]
)
self.assertTrue(parsed_last.enable_return_hidden_states)
self.assertEqual(parsed_last.return_hidden_states_mode, "last")
with self.assertRaisesRegex(
ValueError,
"return_hidden_states_mode must be one of",
):
ServerArgs(
model_path="dummy",
return_hidden_states_mode="lst",
)
def test_config_nested_dict_args_are_json(self):
with tempfile.NamedTemporaryFile(mode="w", suffix=".yaml", delete=False) as f:
f.write("mm-process-config:\n image:\n resize: 128\n")