MCPcopy Create free account
hub / github.com/NVIDIA/TensorRT / execute_inference

Method execute_inference

demo/HuggingFace/BART/frameworks.py:200–300  ·  view source on GitHub ↗
(
        self,
        metadata: NetworkMetadata,
        network_fpaths: NetworkModels,
        inference_input: str,
        timing_profile: TimingProfile,
        use_cpu: bool,
        batch_size: int = 1,
        num_beams: int = 1,
        benchmarking_mode: bool = False,
        benchmarking_args: BARTBenchmarkingArgs = None,
    )

Source from the content-addressed store, hash-verified

198 return tokenizer, BART_torch_encoder, BART_torch_decoder
199
200 def execute_inference(
201 self,
202 metadata: NetworkMetadata,
203 network_fpaths: NetworkModels,
204 inference_input: str,
205 timing_profile: TimingProfile,
206 use_cpu: bool,
207 batch_size: int = 1,
208 num_beams: int = 1,
209 benchmarking_mode: bool = False,
210 benchmarking_args: BARTBenchmarkingArgs = None,
211 ) -> Union[NetworkResult, BenchmarkingResult]:
212
213 tokenizer, BART_torch_encoder, BART_torch_decoder = self.setup_tokenizer_and_model(metadata, network_fpaths)
214
215 # Prepare the input tokens and find output sequence length.
216 if not benchmarking_mode:
217 output_seq_len = BARTModelTRTConfig.MAX_OUTPUT_LENGTH[metadata.variant]
218 input_ids = tokenizer([inference_input] * batch_size, padding=True, return_tensors="pt").input_ids
219 else:
220 max_seq_len = BARTModelTRTConfig.MAX_SEQUENCE_LENGTH[metadata.variant]
221 input_seq_len = benchmarking_args.input_seq_len if benchmarking_args.input_seq_len > 0 else max_seq_len
222 output_seq_len = benchmarking_args.output_seq_len if benchmarking_args.output_seq_len > 0 else max_seq_len
223 input_ids = torch.randint(0, BARTModelTRTConfig.VOCAB_SIZE[metadata.variant], (batch_size, input_seq_len))
224
225 encoder_last_hidden_state, encoder_e2e_time = encoder_inference(
226 BART_torch_encoder, input_ids, timing_profile, use_cuda=(not use_cpu)
227 )
228
229 # Need to feed the decoder a new empty input_ids for text generation.
230 decoder_output_len = output_seq_len // 2 if (not metadata.other.kv_cache) else 1
231 decoder_input_ids = torch.full(
232 (batch_size, decoder_output_len), tokenizer.convert_tokens_to_ids(tokenizer.pad_token), dtype=torch.int32
233 )
234
235 _, decoder_e2e_time = decoder_inference(
236 BART_torch_decoder, decoder_input_ids, encoder_last_hidden_state, timing_profile, use_cuda=(not use_cpu), use_cache=metadata.other.kv_cache
237 )
238
239 if num_beams == 1:
240 decoder_output, full_e2e_runtime = full_inference_greedy(
241 BART_torch_encoder,
242 BART_torch_decoder,
243 input_ids,
244 tokenizer,
245 timing_profile,
246 max_length=output_seq_len,
247 min_length=BARTModelTRTConfig.MIN_OUTPUT_LENGTH[metadata.variant] if not benchmarking_mode else output_seq_len,
248 use_cuda=(not use_cpu),
249 batch_size=batch_size,
250 use_cache=metadata.other.kv_cache,
251 )
252 else:
253 decoder_output, full_e2e_runtime = full_inference_beam(
254 BART_torch_encoder,
255 BART_torch_decoder,
256 input_ids,
257 tokenizer,

Callers 1

run_frameworkMethod · 0.95

Calls 7

encoder_inferenceFunction · 0.90
decoder_inferenceFunction · 0.90
full_inference_greedyFunction · 0.90
full_inference_beamFunction · 0.90
convert_tokens_to_idsMethod · 0.45
decodeMethod · 0.45

Tested by

no test coverage detected