diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index 39bbfaf3c..b4f9d3335 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -196,6 +196,7 @@ class GenerateReqInput(BaseReq): bootstrap_port: Optional[Union[List[Optional[int]], int]] = None bootstrap_room: Optional[Union[List[int], int]] = None bootstrap_pair_key: Optional[Union[List[str], str]] = None + decode_tp_size: Optional[Union[List[Optional[int]], int]] = None # For reasoning reasoning: bool = False @@ -619,6 +620,9 @@ class GenerateReqInput(BaseReq): if self.bootstrap_pair_key is not None else None ), + decode_tp_size=( + self.decode_tp_size[i] if self.decode_tp_size is not None else None + ), validation_time=self.validation_time, data_parallel_rank=( self.data_parallel_rank if self.data_parallel_rank is not None else None @@ -677,6 +681,7 @@ class TokenizedGenerateReqInput(BaseReq): bootstrap_port: Optional[int] = None bootstrap_room: Optional[int] = None bootstrap_pair_key: Optional[str] = None + decode_tp_size: Optional[int] = None # For reasoning reasoning: bool = False diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index 1396a8248..787de1257 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -841,6 +841,9 @@ class TokenizerManager(TokenizerCommunicatorMixin): f"The input_ids {input_ids} contains values greater than the vocab size ({vocab_size})." ) + def _get_sampling_params(self, sampling_kwargs: Dict) -> SamplingParams: + return SamplingParams(**sampling_kwargs) + def _create_tokenized_object( self, obj: Union[GenerateReqInput, EmbeddingReqInput], @@ -858,7 +861,7 @@ class TokenizerManager(TokenizerCommunicatorMixin): sampling_kwargs = {**self.preferred_sampling_params, **obj.sampling_params} else: sampling_kwargs = obj.sampling_params - sampling_params = SamplingParams(**sampling_kwargs) + sampling_params = self._get_sampling_params(sampling_kwargs) sampling_params.normalize(self.tokenizer) sampling_params.verify(self.model_config.vocab_size)