CoolFace
Apppublic

FireRedTeam/FireRedASR

sourceHugging Faceapache-2.0updated 1y agoView on Hugging Face
13likes
speech2text.py106 linesDownload Raw Back to fireredasr
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