Generate output for the given prompt token ids in the synchronous mode. Synchronous generation accepts either single prompt or batched prompts.
(
self,
prompt_token_ids: Union[List[int], List[List[int]]],
sampling_params: Union[SamplingParams, List[SamplingParams]],
query_token_ids: Optional[Union[torch.Tensor, np.ndarray, list]] = None,
lora_request: Optional[Union[LoRARequest, List[LoRARequest]]] = None,
prompt_adapter_request: Optional[Union[
PromptAdapterRequest, List[PromptAdapterRequest]]] = None,
disaggregated_params: Optional[DisaggregatedParams] = None,
)
| 161 | return result |
| 162 | |
| 163 | def generate( |
| 164 | self, |
| 165 | prompt_token_ids: Union[List[int], List[List[int]]], |
| 166 | sampling_params: Union[SamplingParams, List[SamplingParams]], |
| 167 | query_token_ids: Optional[Union[torch.Tensor, np.ndarray, list]] = None, |
| 168 | lora_request: Optional[Union[LoRARequest, List[LoRARequest]]] = None, |
| 169 | prompt_adapter_request: Optional[Union[ |
| 170 | PromptAdapterRequest, List[PromptAdapterRequest]]] = None, |
| 171 | disaggregated_params: Optional[DisaggregatedParams] = None, |
| 172 | ) -> Union[GenerationResult, List[GenerationResult]]: |
| 173 | """Generate output for the given prompt token ids in the synchronous mode. |
| 174 | Synchronous generation accepts either single prompt or batched prompts. |
| 175 | """ |
| 176 | unbatched = isinstance(prompt_token_ids[0], int) |
| 177 | |
| 178 | if unbatched: |
| 179 | prompt_token_ids = [prompt_token_ids] |
| 180 | if query_token_ids: |
| 181 | query_token_ids = [query_token_ids] |
| 182 | |
| 183 | futures = [] |
| 184 | for i, p in enumerate(prompt_token_ids): |
| 185 | if isinstance(sampling_params, list): |
| 186 | sp = sampling_params[i] |
| 187 | else: |
| 188 | sp = sampling_params |
| 189 | if isinstance(lora_request, list): |
| 190 | lora_req = lora_request[i] |
| 191 | else: |
| 192 | lora_req = lora_request |
| 193 | if isinstance(prompt_adapter_request, list): |
| 194 | pa_req = prompt_adapter_request[i] |
| 195 | else: |
| 196 | pa_req = prompt_adapter_request |
| 197 | future = self.generate_async( |
| 198 | p, |
| 199 | sampling_params=sp, |
| 200 | query_token_ids=query_token_ids, |
| 201 | lora_request=lora_req, |
| 202 | prompt_adapter_request=pa_req, |
| 203 | streaming=False, |
| 204 | disaggregated_params=disaggregated_params) |
| 205 | futures.append(future) |
| 206 | |
| 207 | for future in futures: |
| 208 | future.result() |
| 209 | |
| 210 | if unbatched: |
| 211 | futures = futures[0] |
| 212 | |
| 213 | return futures |
| 214 | |
| 215 | def _get_next_client_id(self): |
| 216 | # (self._last_client_id + 1) % UINT64_MAX |
no test coverage detected