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