Init beam search. Args: **kwargs: Additional keyword arguments.
(
self,
**kwargs,
)
| 398 | return loss_ctc, cer_ctc |
| 399 | |
| 400 | def init_beam_search( |
| 401 | self, |
| 402 | **kwargs, |
| 403 | ): |
| 404 | """Init beam search. |
| 405 | |
| 406 | Args: |
| 407 | **kwargs: Additional keyword arguments. |
| 408 | """ |
| 409 | from funasr.models.transformer.search import BeamSearch |
| 410 | from funasr.models.transformer.scorers.ctc import CTCPrefixScorer |
| 411 | from funasr.models.transformer.scorers.length_bonus import LengthBonus |
| 412 | |
| 413 | # 1. Build ASR model |
| 414 | scorers = {} |
| 415 | |
| 416 | if self.ctc != None: |
| 417 | ctc = CTCPrefixScorer(ctc=self.ctc, eos=self.eos) |
| 418 | scorers.update(ctc=ctc) |
| 419 | token_list = kwargs.get("token_list") |
| 420 | scorers.update( |
| 421 | decoder=self.decoder, |
| 422 | length_bonus=LengthBonus(len(token_list)), |
| 423 | ) |
| 424 | |
| 425 | # 3. Build ngram model |
| 426 | # ngram is not supported now |
| 427 | ngram = None |
| 428 | scorers["ngram"] = ngram |
| 429 | |
| 430 | weights = dict( |
| 431 | decoder=1.0 - kwargs.get("decoding_ctc_weight", 0.3), |
| 432 | ctc=kwargs.get("decoding_ctc_weight", 0.3), |
| 433 | lm=kwargs.get("lm_weight", 0.0), |
| 434 | ngram=kwargs.get("ngram_weight", 0.0), |
| 435 | length_bonus=kwargs.get("penalty", 0.0), |
| 436 | ) |
| 437 | beam_search = BeamSearch( |
| 438 | beam_size=kwargs.get("beam_size", 20), |
| 439 | weights=weights, |
| 440 | scorers=scorers, |
| 441 | sos=self.sos, |
| 442 | eos=self.eos, |
| 443 | vocab_size=len(token_list), |
| 444 | token_list=token_list, |
| 445 | pre_beam_score_key=None if self.ctc_weight == 1.0 else "full", |
| 446 | ) |
| 447 | |
| 448 | self.beam_search = beam_search |
| 449 | |
| 450 | def inference( |
| 451 | self, |
no test coverage detected