CoolFace
Apppublic

prabaerode/zero-shot-tts

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
infer_cli.py221 linesDownload Raw Back to infer
1import argparse2import codecs3import os4import re5from importlib.resources import files6from pathlib import Path7 8import numpy as np9import soundfile as sf10import tomli11from cached_path import cached_path12 13from f5_tts.infer.utils_infer import (14    infer_process,15    load_model,16    load_vocoder,17    preprocess_ref_audio_text,18    remove_silence_for_generated_wav,19)20from f5_tts.model import DiT, UNetT21 22parser = argparse.ArgumentParser(23    prog="python3 infer-cli.py",24    description="Commandline interface for E2/F5 TTS with Advanced Batch Processing.",25    epilog="Specify options above to override one or more settings from config.",26)27parser.add_argument(28    "-c",29    "--config",30    help="Configuration file. Default=infer/examples/basic/basic.toml",31    default=os.path.join(files("f5_tts").joinpath("infer/examples/basic"), "basic.toml"),32)33parser.add_argument(34    "-m",35    "--model",36    help="F5-TTS | E2-TTS",37)38parser.add_argument(39    "-p",40    "--ckpt_file",41    help="The Checkpoint .pt",42)43parser.add_argument(44    "-v",45    "--vocab_file",46    help="The vocab .txt",47)48parser.add_argument("-r", "--ref_audio", type=str, help="Reference audio file < 15 seconds.")49parser.add_argument("-s", "--ref_text", type=str, default="666", help="Subtitle for the reference audio.")50parser.add_argument(51    "-t",52    "--gen_text",53    type=str,54    help="Text to generate.",55)56parser.add_argument(57    "-f",58    "--gen_file",59    type=str,60    help="File with text to generate. Ignores --text",61)62parser.add_argument(63    "-o",64    "--output_dir",65    type=str,66    help="Path to output folder..",67)68parser.add_argument(69    "--remove_silence",70    help="Remove silence.",71)72parser.add_argument("--vocoder_name", type=str, default="vocos", choices=["vocos", "bigvgan"], help="vocoder name")73parser.add_argument(74    "--load_vocoder_from_local",75    action="store_true",76    help="load vocoder from local. Default: ../checkpoints/charactr/vocos-mel-24khz",77)78parser.add_argument(79    "--speed",80    type=float,81    default=1.0,82    help="Adjust the speed of the audio generation (default: 1.0)",83)84args = parser.parse_args()85 86config = tomli.load(open(args.config, "rb"))87 88ref_audio = args.ref_audio if args.ref_audio else config["ref_audio"]89ref_text = args.ref_text if args.ref_text != "666" else config["ref_text"]90gen_text = args.gen_text if args.gen_text else config["gen_text"]91gen_file = args.gen_file if args.gen_file else config["gen_file"]92 93# patches for pip pkg user94if "infer/examples/" in ref_audio:95    ref_audio = str(files("f5_tts").joinpath(f"{ref_audio}"))96if "infer/examples/" in gen_file:97    gen_file = str(files("f5_tts").joinpath(f"{gen_file}"))98if "voices" in config:99    for voice in config["voices"]:100        voice_ref_audio = config["voices"][voice]["ref_audio"]101        if "infer/examples/" in voice_ref_audio:102            config["voices"][voice]["ref_audio"] = str(files("f5_tts").joinpath(f"{voice_ref_audio}"))103 104if gen_file:105    gen_text = codecs.open(gen_file, "r", "utf-8").read()106output_dir = args.output_dir if args.output_dir else config["output_dir"]107model = args.model if args.model else config["model"]108ckpt_file = args.ckpt_file if args.ckpt_file else ""109vocab_file = args.vocab_file if args.vocab_file else ""110remove_silence = args.remove_silence if args.remove_silence else config["remove_silence"]111speed = args.speed112wave_path = Path(output_dir) / "infer_cli_out.wav"113# spectrogram_path = Path(output_dir) / "infer_cli_out.png"114if args.vocoder_name == "vocos":115    vocoder_local_path = "../checkpoints/vocos-mel-24khz"116elif args.vocoder_name == "bigvgan":117    vocoder_local_path = "../checkpoints/bigvgan_v2_24khz_100band_256x"118mel_spec_type = args.vocoder_name119 120vocoder = load_vocoder(vocoder_name=mel_spec_type, is_local=args.load_vocoder_from_local, local_path=vocoder_local_path)121 122 123# load models124if model == "F5-TTS":125    model_cls = DiT126    model_cfg = dict(dim=1024, depth=22, heads=16, ff_mult=2, text_dim=512, conv_layers=4)127    if ckpt_file == "":128        if args.vocoder_name == "vocos":129            repo_name = "F5-TTS"130            exp_name = "F5TTS_Base"131            ckpt_step = 1200000132            ckpt_file = str(cached_path(f"hf://SWivid/{repo_name}/{exp_name}/model_{ckpt_step}.safetensors"))133            # ckpt_file = f"ckpts/{exp_name}/model_{ckpt_step}.pt"  # .pt | .safetensors; local path134        elif args.vocoder_name == "bigvgan":135            repo_name = "F5-TTS"136            exp_name = "F5TTS_Base_bigvgan"137            ckpt_step = 1250000138            ckpt_file = str(cached_path(f"hf://SWivid/{repo_name}/{exp_name}/model_{ckpt_step}.pt"))139 140elif model == "E2-TTS":141    model_cls = UNetT142    model_cfg = dict(dim=1024, depth=24, heads=16, ff_mult=4)143    if ckpt_file == "":144        repo_name = "E2-TTS"145        exp_name = "E2TTS_Base"146        ckpt_step = 1200000147        ckpt_file = str(cached_path(f"hf://SWivid/{repo_name}/{exp_name}/model_{ckpt_step}.safetensors"))148        # ckpt_file = f"ckpts/{exp_name}/model_{ckpt_step}.pt"  # .pt | .safetensors; local path149    elif args.vocoder_name == "bigvgan":  # TODO: need to test150        repo_name = "F5-TTS"151        exp_name = "F5TTS_Base_bigvgan"152        ckpt_step = 1250000153        ckpt_file = str(cached_path(f"hf://SWivid/{repo_name}/{exp_name}/model_{ckpt_step}.pt"))154 155 156print(f"Using {model}...")157ema_model = load_model(model_cls, model_cfg, ckpt_file, mel_spec_type=args.vocoder_name, vocab_file=vocab_file)158 159 160def main_process(ref_audio, ref_text, text_gen, model_obj, mel_spec_type, remove_silence, speed):161    main_voice = {"ref_audio": ref_audio, "ref_text": ref_text}162    if "voices" not in config:163        voices = {"main": main_voice}164    else:165        voices = config["voices"]166        voices["main"] = main_voice167    for voice in voices:168        voices[voice]["ref_audio"], voices[voice]["ref_text"] = preprocess_ref_audio_text(169            voices[voice]["ref_audio"], voices[voice]["ref_text"]170        )171        print("Voice:", voice)172        print("Ref_audio:", voices[voice]["ref_audio"])173        print("Ref_text:", voices[voice]["ref_text"])174 175    generated_audio_segments = []176    reg1 = r"(?=\[\w+\])"177    chunks = re.split(reg1, text_gen)178    reg2 = r"\[(\w+)\]"179    for text in chunks:180        if not text.strip():181            continue182        match = re.match(reg2, text)183        if match:184            voice = match[1]185        else:186            print("No voice tag found, using main.")187            voice = "main"188        if voice not in voices:189            print(f"Voice {voice} not found, using main.")190            voice = "main"191        text = re.sub(reg2, "", text)192        gen_text = text.strip()193        ref_audio = voices[voice]["ref_audio"]194        ref_text = voices[voice]["ref_text"]195        print(f"Voice: {voice}")196        audio, final_sample_rate, spectragram = infer_process(197            ref_audio, ref_text, gen_text, model_obj, vocoder, mel_spec_type=mel_spec_type, speed=speed198        )199        generated_audio_segments.append(audio)200 201    if generated_audio_segments:202        final_wave = np.concatenate(generated_audio_segments)203 204        if not os.path.exists(output_dir):205            os.makedirs(output_dir)206 207        with open(wave_path, "wb") as f:208            sf.write(f.name, final_wave, final_sample_rate)209            # Remove silence210            if remove_silence:211                remove_silence_for_generated_wav(f.name)212            print(f.name)213 214 215def main():216    main_process(ref_audio, ref_text, gen_text, ema_model, mel_spec_type, remove_silence, speed)217 218 219if __name__ == "__main__":220    main()221