FireRedTeam/FireRedASR
13
1#!/usr/bin/env python32 3import argparse4import glob5import os6import sys7 8from fireredasr.models.fireredasr import FireRedAsr9 10 11parser = argparse.ArgumentParser()12parser.add_argument('--asr_type', type=str, required=True, choices=["aed", "llm"])13parser.add_argument('--model_dir', type=str, required=True)14 15# Input / Output16parser.add_argument("--wav_path", type=str)17parser.add_argument("--wav_paths", type=str, nargs="*")18parser.add_argument("--wav_dir", type=str)19parser.add_argument("--wav_scp", type=str)20parser.add_argument("--output", type=str)21 22# Decode Options23parser.add_argument('--use_gpu', type=int, default=1)24parser.add_argument("--batch_size", type=int, default=1)25parser.add_argument("--beam_size", type=int, default=1)26parser.add_argument("--decode_max_len", type=int, default=0)27# FireRedASR-AED28parser.add_argument("--nbest", type=int, default=1)29parser.add_argument("--softmax_smoothing", type=float, default=1.0)30parser.add_argument("--aed_length_penalty", type=float, default=0.0)31parser.add_argument("--eos_penalty", type=float, default=1.0)32# FireRedASR-LLM33parser.add_argument("--decode_min_len", type=int, default=0)34parser.add_argument("--repetition_penalty", type=float, default=1.0)35parser.add_argument("--llm_length_penalty", type=float, default=0.0)36parser.add_argument("--temperature", type=float, default=1.0)37 38 39def main(args):40 wavs = get_wav_info(args)41 fout = open(args.output, "w") if args.output else None42 43 model = FireRedAsr.from_pretrained(args.asr_type, args.model_dir)44 45 batch_uttid = []46 batch_wav_path = []47 for i, wav in enumerate(wavs):48 uttid, wav_path = wav49 batch_uttid.append(uttid)50 batch_wav_path.append(wav_path)51 if len(batch_wav_path) < args.batch_size and i != len(wavs) - 1:52 continue53 54 results = model.transcribe(55 batch_uttid,56 batch_wav_path,57 {58 "use_gpu": args.use_gpu,59 "beam_size": args.beam_size,60 "nbest": args.nbest,61 "decode_max_len": args.decode_max_len,62 "softmax_smoothing": args.softmax_smoothing,63 "aed_length_penalty": args.aed_length_penalty,64 "eos_penalty": args.eos_penalty,65 "decode_min_len": args.decode_min_len,66 "repetition_penalty": args.repetition_penalty,67 "llm_length_penalty": args.llm_length_penalty,68 "temperature": args.temperature69 }70 )71 72 for result in results:73 print(result)74 if fout is not None:75 fout.write(f"{result['uttid']}\t{result['text']}\n")76 77 batch_uttid = []78 batch_wav_path = []79 80 81def get_wav_info(args):82 """83 Returns:84 wavs: list of (uttid, wav_path)85 """86 base = lambda p: os.path.basename(p).replace(".wav", "")87 if args.wav_path:88 wavs = [(base(args.wav_path), args.wav_path)]89 elif args.wav_paths and len(args.wav_paths) >= 1:90 wavs = [(base(p), p) for p in sorted(args.wav_paths)]91 elif args.wav_scp:92 wavs = [line.strip().split() for line in open(args.wav_scp)]93 elif args.wav_dir:94 wavs = glob.glob(f"{args.wav_dir}/**/*.wav", recursive=True)95 wavs = [(base(p), p) for p in sorted(wavs)]96 else:97 raise ValueError("Please provide valid wav info")98 print(f"#wavs={len(wavs)}")99 return wavs100 101 102if __name__ == "__main__":103 args = parser.parse_args()104 print(args)105 main(args)106 