prabaerode/zero-shot-tts
0
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 