CoolFace
Apppublic

honey126/VoxAI

sourceHugging Facemitupdated 9mo agoView on Hugging Face
0likes
eval_infer_batch.py199 linesDownload Raw Back to scripts
1import sys2import os3 4sys.path.append(os.getcwd())5 6import time7import random8from tqdm import tqdm9import argparse10 11import torch12import torchaudio13from accelerate import Accelerator14from vocos import Vocos15 16from model import CFM, UNetT, DiT17from model.utils import (18    load_checkpoint,19    get_tokenizer,20    get_seedtts_testset_metainfo,21    get_librispeech_test_clean_metainfo,22    get_inference_prompt,23)24 25accelerator = Accelerator()26device = f"cuda:{accelerator.process_index}"27 28 29# --------------------- Dataset Settings -------------------- #30 31target_sample_rate = 2400032n_mel_channels = 10033hop_length = 25634target_rms = 0.135 36tokenizer = "pinyin"37 38 39# ---------------------- infer setting ---------------------- #40 41parser = argparse.ArgumentParser(description="batch inference")42 43parser.add_argument("-s", "--seed", default=None, type=int)44parser.add_argument("-d", "--dataset", default="Emilia_ZH_EN")45parser.add_argument("-n", "--expname", required=True)46parser.add_argument("-c", "--ckptstep", default=1200000, type=int)47 48parser.add_argument("-nfe", "--nfestep", default=32, type=int)49parser.add_argument("-o", "--odemethod", default="euler")50parser.add_argument("-ss", "--swaysampling", default=-1, type=float)51 52parser.add_argument("-t", "--testset", required=True)53 54args = parser.parse_args()55 56 57seed = args.seed58dataset_name = args.dataset59exp_name = args.expname60ckpt_step = args.ckptstep61ckpt_path = f"ckpts/{exp_name}/model_{ckpt_step}.pt"62 63nfe_step = args.nfestep64ode_method = args.odemethod65sway_sampling_coef = args.swaysampling66 67testset = args.testset68 69 70infer_batch_size = 1  # max frames. 1 for ddp single inference (recommended)71cfg_strength = 2.072speed = 1.073use_truth_duration = False74no_ref_audio = False75 76 77if exp_name == "F5TTS_Base":78    model_cls = DiT79    model_cfg = dict(dim=1024, depth=22, heads=16, ff_mult=2, text_dim=512, conv_layers=4)80 81elif exp_name == "E2TTS_Base":82    model_cls = UNetT83    model_cfg = dict(dim=1024, depth=24, heads=16, ff_mult=4)84 85 86if testset == "ls_pc_test_clean":87    metalst = "data/librispeech_pc_test_clean_cross_sentence.lst"88    librispeech_test_clean_path = "<SOME_PATH>/LibriSpeech/test-clean"  # test-clean path89    metainfo = get_librispeech_test_clean_metainfo(metalst, librispeech_test_clean_path)90 91elif testset == "seedtts_test_zh":92    metalst = "data/seedtts_testset/zh/meta.lst"93    metainfo = get_seedtts_testset_metainfo(metalst)94 95elif testset == "seedtts_test_en":96    metalst = "data/seedtts_testset/en/meta.lst"97    metainfo = get_seedtts_testset_metainfo(metalst)98 99 100# path to save genereted wavs101if seed is None:102    seed = random.randint(-10000, 10000)103output_dir = (104    f"results/{exp_name}_{ckpt_step}/{testset}/"105    f"seed{seed}_{ode_method}_nfe{nfe_step}"106    f"{f'_ss{sway_sampling_coef}' if sway_sampling_coef else ''}"107    f"_cfg{cfg_strength}_speed{speed}"108    f"{'_gt-dur' if use_truth_duration else ''}"109    f"{'_no-ref-audio' if no_ref_audio else ''}"110)111 112 113# -------------------------------------------------#114 115use_ema = True116 117prompts_all = get_inference_prompt(118    metainfo,119    speed=speed,120    tokenizer=tokenizer,121    target_sample_rate=target_sample_rate,122    n_mel_channels=n_mel_channels,123    hop_length=hop_length,124    target_rms=target_rms,125    use_truth_duration=use_truth_duration,126    infer_batch_size=infer_batch_size,127)128 129# Vocoder model130local = False131if local:132    vocos_local_path = "../checkpoints/charactr/vocos-mel-24khz"133    vocos = Vocos.from_hparams(f"{vocos_local_path}/config.yaml")134    state_dict = torch.load(f"{vocos_local_path}/pytorch_model.bin", weights_only=True, map_location=device)135    vocos.load_state_dict(state_dict)136    vocos.eval()137else:138    vocos = Vocos.from_pretrained("charactr/vocos-mel-24khz")139 140# Tokenizer141vocab_char_map, vocab_size = get_tokenizer(dataset_name, tokenizer)142 143# Model144model = CFM(145    transformer=model_cls(**model_cfg, text_num_embeds=vocab_size, mel_dim=n_mel_channels),146    mel_spec_kwargs=dict(147        target_sample_rate=target_sample_rate,148        n_mel_channels=n_mel_channels,149        hop_length=hop_length,150    ),151    odeint_kwargs=dict(152        method=ode_method,153    ),154    vocab_char_map=vocab_char_map,155).to(device)156 157model = load_checkpoint(model, ckpt_path, device, use_ema=use_ema)158 159if not os.path.exists(output_dir) and accelerator.is_main_process:160    os.makedirs(output_dir)161 162# start batch inference163accelerator.wait_for_everyone()164start = time.time()165 166with accelerator.split_between_processes(prompts_all) as prompts:167    for prompt in tqdm(prompts, disable=not accelerator.is_local_main_process):168        utts, ref_rms_list, ref_mels, ref_mel_lens, total_mel_lens, final_text_list = prompt169        ref_mels = ref_mels.to(device)170        ref_mel_lens = torch.tensor(ref_mel_lens, dtype=torch.long).to(device)171        total_mel_lens = torch.tensor(total_mel_lens, dtype=torch.long).to(device)172 173        # Inference174        with torch.inference_mode():175            generated, _ = model.sample(176                cond=ref_mels,177                text=final_text_list,178                duration=total_mel_lens,179                lens=ref_mel_lens,180                steps=nfe_step,181                cfg_strength=cfg_strength,182                sway_sampling_coef=sway_sampling_coef,183                no_ref_audio=no_ref_audio,184                seed=seed,185            )186        # Final result187        for i, gen in enumerate(generated):188            gen = gen[ref_mel_lens[i] : total_mel_lens[i], :].unsqueeze(0)189            gen_mel_spec = gen.permute(0, 2, 1)190            generated_wave = vocos.decode(gen_mel_spec.cpu())191            if ref_rms_list[i] < target_rms:192                generated_wave = generated_wave * ref_rms_list[i] / target_rms193            torchaudio.save(f"{output_dir}/{utts[i]}.wav", generated_wave, target_sample_rate)194 195accelerator.wait_for_everyone()196if accelerator.is_main_process:197    timediff = time.time() - start198    print(f"Done batch inference in {timediff / 60 :.2f} minutes.")199