Add per-request decode tp size (#14678)
Co-authored-by: Byron Hsu <byronhsu1230@gmail.com>
This commit is contained in:
co-authored by
Byron Hsu
parent
0e0b0c0566
commit
ce4e836be5
@@ -196,6 +196,7 @@ class GenerateReqInput(BaseReq):
|
|||||||
bootstrap_port: Optional[Union[List[Optional[int]], int]] = None
|
bootstrap_port: Optional[Union[List[Optional[int]], int]] = None
|
||||||
bootstrap_room: Optional[Union[List[int], int]] = None
|
bootstrap_room: Optional[Union[List[int], int]] = None
|
||||||
bootstrap_pair_key: Optional[Union[List[str], str]] = None
|
bootstrap_pair_key: Optional[Union[List[str], str]] = None
|
||||||
|
decode_tp_size: Optional[Union[List[Optional[int]], int]] = None
|
||||||
|
|
||||||
# For reasoning
|
# For reasoning
|
||||||
reasoning: bool = False
|
reasoning: bool = False
|
||||||
@@ -619,6 +620,9 @@ class GenerateReqInput(BaseReq):
|
|||||||
if self.bootstrap_pair_key is not None
|
if self.bootstrap_pair_key is not None
|
||||||
else 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,
|
validation_time=self.validation_time,
|
||||||
data_parallel_rank=(
|
data_parallel_rank=(
|
||||||
self.data_parallel_rank if self.data_parallel_rank is not None else None
|
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_port: Optional[int] = None
|
||||||
bootstrap_room: Optional[int] = None
|
bootstrap_room: Optional[int] = None
|
||||||
bootstrap_pair_key: Optional[str] = None
|
bootstrap_pair_key: Optional[str] = None
|
||||||
|
decode_tp_size: Optional[int] = None
|
||||||
|
|
||||||
# For reasoning
|
# For reasoning
|
||||||
reasoning: bool = False
|
reasoning: bool = False
|
||||||
|
|||||||
@@ -841,6 +841,9 @@ class TokenizerManager(TokenizerCommunicatorMixin):
|
|||||||
f"The input_ids {input_ids} contains values greater than the vocab size ({vocab_size})."
|
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(
|
def _create_tokenized_object(
|
||||||
self,
|
self,
|
||||||
obj: Union[GenerateReqInput, EmbeddingReqInput],
|
obj: Union[GenerateReqInput, EmbeddingReqInput],
|
||||||
@@ -858,7 +861,7 @@ class TokenizerManager(TokenizerCommunicatorMixin):
|
|||||||
sampling_kwargs = {**self.preferred_sampling_params, **obj.sampling_params}
|
sampling_kwargs = {**self.preferred_sampling_params, **obj.sampling_params}
|
||||||
else:
|
else:
|
||||||
sampling_kwargs = obj.sampling_params
|
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.normalize(self.tokenizer)
|
||||||
sampling_params.verify(self.model_config.vocab_size)
|
sampling_params.verify(self.model_config.vocab_size)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user