Run inference on input data. Args: data_in: Input data (audio samples, file paths, or text). data_lengths: Lengths of each input sample in the batch. key: Sample identifiers. tokenizer: Tokenizer instance for text e
(
self,
data_in,
data_lengths=None,
key: list = None,
tokenizer=None,
frontend=None,
**kwargs,
)
| 512 | self.beam_search = beam_search |
| 513 | |
| 514 | def inference( |
| 515 | self, |
| 516 | data_in, |
| 517 | data_lengths=None, |
| 518 | key: list = None, |
| 519 | tokenizer=None, |
| 520 | frontend=None, |
| 521 | **kwargs, |
| 522 | ): |
| 523 | |
| 524 | """Run inference on input data. |
| 525 | |
| 526 | Args: |
| 527 | data_in: Input data (audio samples, file paths, or text). |
| 528 | data_lengths: Lengths of each input sample in the batch. |
| 529 | key: Sample identifiers. |
| 530 | tokenizer: Tokenizer instance for text encoding/decoding. |
| 531 | frontend: Audio frontend for feature extraction. |
| 532 | **kwargs: Additional keyword arguments. |
| 533 | """ |
| 534 | if kwargs.get("batch_size", 1) > 1: |
| 535 | return self.inference_batch_ctc( |
| 536 | data_in, data_lengths=data_lengths, key=key, |
| 537 | tokenizer=tokenizer, frontend=frontend, **kwargs |
| 538 | ) |
| 539 | |
| 540 | # init beamsearch |
| 541 | if self.beam_search is None: |
| 542 | logging.info("enable beam_search") |
| 543 | self.init_beam_search(**kwargs) |
| 544 | self.nbest = kwargs.get("nbest", 1) |
| 545 | |
| 546 | meta_data = {} |
| 547 | if ( |
| 548 | isinstance(data_in, torch.Tensor) and kwargs.get("data_type", "sound") == "fbank" |
| 549 | ): # fbank |
| 550 | speech, speech_lengths = data_in, data_lengths |
| 551 | if len(speech.shape) < 3: |
| 552 | speech = speech[None, :, :] |
| 553 | if speech_lengths is None: |
| 554 | speech_lengths = speech.shape[1] |
| 555 | else: |
| 556 | # extract fbank feats |
| 557 | time1 = time.perf_counter() |
| 558 | audio_sample_list = load_audio_text_image_video( |
| 559 | data_in, |
| 560 | fs=frontend.fs, |
| 561 | audio_fs=kwargs.get("fs", 16000), |
| 562 | data_type=kwargs.get("data_type", "sound"), |
| 563 | tokenizer=tokenizer, |
| 564 | ) |
| 565 | time2 = time.perf_counter() |
| 566 | meta_data["load_data"] = f"{time2 - time1:0.3f}" |
| 567 | speech, speech_lengths = extract_fbank( |
| 568 | audio_sample_list, data_type=kwargs.get("data_type", "sound"), frontend=frontend |
| 569 | ) |
| 570 | time3 = time.perf_counter() |
| 571 | meta_data["extract_feat"] = f"{time3 - time2:0.3f}" |
nothing calls this directly
no test coverage detected