| 11 | |
| 12 | |
| 13 | class 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 |
nothing calls this directly
no outgoing calls
no test coverage detected