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

Method execute_inference

demo/HuggingFace/T5/trt.py:505–611  ·  view source on GitHub ↗
(
        self,
        metadata: NetworkMetadata,
        onnx_fpaths: Dict[str, NetworkModel],
        inference_input: str,
        timing_profile: TimingProfile,
        batch_size: int = 1,
        num_beams: int = 1,
        benchmarking_mode: bool = False,
        benchmarking_args: T5TRTBenchmarkingArgs = None,
    )

Source from the content-addressed store, hash-verified

503 return decoder_output
504
505 def execute_inference(
506 self,
507 metadata: NetworkMetadata,
508 onnx_fpaths: Dict[str, NetworkModel],
509 inference_input: str,
510 timing_profile: TimingProfile,
511 batch_size: int = 1,
512 num_beams: int = 1,
513 benchmarking_mode: bool = False,
514 benchmarking_args: T5TRTBenchmarkingArgs = None,
515 ) -> Union[NetworkResult, BenchmarkingResult]:
516
517 tokenizer = T5Tokenizer.from_pretrained(metadata.variant)
518 hf_config = self.t5_trt_decoder.config
519 # Prepare the input tokens and find out output sequence length.
520 if not benchmarking_mode:
521 output_seq_len = T5ModelTRTConfig.MAX_OUTPUT_LENGTH[metadata.variant]
522 input_ids = tokenizer([inference_input] * batch_size, padding=True, return_tensors="pt").input_ids
523 else:
524 input_seq_len = benchmarking_args.input_seq_len
525 output_seq_len = benchmarking_args.output_seq_len
526
527 input_ids = torch.randint(0, hf_config.vocab_size, (batch_size, input_seq_len))
528
529 encoder_last_hidden_state, encoder_e2e_time = encoder_inference(
530 self.t5_trt_encoder, input_ids, timing_profile
531 )
532
533 # Need to feed the decoder a new empty input_ids for text generation.
534 decoder_output_len = output_seq_len // 2 if (not metadata.other.kv_cache) else 1
535
536 decoder_input_ids = torch.full(
537 (batch_size, decoder_output_len), tokenizer.convert_tokens_to_ids(tokenizer.pad_token), dtype=torch.int32
538 )
539
540 _, decoder_e2e_time = decoder_inference(
541 self.t5_trt_decoder,
542 expand_inputs_for_beam_search(decoder_input_ids, num_beams) if num_beams > 1 else decoder_input_ids,
543 expand_inputs_for_beam_search(encoder_last_hidden_state, num_beams) if num_beams > 1 else encoder_last_hidden_state,
544 timing_profile,
545 use_cache=metadata.other.kv_cache,
546 )
547
548 self.t5_trt_decoder.reset()
549
550 decoder_output, full_e2e_runtime = full_inference(
551 self.t5_trt_encoder,
552 self.t5_trt_decoder,
553 input_ids,
554 tokenizer,
555 timing_profile,
556 max_length=output_seq_len,
557 min_length=T5ModelTRTConfig.MIN_OUTPUT_LENGTH[metadata.variant] if not benchmarking_mode else output_seq_len,
558 batch_size=batch_size,
559 use_cache=metadata.other.kv_cache,
560 num_beams = num_beams,
561 )
562

Callers 1

run_trtMethod · 0.95

Calls 8

encoder_inferenceFunction · 0.90
decoder_inferenceFunction · 0.90
full_inferenceFunction · 0.90
convert_tokens_to_idsMethod · 0.45
resetMethod · 0.45
valuesMethod · 0.45
decodeMethod · 0.45

Tested by

no test coverage detected