MCPcopy Create free account
hub / github.com/NVIDIA/TensorRT-LLM / generate

Method generate

tensorrt_llm/executor/executor.py:163–213  ·  view source on GitHub ↗

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,
    )

Source from the content-addressed store, hash-verified

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

Callers 1

compute_scoreMethod · 0.45

Calls 3

generate_asyncMethod · 0.95
appendMethod · 0.45
resultMethod · 0.45

Tested by

no test coverage detected