| 1 | import random |
| 2 | import sys |
| 3 | from importlib.resources import files |
| 4 | |
| 5 | import soundfile as sf |
| 6 | import tqdm |
| 7 | from cached_path import cached_path |
| 8 | from hydra.utils import get_class |
| 9 | from omegaconf import OmegaConf |
| 10 | |
| 11 | from f5_tts.infer.utils_infer import ( |
| 12 | infer_process, |
| 13 | load_model, |
| 14 | load_vocoder, |
| 15 | preprocess_ref_audio_text, |
| 16 | remove_silence_for_generated_wav, |
| 17 | save_spectrogram, |
| 18 | transcribe, |
| 19 | ) |
| 20 | from f5_tts.model.utils import seed_everything |
| 21 | |
| 22 | |
| 23 | class F5TTS: |
| 24 | def __init__( |
| 25 | self, |
| 26 | model="F5TTS_v1_Base", |
| 27 | ckpt_file="", |
| 28 | vocab_file="", |
| 29 | ode_method="euler", |
| 30 | use_ema=True, |
| 31 | vocoder_local_path=None, |
| 32 | device=None, |
| 33 | hf_cache_dir=None, |
| 34 | ): |
| 35 | model_cfg = OmegaConf.load(str(files("f5_tts").joinpath(f"configs/{model}.yaml"))) |
| 36 | model_cls = get_class(f"f5_tts.model.{model_cfg.model.backbone}") |
| 37 | model_arc = model_cfg.model.arch |
| 38 | |
| 39 | self.mel_spec_type = model_cfg.model.mel_spec.mel_spec_type |
| 40 | self.target_sample_rate = model_cfg.model.mel_spec.target_sample_rate |
| 41 | |
| 42 | self.ode_method = ode_method |
| 43 | self.use_ema = use_ema |
| 44 | |
| 45 | if device is not None: |
| 46 | self.device = device |
| 47 | else: |
| 48 | import torch |
| 49 | |
| 50 | self.device = ( |
| 51 | "cuda" |
| 52 | if torch.cuda.is_available() |
| 53 | else "xpu" |
| 54 | if torch.xpu.is_available() |
| 55 | else "mps" |
| 56 | if torch.backends.mps.is_available() |
| 57 | else "cpu" |
| 58 | ) |
| 59 | |
| 60 | # Load models |
| 61 | self.vocoder = load_vocoder( |
| 62 | self.mel_spec_type, vocoder_local_path is not None, vocoder_local_path, self.device, hf_cache_dir |
| 63 | ) |
| 64 | |
| 65 | repo_name, ckpt_step, ckpt_type = "F5-TTS", 1250000, "safetensors" |
| 66 | |
| 67 | # override for previous models |
| 68 | if model == "F5TTS_Base": |
| 69 | if self.mel_spec_type == "vocos": |
| 70 | ckpt_step = 1200000 |
| 71 | elif self.mel_spec_type == "bigvgan": |
| 72 | model = "F5TTS_Base_bigvgan" |
| 73 | ckpt_type = "pt" |
| 74 | elif model == "E2TTS_Base": |
| 75 | repo_name = "E2-TTS" |
| 76 | ckpt_step = 1200000 |
| 77 | |
| 78 | if not ckpt_file: |
| 79 | ckpt_file = str( |
| 80 | cached_path(f"hf://SWivid/{repo_name}/{model}/model_{ckpt_step}.{ckpt_type}", cache_dir=hf_cache_dir) |
| 81 | ) |
| 82 | self.ema_model = load_model( |
| 83 | model_cls, model_arc, ckpt_file, self.mel_spec_type, vocab_file, self.ode_method, self.use_ema, self.device |
| 84 | ) |
| 85 | |
| 86 | def transcribe(self, ref_audio, language=None): |
| 87 | return transcribe(ref_audio, language) |
| 88 | |
| 89 | def export_wav(self, wav, file_wave, remove_silence=False): |
| 90 | sf.write(file_wave, wav, self.target_sample_rate) |
| 91 | |
| 92 | if remove_silence: |
| 93 | remove_silence_for_generated_wav(file_wave) |
| 94 | |
| 95 | def export_spectrogram(self, spec, file_spec): |
| 96 | save_spectrogram(spec, file_spec) |
| 97 | |
| 98 | def infer( |
| 99 | self, |
| 100 | ref_file, |
| 101 | ref_text, |
| 102 | gen_text, |
| 103 | show_info=print, |
| 104 | progress=tqdm, |
| 105 | target_rms=0.1, |
| 106 | cross_fade_duration=0.15, |
| 107 | sway_sampling_coef=-1, |
| 108 | cfg_strength=2, |
| 109 | nfe_step=32, |
| 110 | speed=1.0, |
| 111 | fix_duration=None, |
| 112 | remove_silence=False, |
| 113 | file_wave=None, |
| 114 | file_spec=None, |
| 115 | seed=None, |
| 116 | ): |
| 117 | if seed is None: |
| 118 | seed = random.randint(0, sys.maxsize) |
| 119 | seed_everything(seed) |
| 120 | self.seed = seed |
| 121 | |
| 122 | ref_file, ref_text = preprocess_ref_audio_text(ref_file, ref_text, show_info=show_info) |
| 123 | |
| 124 | wav, sr, spec = infer_process( |
| 125 | ref_file, |
| 126 | ref_text, |
| 127 | gen_text, |
| 128 | self.ema_model, |
| 129 | self.vocoder, |
| 130 | self.mel_spec_type, |
| 131 | show_info=show_info, |
| 132 | progress=progress, |
| 133 | target_rms=target_rms, |
| 134 | cross_fade_duration=cross_fade_duration, |
| 135 | nfe_step=nfe_step, |
| 136 | cfg_strength=cfg_strength, |
| 137 | sway_sampling_coef=sway_sampling_coef, |
| 138 | speed=speed, |
| 139 | fix_duration=fix_duration, |
| 140 | device=self.device, |
| 141 | ) |
| 142 | |
| 143 | if file_wave is not None: |
| 144 | self.export_wav(wav, file_wave, remove_silence) |
| 145 | |
| 146 | if file_spec is not None: |
| 147 | self.export_spectrogram(spec, file_spec) |
| 148 | |
| 149 | return wav, sr, spec |
| 150 | |
| 151 | |
| 152 | if __name__ == "__main__": |
| 153 | f5tts = F5TTS() |
| 154 | |
| 155 | wav, sr, spec = f5tts.infer( |
| 156 | ref_file=str(files("f5_tts").joinpath("infer/examples/basic/basic_ref_en.wav")), |
| 157 | ref_text="Some call me nature, others call me mother nature.", |
| 158 | gen_text="I don't really care what you call me. I've been a silent spectator, watching species evolve, empires rise and fall. But always remember, I am mighty and enduring.", |
| 159 | file_wave=str(files("f5_tts").joinpath("../../tests/api_out.wav")), |
| 160 | file_spec=str(files("f5_tts").joinpath("../../tests/api_out.png")), |
| 161 | seed=None, |
| 162 | ) |
| 163 | |
| 164 | print("seed :", f5tts.seed) |
| 165 |