| 1 | # Evaluate with Seed-TTS testset |
| 2 | |
| 3 | import argparse |
| 4 | import ast |
| 5 | import json |
| 6 | import os |
| 7 | import sys |
| 8 | |
| 9 | |
| 10 | sys.path.append(os.getcwd()) |
| 11 | |
| 12 | import multiprocessing as mp |
| 13 | from importlib.resources import files |
| 14 | |
| 15 | import numpy as np |
| 16 | |
| 17 | from f5_tts.eval.utils_eval import get_seed_tts_test, run_asr_wer, run_sim |
| 18 | |
| 19 | |
| 20 | rel_path = str(files("f5_tts").joinpath("../../")) |
| 21 | |
| 22 | |
| 23 | def get_args(): |
| 24 | parser = argparse.ArgumentParser() |
| 25 | parser.add_argument("-e", "--eval_task", type=str, default="wer", choices=["sim", "wer"]) |
| 26 | parser.add_argument("-l", "--lang", type=str, default="en", choices=["zh", "en"]) |
| 27 | parser.add_argument("-g", "--gen_wav_dir", type=str, required=True) |
| 28 | parser.add_argument( |
| 29 | "-n", "--gpu_nums", type=str, default="8", help="Number of GPUs to use (e.g., 8) or GPU list (e.g., [0,1,2,3])" |
| 30 | ) |
| 31 | parser.add_argument("--local", action="store_true", help="Use local custom checkpoint directory") |
| 32 | return parser.parse_args() |
| 33 | |
| 34 | |
| 35 | def parse_gpu_nums(gpu_nums_str): |
| 36 | try: |
| 37 | if gpu_nums_str.startswith("[") and gpu_nums_str.endswith("]"): |
| 38 | gpu_list = ast.literal_eval(gpu_nums_str) |
| 39 | if isinstance(gpu_list, list): |
| 40 | return gpu_list |
| 41 | return list(range(int(gpu_nums_str))) |
| 42 | except (ValueError, SyntaxError): |
| 43 | raise argparse.ArgumentTypeError( |
| 44 | f"Invalid GPU specification: {gpu_nums_str}. Use a number (e.g., 8) or a list (e.g., [0,1,2,3])" |
| 45 | ) |
| 46 | |
| 47 | |
| 48 | def main(): |
| 49 | args = get_args() |
| 50 | eval_task = args.eval_task |
| 51 | lang = args.lang |
| 52 | gen_wav_dir = args.gen_wav_dir |
| 53 | metalst = rel_path + f"/data/seedtts_testset/{lang}/meta.lst" # seed-tts testset |
| 54 | |
| 55 | # NOTE. paraformer-zh result will be slightly different according to the number of gpus, cuz batchsize is different |
| 56 | # zh 1.254 seems a result of 4 workers wer_seed_tts |
| 57 | gpus = parse_gpu_nums(args.gpu_nums) |
| 58 | test_set = get_seed_tts_test(metalst, gen_wav_dir, gpus) |
| 59 | |
| 60 | local = args.local |
| 61 | if local: # use local custom checkpoint dir |
| 62 | if lang == "zh": |
| 63 | asr_ckpt_dir = "../checkpoints/funasr" # paraformer-zh dir under funasr |
| 64 | elif lang == "en": |
| 65 | asr_ckpt_dir = "../checkpoints/Systran/faster-whisper-large-v3" |
| 66 | else: |
| 67 | asr_ckpt_dir = "" # auto download to cache dir |
| 68 | wavlm_ckpt_dir = "../checkpoints/UniSpeech/wavlm_large_finetune.pth" |
| 69 | |
| 70 | # -------------------------------------------------------------------------- |
| 71 | |
| 72 | full_results = [] |
| 73 | metrics = [] |
| 74 | |
| 75 | if eval_task == "wer": |
| 76 | with mp.Pool(processes=len(gpus)) as pool: |
| 77 | args = [(rank, lang, sub_test_set, asr_ckpt_dir) for (rank, sub_test_set) in test_set] |
| 78 | results = pool.map(run_asr_wer, args) |
| 79 | for r in results: |
| 80 | full_results.extend(r) |
| 81 | elif eval_task == "sim": |
| 82 | with mp.Pool(processes=len(gpus)) as pool: |
| 83 | args = [(rank, sub_test_set, wavlm_ckpt_dir) for (rank, sub_test_set) in test_set] |
| 84 | results = pool.map(run_sim, args) |
| 85 | for r in results: |
| 86 | full_results.extend(r) |
| 87 | else: |
| 88 | raise ValueError(f"Unknown metric type: {eval_task}") |
| 89 | |
| 90 | result_path = f"{gen_wav_dir}/_{eval_task}_results.jsonl" |
| 91 | with open(result_path, "w") as f: |
| 92 | for line in full_results: |
| 93 | metrics.append(line[eval_task]) |
| 94 | f.write(json.dumps(line, ensure_ascii=False) + "\n") |
| 95 | metric = round(np.mean(metrics), 5) |
| 96 | f.write(f"\n{eval_task.upper()}: {metric}\n") |
| 97 | |
| 98 | print(f"\nTotal {len(metrics)} samples") |
| 99 | print(f"{eval_task.upper()}: {metric}") |
| 100 | print(f"{eval_task.upper()} results saved to {result_path}") |
| 101 | |
| 102 | |
| 103 | if __name__ == "__main__": |
| 104 | main() |
| 105 |