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

Method execute_inference

demo/HuggingFace/BART/trt.py:684–802  ·  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: BARTTRTBenchmarkingArgs = None,
    )

Source from the content-addressed store, hash-verified

682 self.BART_trt_decoder.use_non_kv_engine = self.metadata.other.kv_cache
683
684 def execute_inference(
685 self,
686 metadata: NetworkMetadata,
687 onnx_fpaths: Dict[str, NetworkModel],
688 inference_input: str,
689 timing_profile: TimingProfile,
690 batch_size: int = 1,
691 num_beams: int = 1,
692 benchmarking_mode: bool = False,
693 benchmarking_args: BARTTRTBenchmarkingArgs = None,
694 ) -> Union[NetworkResult, BenchmarkingResult]:
695 if "mbart" not in metadata.variant:
696 tokenizer = BartTokenizer.from_pretrained(metadata.variant)
697 else:
698 tokenizer = MBart50Tokenizer.from_pretrained(metadata.variant, src_lang="en_XX")
699
700 # Prepare the input tokens and find output sequence length.
701 if not benchmarking_mode:
702 output_seq_len = BARTModelTRTConfig.MAX_OUTPUT_LENGTH[metadata.variant]
703 input_ids = tokenizer([inference_input] * batch_size, padding=True, return_tensors="pt").input_ids
704 else:
705 input_seq_len = benchmarking_args.input_seq_len
706 output_seq_len = benchmarking_args.output_seq_len
707
708 input_ids = torch.randint(0, BARTModelTRTConfig.VOCAB_SIZE[metadata.variant], (batch_size, input_seq_len))
709
710 encoder_last_hidden_state, encoder_e2e_time = encoder_inference(
711 self.BART_trt_encoder, input_ids, timing_profile
712 )
713
714 # Need to feed the decoder a new empty input_ids for text generation.
715 decoder_output_len = output_seq_len // 2 if (not metadata.other.kv_cache) else 1
716 decoder_input_ids = torch.full(
717 (batch_size, decoder_output_len), tokenizer.convert_tokens_to_ids(tokenizer.pad_token), dtype=torch.int32
718 )
719
720 _, decoder_e2e_time = decoder_inference(
721 self.BART_trt_decoder,
722 expand_inputs_for_beam_search(decoder_input_ids, num_beams) if num_beams > 1 else decoder_input_ids,
723 expand_inputs_for_beam_search(encoder_last_hidden_state, num_beams) if num_beams > 1 else encoder_last_hidden_state,
724 timing_profile,
725 use_cache=metadata.other.kv_cache,
726 )
727
728 if num_beams == 1:
729 decoder_output, full_e2e_runtime = full_inference_greedy(
730 self.BART_trt_encoder,
731 self.BART_trt_decoder,
732 input_ids,
733 tokenizer,
734 timing_profile,
735 max_length=output_seq_len,
736 min_length=BARTModelTRTConfig.MIN_OUTPUT_LENGTH[metadata.variant] if not benchmarking_mode else output_seq_len,
737 batch_size=batch_size,
738 use_cache=metadata.other.kv_cache,
739 )
740 else:
741 decoder_output, full_e2e_runtime = full_inference_beam(

Callers 1

run_trtMethod · 0.95

Calls 8

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

Tested by

no test coverage detected