| 1 | import math |
| 2 | import os |
| 3 | import random |
| 4 | import string |
| 5 | from pathlib import Path |
| 6 | |
| 7 | import torch |
| 8 | import torch.nn.functional as F |
| 9 | import torchaudio |
| 10 | from tqdm import tqdm |
| 11 | |
| 12 | from f5_tts.eval.ecapa_tdnn import ECAPA_TDNN_SMALL |
| 13 | from f5_tts.model.modules import MelSpec |
| 14 | from f5_tts.model.utils import convert_char_to_pinyin |
| 15 | |
| 16 | |
| 17 | # seedtts testset metainfo: utt, prompt_text, prompt_wav, gt_text, gt_wav |
| 18 | def get_seedtts_testset_metainfo(metalst): |
| 19 | f = open(metalst) |
| 20 | lines = f.readlines() |
| 21 | f.close() |
| 22 | metainfo = [] |
| 23 | for line in lines: |
| 24 | if len(line.strip().split("|")) == 5: |
| 25 | utt, prompt_text, prompt_wav, gt_text, gt_wav = line.strip().split("|") |
| 26 | elif len(line.strip().split("|")) == 4: |
| 27 | utt, prompt_text, prompt_wav, gt_text = line.strip().split("|") |
| 28 | gt_wav = os.path.join(os.path.dirname(metalst), "wavs", utt + ".wav") |
| 29 | if not os.path.isabs(prompt_wav): |
| 30 | prompt_wav = os.path.join(os.path.dirname(metalst), prompt_wav) |
| 31 | metainfo.append((utt, prompt_text, prompt_wav, gt_text, gt_wav)) |
| 32 | return metainfo |
| 33 | |
| 34 | |
| 35 | # librispeech test-clean metainfo: gen_utt, ref_txt, ref_wav, gen_txt, gen_wav |
| 36 | def get_librispeech_test_clean_metainfo(metalst, librispeech_test_clean_path): |
| 37 | f = open(metalst) |
| 38 | lines = f.readlines() |
| 39 | f.close() |
| 40 | metainfo = [] |
| 41 | for line in lines: |
| 42 | ref_utt, ref_dur, ref_txt, gen_utt, gen_dur, gen_txt = line.strip().split("\t") |
| 43 | |
| 44 | # ref_txt = ref_txt[0] + ref_txt[1:].lower() + '.' # if use librispeech test-clean (no-pc) |
| 45 | ref_spk_id, ref_chaptr_id, _ = ref_utt.split("-") |
| 46 | ref_wav = os.path.join(librispeech_test_clean_path, ref_spk_id, ref_chaptr_id, ref_utt + ".flac") |
| 47 | |
| 48 | # gen_txt = gen_txt[0] + gen_txt[1:].lower() + '.' # if use librispeech test-clean (no-pc) |
| 49 | gen_spk_id, gen_chaptr_id, _ = gen_utt.split("-") |
| 50 | gen_wav = os.path.join(librispeech_test_clean_path, gen_spk_id, gen_chaptr_id, gen_utt + ".flac") |
| 51 | |
| 52 | metainfo.append((gen_utt, ref_txt, ref_wav, " " + gen_txt, gen_wav)) |
| 53 | |
| 54 | return metainfo |
| 55 | |
| 56 | |
| 57 | # padded to max length mel batch |
| 58 | def padded_mel_batch(ref_mels): |
| 59 | max_mel_length = torch.LongTensor([mel.shape[-1] for mel in ref_mels]).amax() |
| 60 | padded_ref_mels = [] |
| 61 | for mel in ref_mels: |
| 62 | padded_ref_mel = F.pad(mel, (0, max_mel_length - mel.shape[-1]), value=0) |
| 63 | padded_ref_mels.append(padded_ref_mel) |
| 64 | padded_ref_mels = torch.stack(padded_ref_mels) |
| 65 | padded_ref_mels = padded_ref_mels.permute(0, 2, 1) |
| 66 | return padded_ref_mels |
| 67 | |
| 68 | |
| 69 | # get prompts from metainfo containing: utt, prompt_text, prompt_wav, gt_text, gt_wav |
| 70 | |
| 71 | |
| 72 | def get_inference_prompt( |
| 73 | metainfo, |
| 74 | speed=1.0, |
| 75 | tokenizer="pinyin", |
| 76 | polyphone=True, |
| 77 | target_sample_rate=24000, |
| 78 | n_fft=1024, |
| 79 | win_length=1024, |
| 80 | n_mel_channels=100, |
| 81 | hop_length=256, |
| 82 | mel_spec_type="vocos", |
| 83 | target_rms=0.1, |
| 84 | use_truth_duration=False, |
| 85 | infer_batch_size=1, |
| 86 | num_buckets=200, |
| 87 | min_secs=3, |
| 88 | max_secs=40, |
| 89 | ): |
| 90 | prompts_all = [] |
| 91 | |
| 92 | min_tokens = min_secs * target_sample_rate // hop_length |
| 93 | max_tokens = max_secs * target_sample_rate // hop_length |
| 94 | |
| 95 | batch_accum = [0] * num_buckets |
| 96 | utts, ref_rms_list, ref_mels, ref_mel_lens, total_mel_lens, final_text_list = ( |
| 97 | [[] for _ in range(num_buckets)] for _ in range(6) |
| 98 | ) |
| 99 | |
| 100 | mel_spectrogram = MelSpec( |
| 101 | n_fft=n_fft, |
| 102 | hop_length=hop_length, |
| 103 | win_length=win_length, |
| 104 | n_mel_channels=n_mel_channels, |
| 105 | target_sample_rate=target_sample_rate, |
| 106 | mel_spec_type=mel_spec_type, |
| 107 | ) |
| 108 | |
| 109 | for utt, prompt_text, prompt_wav, gt_text, gt_wav in tqdm(metainfo, desc="Processing prompts..."): |
| 110 | # Audio |
| 111 | ref_audio, ref_sr = torchaudio.load(prompt_wav) |
| 112 | ref_rms = torch.sqrt(torch.mean(torch.square(ref_audio))) |
| 113 | if ref_rms < target_rms: |
| 114 | ref_audio = ref_audio * target_rms / ref_rms |
| 115 | assert ref_audio.shape[-1] > 5000, f"Empty prompt wav: {prompt_wav}, or torchaudio backend issue." |
| 116 | if ref_sr != target_sample_rate: |
| 117 | resampler = torchaudio.transforms.Resample(ref_sr, target_sample_rate) |
| 118 | ref_audio = resampler(ref_audio) |
| 119 | |
| 120 | # Text |
| 121 | if len(prompt_text[-1].encode("utf-8")) == 1: |
| 122 | prompt_text = prompt_text + " " |
| 123 | text = [prompt_text + gt_text] |
| 124 | if tokenizer == "pinyin": |
| 125 | text_list = convert_char_to_pinyin(text, polyphone=polyphone) |
| 126 | else: |
| 127 | text_list = text |
| 128 | |
| 129 | # to mel spectrogram |
| 130 | ref_mel = mel_spectrogram(ref_audio) |
| 131 | ref_mel = ref_mel.squeeze(0) |
| 132 | |
| 133 | # Duration, mel frame length |
| 134 | ref_mel_len = ref_mel.shape[-1] |
| 135 | |
| 136 | if use_truth_duration: |
| 137 | gt_audio, gt_sr = torchaudio.load(gt_wav) |
| 138 | if gt_sr != target_sample_rate: |
| 139 | resampler = torchaudio.transforms.Resample(gt_sr, target_sample_rate) |
| 140 | gt_audio = resampler(gt_audio) |
| 141 | total_mel_len = ref_mel_len + int(gt_audio.shape[-1] / hop_length / speed) |
| 142 | |
| 143 | # # test vocoder resynthesis |
| 144 | # ref_audio = gt_audio |
| 145 | else: |
| 146 | ref_text_len = len(prompt_text.encode("utf-8")) |
| 147 | gen_text_len = len(gt_text.encode("utf-8")) |
| 148 | total_mel_len = ref_mel_len + int(ref_mel_len / ref_text_len * gen_text_len / speed) |
| 149 | |
| 150 | # deal with batch |
| 151 | assert infer_batch_size > 0, "infer_batch_size should be greater than 0." |
| 152 | assert min_tokens <= total_mel_len <= max_tokens, ( |
| 153 | f"Audio {utt} has duration {total_mel_len * hop_length // target_sample_rate}s out of range [{min_secs}, {max_secs}]." |
| 154 | ) |
| 155 | bucket_i = math.floor((total_mel_len - min_tokens) / (max_tokens - min_tokens + 1) * num_buckets) |
| 156 | |
| 157 | utts[bucket_i].append(utt) |
| 158 | ref_rms_list[bucket_i].append(ref_rms) |
| 159 | ref_mels[bucket_i].append(ref_mel) |
| 160 | ref_mel_lens[bucket_i].append(ref_mel_len) |
| 161 | total_mel_lens[bucket_i].append(total_mel_len) |
| 162 | final_text_list[bucket_i].extend(text_list) |
| 163 | |
| 164 | batch_accum[bucket_i] += total_mel_len |
| 165 | |
| 166 | if batch_accum[bucket_i] >= infer_batch_size: |
| 167 | # print(f"\n{len(ref_mels[bucket_i][0][0])}\n{ref_mel_lens[bucket_i]}\n{total_mel_lens[bucket_i]}") |
| 168 | prompts_all.append( |
| 169 | ( |
| 170 | utts[bucket_i], |
| 171 | ref_rms_list[bucket_i], |
| 172 | padded_mel_batch(ref_mels[bucket_i]), |
| 173 | ref_mel_lens[bucket_i], |
| 174 | total_mel_lens[bucket_i], |
| 175 | final_text_list[bucket_i], |
| 176 | ) |
| 177 | ) |
| 178 | batch_accum[bucket_i] = 0 |
| 179 | ( |
| 180 | utts[bucket_i], |
| 181 | ref_rms_list[bucket_i], |
| 182 | ref_mels[bucket_i], |
| 183 | ref_mel_lens[bucket_i], |
| 184 | total_mel_lens[bucket_i], |
| 185 | final_text_list[bucket_i], |
| 186 | ) = [], [], [], [], [], [] |
| 187 | |
| 188 | # add residual |
| 189 | for bucket_i, bucket_frames in enumerate(batch_accum): |
| 190 | if bucket_frames > 0: |
| 191 | prompts_all.append( |
| 192 | ( |
| 193 | utts[bucket_i], |
| 194 | ref_rms_list[bucket_i], |
| 195 | padded_mel_batch(ref_mels[bucket_i]), |
| 196 | ref_mel_lens[bucket_i], |
| 197 | total_mel_lens[bucket_i], |
| 198 | final_text_list[bucket_i], |
| 199 | ) |
| 200 | ) |
| 201 | # not only leave easy work for last workers |
| 202 | random.seed(666) |
| 203 | random.shuffle(prompts_all) |
| 204 | |
| 205 | return prompts_all |
| 206 | |
| 207 | |
| 208 | # get wav_res_ref_text of seed-tts test metalst |
| 209 | # https://github.com/BytedanceSpeech/seed-tts-eval |
| 210 | |
| 211 | |
| 212 | def get_seed_tts_test(metalst, gen_wav_dir, gpus): |
| 213 | f = open(metalst) |
| 214 | lines = f.readlines() |
| 215 | f.close() |
| 216 | |
| 217 | test_set_ = [] |
| 218 | for line in tqdm(lines): |
| 219 | if len(line.strip().split("|")) == 5: |
| 220 | utt, prompt_text, prompt_wav, gt_text, gt_wav = line.strip().split("|") |
| 221 | elif len(line.strip().split("|")) == 4: |
| 222 | utt, prompt_text, prompt_wav, gt_text = line.strip().split("|") |
| 223 | |
| 224 | if not os.path.exists(os.path.join(gen_wav_dir, utt + ".wav")): |
| 225 | continue |
| 226 | gen_wav = os.path.join(gen_wav_dir, utt + ".wav") |
| 227 | if not os.path.isabs(prompt_wav): |
| 228 | prompt_wav = os.path.join(os.path.dirname(metalst), prompt_wav) |
| 229 | |
| 230 | test_set_.append((gen_wav, prompt_wav, gt_text)) |
| 231 | |
| 232 | num_jobs = len(gpus) |
| 233 | if num_jobs == 1: |
| 234 | return [(gpus[0], test_set_)] |
| 235 | |
| 236 | wav_per_job = len(test_set_) // num_jobs + 1 |
| 237 | test_set = [] |
| 238 | for i in range(num_jobs): |
| 239 | test_set.append((gpus[i], test_set_[i * wav_per_job : (i + 1) * wav_per_job])) |
| 240 | |
| 241 | return test_set |
| 242 | |
| 243 | |
| 244 | # get librispeech test-clean cross sentence test |
| 245 | |
| 246 | |
| 247 | def get_librispeech_test(metalst, gen_wav_dir, gpus, librispeech_test_clean_path, eval_ground_truth=False): |
| 248 | f = open(metalst) |
| 249 | lines = f.readlines() |
| 250 | f.close() |
| 251 | |
| 252 | test_set_ = [] |
| 253 | for line in tqdm(lines): |
| 254 | ref_utt, ref_dur, ref_txt, gen_utt, gen_dur, gen_txt = line.strip().split("\t") |
| 255 | |
| 256 | if eval_ground_truth: |
| 257 | gen_spk_id, gen_chaptr_id, _ = gen_utt.split("-") |
| 258 | gen_wav = os.path.join(librispeech_test_clean_path, gen_spk_id, gen_chaptr_id, gen_utt + ".flac") |
| 259 | else: |
| 260 | if not os.path.exists(os.path.join(gen_wav_dir, gen_utt + ".wav")): |
| 261 | raise FileNotFoundError(f"Generated wav not found: {gen_utt}") |
| 262 | gen_wav = os.path.join(gen_wav_dir, gen_utt + ".wav") |
| 263 | |
| 264 | ref_spk_id, ref_chaptr_id, _ = ref_utt.split("-") |
| 265 | ref_wav = os.path.join(librispeech_test_clean_path, ref_spk_id, ref_chaptr_id, ref_utt + ".flac") |
| 266 | |
| 267 | test_set_.append((gen_wav, ref_wav, gen_txt)) |
| 268 | |
| 269 | num_jobs = len(gpus) |
| 270 | if num_jobs == 1: |
| 271 | return [(gpus[0], test_set_)] |
| 272 | |
| 273 | wav_per_job = len(test_set_) // num_jobs + 1 |
| 274 | test_set = [] |
| 275 | for i in range(num_jobs): |
| 276 | test_set.append((gpus[i], test_set_[i * wav_per_job : (i + 1) * wav_per_job])) |
| 277 | |
| 278 | return test_set |
| 279 | |
| 280 | |
| 281 | # load asr model |
| 282 | |
| 283 | |
| 284 | def load_asr_model(lang, ckpt_dir=""): |
| 285 | if lang == "zh": |
| 286 | from funasr import AutoModel |
| 287 | |
| 288 | model = AutoModel( |
| 289 | model=os.path.join(ckpt_dir, "paraformer-zh"), |
| 290 | # vad_model = os.path.join(ckpt_dir, "fsmn-vad"), |
| 291 | # punc_model = os.path.join(ckpt_dir, "ct-punc"), |
| 292 | # spk_model = os.path.join(ckpt_dir, "cam++"), |
| 293 | disable_update=True, |
| 294 | ) # following seed-tts setting |
| 295 | elif lang == "en": |
| 296 | from faster_whisper import WhisperModel |
| 297 | |
| 298 | model_size = "large-v3" if ckpt_dir == "" else ckpt_dir |
| 299 | model = WhisperModel(model_size, device="cuda", compute_type="float16") |
| 300 | return model |
| 301 | |
| 302 | |
| 303 | # WER Evaluation, the way Seed-TTS does |
| 304 | |
| 305 | |
| 306 | def run_asr_wer(args): |
| 307 | rank, lang, test_set, ckpt_dir = args |
| 308 | |
| 309 | if lang == "zh": |
| 310 | import zhconv |
| 311 | |
| 312 | torch.cuda.set_device(rank) |
| 313 | elif lang == "en": |
| 314 | os.environ["CUDA_VISIBLE_DEVICES"] = str(rank) |
| 315 | else: |
| 316 | raise NotImplementedError( |
| 317 | "lang support only 'zh' (funasr paraformer-zh), 'en' (faster-whisper-large-v3), for now." |
| 318 | ) |
| 319 | |
| 320 | asr_model = load_asr_model(lang, ckpt_dir=ckpt_dir) |
| 321 | |
| 322 | from zhon.hanzi import punctuation |
| 323 | |
| 324 | punctuation_all = punctuation + string.punctuation |
| 325 | wer_results = [] |
| 326 | |
| 327 | from jiwer import process_words |
| 328 | |
| 329 | for gen_wav, prompt_wav, truth in tqdm(test_set): |
| 330 | if lang == "zh": |
| 331 | res = asr_model.generate(input=gen_wav, batch_size_s=300, disable_pbar=True) |
| 332 | hypo = res[0]["text"] |
| 333 | hypo = zhconv.convert(hypo, "zh-cn") |
| 334 | elif lang == "en": |
| 335 | segments, _ = asr_model.transcribe(gen_wav, beam_size=5, language="en") |
| 336 | hypo = "" |
| 337 | for segment in segments: |
| 338 | hypo = hypo + " " + segment.text |
| 339 | |
| 340 | raw_truth = truth |
| 341 | raw_hypo = hypo |
| 342 | |
| 343 | for x in punctuation_all: |
| 344 | truth = truth.replace(x, "") |
| 345 | hypo = hypo.replace(x, "") |
| 346 | |
| 347 | truth = truth.replace(" ", " ") |
| 348 | hypo = hypo.replace(" ", " ") |
| 349 | |
| 350 | if lang == "zh": |
| 351 | truth = " ".join([x for x in truth]) |
| 352 | hypo = " ".join([x for x in hypo]) |
| 353 | elif lang == "en": |
| 354 | truth = truth.lower() |
| 355 | hypo = hypo.lower() |
| 356 | |
| 357 | measures = process_words(truth, hypo) |
| 358 | wer = measures.wer |
| 359 | |
| 360 | # ref_list = truth.split(" ") |
| 361 | # subs = measures.substitutions / len(ref_list) |
| 362 | # dele = measures.deletions / len(ref_list) |
| 363 | # inse = measures.insertions / len(ref_list) |
| 364 | |
| 365 | wer_results.append( |
| 366 | { |
| 367 | "wav": Path(gen_wav).stem, |
| 368 | "truth": raw_truth, |
| 369 | "hypo": raw_hypo, |
| 370 | "wer": wer, |
| 371 | } |
| 372 | ) |
| 373 | |
| 374 | return wer_results |
| 375 | |
| 376 | |
| 377 | # SIM Evaluation |
| 378 | |
| 379 | |
| 380 | def run_sim(args): |
| 381 | rank, test_set, ckpt_dir = args |
| 382 | device = f"cuda:{rank}" |
| 383 | |
| 384 | model = ECAPA_TDNN_SMALL(feat_dim=1024, feat_type="wavlm_large", config_path=None) |
| 385 | state_dict = torch.load(ckpt_dir, weights_only=True, map_location=lambda storage, loc: storage) |
| 386 | model.load_state_dict(state_dict["model"], strict=False) |
| 387 | |
| 388 | use_gpu = True if torch.cuda.is_available() else False |
| 389 | if use_gpu: |
| 390 | model = model.cuda(device) |
| 391 | model.eval() |
| 392 | |
| 393 | sim_results = [] |
| 394 | for gen_wav, prompt_wav, truth in tqdm(test_set): |
| 395 | wav1, sr1 = torchaudio.load(gen_wav) |
| 396 | wav2, sr2 = torchaudio.load(prompt_wav) |
| 397 | |
| 398 | if use_gpu: |
| 399 | wav1 = wav1.cuda(device) |
| 400 | wav2 = wav2.cuda(device) |
| 401 | |
| 402 | if sr1 != 16000: |
| 403 | resample1 = torchaudio.transforms.Resample(orig_freq=sr1, new_freq=16000) |
| 404 | if use_gpu: |
| 405 | resample1 = resample1.cuda(device) |
| 406 | wav1 = resample1(wav1) |
| 407 | if sr2 != 16000: |
| 408 | resample2 = torchaudio.transforms.Resample(orig_freq=sr2, new_freq=16000) |
| 409 | if use_gpu: |
| 410 | resample2 = resample2.cuda(device) |
| 411 | wav2 = resample2(wav2) |
| 412 | |
| 413 | with torch.no_grad(): |
| 414 | emb1 = model(wav1) |
| 415 | emb2 = model(wav2) |
| 416 | |
| 417 | sim = F.cosine_similarity(emb1, emb2)[0].item() |
| 418 | # print(f"VSim score between two audios: {sim:.4f} (-1.0, 1.0).") |
| 419 | sim_results.append( |
| 420 | { |
| 421 | "wav": Path(gen_wav).stem, |
| 422 | "sim": sim, |
| 423 | } |
| 424 | ) |
| 425 | |
| 426 | return sim_results |
| 427 |