MCPcopy Create free account
hub / github.com/modelscope/FunASR / init_beam_search

Method init_beam_search

funasr/models/lcbnet/model.py:400–448  ·  view source on GitHub ↗

Init beam search. Args: **kwargs: Additional keyword arguments.

(
        self,
        **kwargs,
    )

Source from the content-addressed store, hash-verified

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,

Callers 1

inferenceMethod · 0.95

Calls 4

CTCPrefixScorerClass · 0.90
LengthBonusClass · 0.90
BeamSearchClass · 0.90
updateMethod · 0.45

Tested by

no test coverage detected