(
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,
)
| 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( |
no test coverage detected