MCPcopy Create free account
hub / github.com/FireRedTeam/FireRedASR / FireRedAsr

Class FireRedAsr

fireredasr/models/fireredasr.py:13–105  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

11
12
13class FireRedAsr:
14 @classmethod
15 def from_pretrained(cls, asr_type, model_dir):
16 assert asr_type in ["aed", "llm"]
17
18 cmvn_path = os.path.join(model_dir, "cmvn.ark")
19 feat_extractor = ASRFeatExtractor(cmvn_path)
20
21 if asr_type == "aed":
22 model_path = os.path.join(model_dir, "model.pth.tar")
23 dict_path =os.path.join(model_dir, "dict.txt")
24 spm_model = os.path.join(model_dir, "train_bpe1000.model")
25 model = load_fireredasr_aed_model(model_path)
26 tokenizer = ChineseCharEnglishSpmTokenizer(dict_path, spm_model)
27 elif asr_type == "llm":
28 model_path = os.path.join(model_dir, "model.pth.tar")
29 encoder_path = os.path.join(model_dir, "asr_encoder.pth.tar")
30 llm_dir = os.path.join(model_dir, "Qwen2-7B-Instruct")
31 model, tokenizer = load_firered_llm_model_and_tokenizer(
32 model_path, encoder_path, llm_dir)
33 model.eval()
34 return cls(asr_type, feat_extractor, model, tokenizer)
35
36 def __init__(self, asr_type, feat_extractor, model, tokenizer):
37 self.asr_type = asr_type
38 self.feat_extractor = feat_extractor
39 self.model = model
40 self.tokenizer = tokenizer
41
42 @torch.no_grad()
43 def transcribe(self, batch_uttid, batch_wav_path, args={}):
44 feats, lengths, durs = self.feat_extractor(batch_wav_path)
45 total_dur = sum(durs)
46 if args.get("use_gpu", False):
47 feats, lengths = feats.cuda(), lengths.cuda()
48 self.model.cuda()
49 else:
50 self.model.cpu()
51
52 if self.asr_type == "aed":
53 start_time = time.time()
54
55 hyps = self.model.transcribe(
56 feats, lengths,
57 args.get("beam_size", 1),
58 args.get("nbest", 1),
59 args.get("decode_max_len", 0),
60 args.get("softmax_smoothing", 1.0),
61 args.get("aed_length_penalty", 0.0),
62 args.get("eos_penalty", 1.0)
63 )
64
65 elapsed = time.time() - start_time
66 rtf= elapsed / total_dur if total_dur > 0 else 0
67
68 results = []
69 for uttid, wav, hyp in zip(batch_uttid, batch_wav_path, hyps):
70 hyp = hyp[0] # only return 1-best

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected