Init beam search. Args: **kwargs: Additional keyword arguments.
(
self,
**kwargs,
)
| 541 | ) |
| 542 | |
| 543 | def init_beam_search( |
| 544 | self, |
| 545 | **kwargs, |
| 546 | ): |
| 547 | |
| 548 | """Init beam search. |
| 549 | |
| 550 | Args: |
| 551 | **kwargs: Additional keyword arguments. |
| 552 | """ |
| 553 | from funasr.models.scama.beam_search import BeamSearchScamaStreaming |
| 554 | |
| 555 | from funasr.models.transformer.scorers.ctc import CTCPrefixScorer |
| 556 | from funasr.models.transformer.scorers.length_bonus import LengthBonus |
| 557 | |
| 558 | # 1. Build ASR model |
| 559 | scorers = {} |
| 560 | |
| 561 | if self.ctc != None: |
| 562 | ctc = CTCPrefixScorer(ctc=self.ctc, eos=self.eos) |
| 563 | scorers.update(ctc=ctc) |
| 564 | token_list = kwargs.get("token_list") |
| 565 | scorers.update( |
| 566 | decoder=self.decoder, |
| 567 | length_bonus=LengthBonus(len(token_list)), |
| 568 | ) |
| 569 | |
| 570 | # 3. Build ngram model |
| 571 | # ngram is not supported now |
| 572 | ngram = None |
| 573 | scorers["ngram"] = ngram |
| 574 | |
| 575 | weights = dict( |
| 576 | decoder=1.0 - kwargs.get("decoding_ctc_weight", 0.0), |
| 577 | ctc=kwargs.get("decoding_ctc_weight", 0.0), |
| 578 | lm=kwargs.get("lm_weight", 0.0), |
| 579 | ngram=kwargs.get("ngram_weight", 0.0), |
| 580 | length_bonus=kwargs.get("penalty", 0.0), |
| 581 | ) |
| 582 | |
| 583 | beam_search = BeamSearchScamaStreaming( |
| 584 | beam_size=kwargs.get("beam_size", 2), |
| 585 | weights=weights, |
| 586 | scorers=scorers, |
| 587 | sos=self.sos, |
| 588 | eos=self.eos, |
| 589 | vocab_size=len(token_list), |
| 590 | token_list=token_list, |
| 591 | pre_beam_score_key=None if self.ctc_weight == 1.0 else "full", |
| 592 | ) |
| 593 | # beam_search.to(device=kwargs.get("device", "cpu"), dtype=getattr(torch, kwargs.get("dtype", "float32"))).eval() |
| 594 | # for scorer in scorers.values(): |
| 595 | # if isinstance(scorer, torch.nn.Module): |
| 596 | # scorer.to(device=kwargs.get("device", "cpu"), dtype=getattr(torch, kwargs.get("dtype", "float32"))).eval() |
| 597 | self.beam_search = beam_search |
| 598 | |
| 599 | def generate_chunk( |
| 600 | self, |
no test coverage detected