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,
)
| 532 | self.beam_search = beam_search |
| 533 | |
| 534 | def inference( |
| 535 | self, |
| 536 | data_in, |
| 537 | data_lengths=None, |
| 538 | key: list = None, |
| 539 | tokenizer=None, |
| 540 | frontend=None, |
| 541 | **kwargs, |
| 542 | ): |
| 543 | # init beamsearch |
| 544 | """Run inference on input data. |
| 545 | |
| 546 | Args: |
| 547 | data_in: Input data (audio samples, file paths, or text). |
| 548 | data_lengths: Lengths of each input sample in the batch. |
| 549 | key: Sample identifiers. |
| 550 | tokenizer: Tokenizer instance for text encoding/decoding. |
| 551 | frontend: Audio frontend for feature extraction. |
| 552 | **kwargs: Additional keyword arguments. |
| 553 | """ |
| 554 | is_use_ctc = kwargs.get("decoding_ctc_weight", 0.0) > 0.00001 and self.ctc != None |
| 555 | is_use_lm = ( |
| 556 | kwargs.get("lm_weight", 0.0) > 0.00001 and kwargs.get("lm_file", None) is not None |
| 557 | ) |
| 558 | pred_timestamp = kwargs.get("pred_timestamp", False) |
| 559 | if self.beam_search is None and (is_use_lm or is_use_ctc): |
| 560 | logging.info("enable beam_search") |
| 561 | self.init_beam_search(**kwargs) |
| 562 | self.nbest = kwargs.get("nbest", 1) |
| 563 | |
| 564 | meta_data = {} |
| 565 | if ( |
| 566 | isinstance(data_in, torch.Tensor) and kwargs.get("data_type", "sound") == "fbank" |
| 567 | ): # fbank |
| 568 | speech, speech_lengths = data_in, data_lengths |
| 569 | if len(speech.shape) < 3: |
| 570 | speech = speech[None, :, :] |
| 571 | if speech_lengths is not None: |
| 572 | speech_lengths = speech_lengths.squeeze(-1) |
| 573 | else: |
| 574 | speech_lengths = speech.shape[1] |
| 575 | else: |
| 576 | # extract fbank feats |
| 577 | time1 = time.perf_counter() |
| 578 | audio_sample_list = load_audio_text_image_video( |
| 579 | data_in, |
| 580 | fs=frontend.fs, |
| 581 | audio_fs=kwargs.get("fs", 16000), |
| 582 | data_type=kwargs.get("data_type", "sound"), |
| 583 | tokenizer=tokenizer, |
| 584 | ) |
| 585 | time2 = time.perf_counter() |
| 586 | meta_data["load_data"] = f"{time2 - time1:0.3f}" |
| 587 | speech, speech_lengths = extract_fbank( |
| 588 | audio_sample_list, data_type=kwargs.get("data_type", "sound"), frontend=frontend |
| 589 | ) |
| 590 | time3 = time.perf_counter() |
| 591 | meta_data["extract_feat"] = f"{time3 - time2:0.3f}" |
nothing calls this directly
no test coverage detected