[Feature] Support return_hidden_states="last" (#30177)
Co-authored-by: litao.dream <litao.dream@bytedance.com>
This commit is contained in:
@@ -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()
|
||||
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
],
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user