prabaerode/zero-shot-tts
0
1import argparse2import codecs3import re4from pathlib import Path5 6import numpy as np7import soundfile as sf8import tomli9from cached_path import cached_path10 11from model import DiT, UNetT12from model.utils_infer import (13 load_vocoder,14 load_model,15 preprocess_ref_audio_text,16 infer_process,17 remove_silence_for_generated_wav,18)19 20 21parser = argparse.ArgumentParser(22 prog="python3 inference-cli.py",23 description="Commandline interface for E2/F5 TTS with Advanced Batch Processing.",24 epilog="Specify options above to override one or more settings from config.",25)26parser.add_argument(27 "-c",28 "--config",29 help="Configuration file. Default=cli-config.toml",30 default="inference-cli.toml",31)32parser.add_argument(33 "-m",34 "--model",35 help="F5-TTS | E2-TTS",36)37parser.add_argument(38 "-p",39 "--ckpt_file",40 help="The Checkpoint .pt",41)42parser.add_argument(43 "-v",44 "--vocab_file",45 help="The vocab .txt",46)47parser.add_argument("-r", "--ref_audio", type=str, help="Reference audio file < 15 seconds.")48parser.add_argument("-s", "--ref_text", type=str, default="666", help="Subtitle for the reference audio.")49parser.add_argument(50 "-t",51 "--gen_text",52 type=str,53 help="Text to generate.",54)55parser.add_argument(56 "-f",57 "--gen_file",58 type=str,59 help="File with text to generate. Ignores --text",60)61parser.add_argument(62 "-o",63 "--output_dir",64 type=str,65 help="Path to output folder..",66)67parser.add_argument(68 "--remove_silence",69 help="Remove silence.",70)71parser.add_argument(72 "--load_vocoder_from_local",73 action="store_true",74 help="load vocoder from local. Default: ../checkpoints/charactr/vocos-mel-24khz",75)76args = parser.parse_args()77 78config = tomli.load(open(args.config, "rb"))79 80ref_audio = args.ref_audio if args.ref_audio else config["ref_audio"]81ref_text = args.ref_text if args.ref_text != "666" else config["ref_text"]82gen_text = args.gen_text if args.gen_text else config["gen_text"]83gen_file = args.gen_file if args.gen_file else config["gen_file"]84if gen_file:85 gen_text = codecs.open(gen_file, "r", "utf-8").read()86output_dir = args.output_dir if args.output_dir else config["output_dir"]87model = args.model if args.model else config["model"]88ckpt_file = args.ckpt_file if args.ckpt_file else ""89vocab_file = args.vocab_file if args.vocab_file else ""90remove_silence = args.remove_silence if args.remove_silence else config["remove_silence"]91wave_path = Path(output_dir) / "out.wav"92spectrogram_path = Path(output_dir) / "out.png"93vocos_local_path = "../checkpoints/charactr/vocos-mel-24khz"94 95vocos = load_vocoder(is_local=args.load_vocoder_from_local, local_path=vocos_local_path)96 97 98# load models99if model == "F5-TTS":100 model_cls = DiT101 model_cfg = dict(dim=1024, depth=22, heads=16, ff_mult=2, text_dim=512, conv_layers=4)102 if ckpt_file == "":103 repo_name = "F5-TTS"104 exp_name = "F5TTS_Base"105 ckpt_step = 1200000106 ckpt_file = str(cached_path(f"hf://SWivid/{repo_name}/{exp_name}/model_{ckpt_step}.safetensors"))107 # ckpt_file = f"ckpts/{exp_name}/model_{ckpt_step}.pt" # .pt | .safetensors; local path108 109elif model == "E2-TTS":110 model_cls = UNetT111 model_cfg = dict(dim=1024, depth=24, heads=16, ff_mult=4)112 if ckpt_file == "":113 repo_name = "E2-TTS"114 exp_name = "E2TTS_Base"115 ckpt_step = 1200000116 ckpt_file = str(cached_path(f"hf://SWivid/{repo_name}/{exp_name}/model_{ckpt_step}.safetensors"))117 # ckpt_file = f"ckpts/{exp_name}/model_{ckpt_step}.pt" # .pt | .safetensors; local path118 119print(f"Using {model}...")120ema_model = load_model(model_cls, model_cfg, ckpt_file, vocab_file)121 122 123def main_process(ref_audio, ref_text, text_gen, model_obj, remove_silence):124 main_voice = {"ref_audio": ref_audio, "ref_text": ref_text}125 if "voices" not in config:126 voices = {"main": main_voice}127 else:128 voices = config["voices"]129 voices["main"] = main_voice130 for voice in voices:131 voices[voice]["ref_audio"], voices[voice]["ref_text"] = preprocess_ref_audio_text(132 voices[voice]["ref_audio"], voices[voice]["ref_text"]133 )134 print("Voice:", voice)135 print("Ref_audio:", voices[voice]["ref_audio"])136 print("Ref_text:", voices[voice]["ref_text"])137 138 generated_audio_segments = []139 reg1 = r"(?=\[\w+\])"140 chunks = re.split(reg1, text_gen)141 reg2 = r"\[(\w+)\]"142 for text in chunks:143 match = re.match(reg2, text)144 if match:145 voice = match[1]146 else:147 print("No voice tag found, using main.")148 voice = "main"149 if voice not in voices:150 print(f"Voice {voice} not found, using main.")151 voice = "main"152 text = re.sub(reg2, "", text)153 gen_text = text.strip()154 ref_audio = voices[voice]["ref_audio"]155 ref_text = voices[voice]["ref_text"]156 print(f"Voice: {voice}")157 audio, final_sample_rate, spectragram = infer_process(ref_audio, ref_text, gen_text, model_obj)158 generated_audio_segments.append(audio)159 160 if generated_audio_segments:161 final_wave = np.concatenate(generated_audio_segments)162 with open(wave_path, "wb") as f:163 sf.write(f.name, final_wave, final_sample_rate)164 # Remove silence165 if remove_silence:166 remove_silence_for_generated_wav(f.name)167 print(f.name)168 169 170main_process(ref_audio, ref_text, gen_text, ema_model, remove_silence)171 