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