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