CoolFace
Apppublic

ThreadAbort/E2-F5-TTS

sourceHugging Faceupdated 1y agoView on Hugging Face
26likes
test_infer_batch.py203 linesDownload Raw Back to root
1import os2import time3import random4from tqdm import tqdm5import argparse6 7import torch8import torchaudio9from accelerate import Accelerator10from einops import rearrange11from ema_pytorch import EMA12from vocos import Vocos13 14from model import CFM, UNetT, DiT15from model.utils import (16    get_tokenizer, 17    get_seedtts_testset_metainfo, 18    get_librispeech_test_clean_metainfo, 19    get_inference_prompt,20)21 22accelerator = Accelerator()23device = f"cuda:{accelerator.process_index}"24 25 26# --------------------- Dataset Settings -------------------- #27 28target_sample_rate = 2400029n_mel_channels = 10030hop_length = 25631target_rms = 0.132 33tokenizer = "pinyin"34 35 36# ---------------------- infer setting ---------------------- #37 38parser = argparse.ArgumentParser(description="batch inference")39 40parser.add_argument('-s', '--seed', default=None, type=int)41parser.add_argument('-d', '--dataset', default="Emilia_ZH_EN")42parser.add_argument('-n', '--expname', required=True)43parser.add_argument('-c', '--ckptstep', default=1200000, type=int)44 45parser.add_argument('-nfe', '--nfestep', default=32, type=int)46parser.add_argument('-o', '--odemethod', default="euler")47parser.add_argument('-ss', '--swaysampling', default=-1, type=float)48 49parser.add_argument('-t', '--testset', required=True)50 51args = parser.parse_args()52 53 54seed = args.seed55dataset_name = args.dataset56exp_name = args.expname57ckpt_step = args.ckptstep58checkpoint = torch.load(f"ckpts/{exp_name}/model_{ckpt_step}.pt", map_location=device)59 60nfe_step = args.nfestep61ode_method = args.odemethod62sway_sampling_coef = args.swaysampling63 64testset = args.testset65 66 67infer_batch_size = 1  # max frames. 1 for ddp single inference (recommended)68cfg_strength = 2.69speed = 1.70use_truth_duration = False71no_ref_audio = False72 73 74if exp_name == "F5TTS_Base":75    model_cls = DiT76    model_cfg = dict(dim = 1024, depth = 22, heads = 16, ff_mult = 2, text_dim = 512, conv_layers = 4)77 78elif exp_name == "E2TTS_Base":79    model_cls = UNetT80    model_cfg = dict(dim = 1024, depth = 24, heads = 16, ff_mult = 4)81 82 83if testset == "ls_pc_test_clean":84    metalst = "data/librispeech_pc_test_clean_cross_sentence.lst"85    librispeech_test_clean_path = "<SOME_PATH>/LibriSpeech/test-clean"  # test-clean path86    metainfo = get_librispeech_test_clean_metainfo(metalst, librispeech_test_clean_path)87    88elif testset == "seedtts_test_zh":89    metalst = "data/seedtts_testset/zh/meta.lst"90    metainfo = get_seedtts_testset_metainfo(metalst)91 92elif testset == "seedtts_test_en":93    metalst = "data/seedtts_testset/en/meta.lst"94    metainfo = get_seedtts_testset_metainfo(metalst)95 96 97# path to save genereted wavs98if seed is None: seed = random.randint(-10000, 10000)99output_dir = f"results/{exp_name}_{ckpt_step}/{testset}/" \100    f"seed{seed}_{ode_method}_nfe{nfe_step}" \101    f"{f'_ss{sway_sampling_coef}' if sway_sampling_coef else ''}" \102    f"_cfg{cfg_strength}_speed{speed}" \103    f"{'_gt-dur' if use_truth_duration else ''}" \104    f"{'_no-ref-audio' if no_ref_audio else ''}"105 106 107# -------------------------------------------------#108 109use_ema = True110 111prompts_all = get_inference_prompt(112    metainfo, 113    speed = speed, 114    tokenizer = tokenizer, 115    target_sample_rate = target_sample_rate, 116    n_mel_channels = n_mel_channels,117    hop_length = hop_length,118    target_rms = target_rms,119    use_truth_duration = use_truth_duration,120    infer_batch_size = infer_batch_size,121)122 123# Vocoder model124local = False125if local:126    vocos_local_path = "../checkpoints/charactr/vocos-mel-24khz"127    vocos = Vocos.from_hparams(f"{vocos_local_path}/config.yaml")128    state_dict = torch.load(f"{vocos_local_path}/pytorch_model.bin", map_location=device)129    vocos.load_state_dict(state_dict)130    vocos.eval()131else:132    vocos = Vocos.from_pretrained("charactr/vocos-mel-24khz")133 134# Tokenizer135vocab_char_map, vocab_size = get_tokenizer(dataset_name, tokenizer)136 137# Model138model = CFM(139    transformer = model_cls(140        **model_cfg,141        text_num_embeds = vocab_size, 142        mel_dim = n_mel_channels143    ),144    mel_spec_kwargs = dict(145        target_sample_rate = target_sample_rate, 146        n_mel_channels = n_mel_channels,147        hop_length = hop_length,148    ),149    odeint_kwargs = dict(150        method = ode_method,151    ),152    vocab_char_map = vocab_char_map,153).to(device)154 155if use_ema == True:156    ema_model = EMA(model, include_online_model = False).to(device)157    ema_model.load_state_dict(checkpoint['ema_model_state_dict'])158    ema_model.copy_params_from_ema_to_model()159else:160    model.load_state_dict(checkpoint['model_state_dict'])161 162if not os.path.exists(output_dir) and accelerator.is_main_process:163    os.makedirs(output_dir)164 165# start batch inference166accelerator.wait_for_everyone()167start = time.time()168 169with accelerator.split_between_processes(prompts_all) as prompts:170 171    for prompt in tqdm(prompts, disable=not accelerator.is_local_main_process):172        utts, ref_rms_list, ref_mels, ref_mel_lens, total_mel_lens, final_text_list = prompt173        ref_mels = ref_mels.to(device)174        ref_mel_lens = torch.tensor(ref_mel_lens, dtype = torch.long).to(device)175        total_mel_lens = torch.tensor(total_mel_lens, dtype = torch.long).to(device)176        177        # Inference178        with torch.inference_mode():179            generated, _ = model.sample(180                cond = ref_mels,181                text = final_text_list,182                duration = total_mel_lens,183                lens = ref_mel_lens,184                steps = nfe_step,185                cfg_strength = cfg_strength,186                sway_sampling_coef = sway_sampling_coef,187                no_ref_audio = no_ref_audio,188                seed = seed,189            )190        # Final result191        for i, gen in enumerate(generated):192            gen = gen[ref_mel_lens[i]:total_mel_lens[i], :].unsqueeze(0)193            gen_mel_spec = rearrange(gen, '1 n d -> 1 d n')194            generated_wave = vocos.decode(gen_mel_spec.cpu())195            if ref_rms_list[i] < target_rms:196                generated_wave = generated_wave * ref_rms_list[i] / target_rms197            torchaudio.save(f"{output_dir}/{utts[i]}.wav", generated_wave, target_sample_rate)198 199accelerator.wait_for_everyone()200if accelerator.is_main_process:201    timediff = time.time() - start202    print(f"Done batch inference in {timediff / 60 :.2f} minutes.")203