| |
| try: |
| import spaces |
| USING_SPACES = True |
| except ImportError: |
| USING_SPACES = False |
|
|
| |
| import re |
| import gradio as gr |
| import numpy as np |
| import tempfile |
| from tqdm import tqdm |
| from einops import rearrange |
| from pydub import AudioSegment, silence |
| from model import UNetT, DiT |
| from cached_path import cached_path |
| from model.utils import ( |
| get_tokenizer, |
| convert_char_to_pinyin, |
| ) |
| from infer.utils_infer import ( |
| load_vocoder, |
| load_model, |
| remove_silence_edges, |
| remove_silence_for_generated_wav, |
| save_spectrogram, |
| ) |
| from tokenizers import Tokenizer |
| from phonemizer import phonemize |
|
|
| from transformers import pipeline |
| import click |
| import soundfile as sf |
|
|
| |
| import torch |
| import torchaudio |
|
|
| |
| def gpu_decorator(func): |
| if USING_SPACES: |
| return spaces.GPU(func) |
| else: |
| return func |
|
|
| |
| device = ( |
| "cuda" |
| if torch.cuda.is_available() |
| else "mps" if torch.backends.mps.is_available() else "cpu" |
| ) |
|
|
| |
| if device == "cuda": |
| dtype = torch.float16 |
| elif device == "cpu": |
| dtype = torch.float32 |
| else: |
| dtype = torch.float32 |
|
|
| |
| device = torch.device(device) |
| print(f"Using device: {device}, dtype: {dtype}") |
|
|
| pipe = pipeline( |
| "automatic-speech-recognition", |
| model="openai/whisper-large-v3-turbo", |
| torch_dtype=dtype, |
| device=device, |
| ) |
| vocos = load_vocoder() |
|
|
| |
| target_sample_rate = 24000 |
| n_mel_channels = 100 |
| hop_length = 256 |
| target_rms = 0.1 |
| nfe_step = 32 |
| cfg_strength = 2.0 |
| ode_method = "euler" |
| sway_sampling_coef = -1.0 |
| speed = 1 |
| fix_duration = None |
| ref_language = "en-us" |
| language = "en-us" |
|
|
| DEFAULT_TTS_MODEL = "F5-TTS" |
| tts_model_choice = DEFAULT_TTS_MODEL |
|
|
| def load_custom(ckpt_path: str, vocab_path="", model_cfg=None): |
| ckpt_path, vocab_path = ckpt_path.strip(), vocab_path.strip() |
| if ckpt_path.startswith("hf://"): |
| ckpt_path = str(cached_path(ckpt_path)) |
| if vocab_path.startswith("hf://"): |
| vocab_path = str(cached_path(vocab_path)) |
| if model_cfg is None: |
| model_cfg = dict(dim=1024, depth=22, heads=16, ff_mult=2, text_dim=512, conv_layers=4) |
| return load_model(DiT, model_cfg, ckpt_path, vocab_file=vocab_path) |
|
|
| |
| F5TTS_model_cfg = dict(dim=1024, depth=22, heads=16, ff_mult=2, text_dim=512, conv_layers=4) |
| E2TTS_model_cfg = dict(dim=1024, depth=24, heads=16, ff_mult=4) |
|
|
| F5TTS_ema_model = load_custom( |
| "hf://Gregniuki/F5-tts_English_German_Polish/multi3/model_900000.pt", "", F5TTS_model_cfg |
| ) |
|
|
| def chunk_text(text, max_chars): |
| |
| if max_chars > 180: |
| max_chars = 180 |
| if max_chars < 50: |
| max_chars = 50 |
| |
| split_after_space_chars = max_chars + int(max_chars * 0.1) |
| chunks = [] |
| current_chunk = "" |
| |
| sentences = re.split(r"(?<=[;:,.。、!?])\s+|(?<=[;:,.。、!?])", text) |
|
|
| for sentence in sentences: |
| if len(current_chunk) + len(sentence) + 1 <= max_chars: |
| current_chunk += sentence + " " |
| else: |
| while len(current_chunk) > split_after_space_chars: |
| split_index = current_chunk.rfind(" ", 0, split_after_space_chars) |
| if split_index == -1: |
| split_index = split_after_space_chars |
| chunks.append(current_chunk[:split_index].strip()) |
| current_chunk = current_chunk[split_index:].strip() |
| |
| if current_chunk: |
| chunks.append(current_chunk.strip()) |
| current_chunk = sentence + " " |
|
|
| while len(current_chunk) > split_after_space_chars: |
| split_index = current_chunk.rfind(" ", 0, split_after_space_chars) |
| if split_index == -1: |
| split_index = split_after_space_chars |
| chunks.append(current_chunk[:split_index].strip()) |
| current_chunk = current_chunk[split_index:].strip() |
|
|
| if current_chunk: |
| chunks.append(current_chunk.strip()) |
|
|
| return chunks |
|
|
| def text_to_ipa(text, language=language): |
| try: |
| ipa_text = phonemize( |
| text, |
| language=language, |
| backend='espeak', |
| strip=False, |
| preserve_punctuation=True, |
| with_stress=True |
| ) |
| ipa_text = re.sub(r'\([a-z]{2,3}\)', '', ipa_text) |
| ipa_text = re.sub(r'tʃˈaɪniːzlˈe̞tə', '', ipa_text) |
| ipa_text = re.sub(r'tʃˈaɪniːzɭˈetə', '', ipa_text) |
| ipa_text = re.sub(r'dʒˈapəniːzlˈe̞tə', '', ipa_text) |
| ipa_text = re.sub(r'dʒˈapəniːzɭˈetə', '', ipa_text) |
| return ipa_text |
| except Exception as e: |
| print(f"Error processing text: {text}. Error: {e}") |
| return None |
|
|
| @gpu_decorator |
| def infer_batch(ref_audio, ref_text, gen_text_batches, exp_name, remove_silence, cross_fade_duration=0.15, progress=gr.Progress()): |
| if exp_name == "Multi": |
| ema_model = F5TTS_ema_model |
|
|
| audio, sr = ref_audio |
| if audio.shape[0] > 1: |
| audio = torch.mean(audio, dim=0, keepdim=True) |
|
|
| rms = torch.sqrt(torch.mean(torch.square(audio))) |
| if rms < target_rms: |
| audio = audio * target_rms / rms |
| if sr != target_sample_rate: |
| resampler = torchaudio.transforms.Resample(sr, target_sample_rate) |
| audio = resampler(audio) |
| |
| audio = audio.to(device) |
| tokenizer = Tokenizer.from_file("data/Emilia_ZH_EN_pinyin/tokenizer.json") |
|
|
| generated_waves = [] |
| spectrograms = [] |
| punctuation_weights = {",": 0, ".": 0, " ": 0} |
| |
| progress_bar = tqdm(gen_text_batches) |
| ipa_text_ref = text_to_ipa(ref_text, language=ref_language) |
|
|
| for i, gen_text in enumerate(progress_bar): |
| ipa_text_gen = text_to_ipa(gen_text, language=language) |
| text_list = ipa_text_ref + ipa_text_gen |
| encoding = tokenizer.encode(text_list) |
| tokens = encoding.tokens |
| text_list = ' '.join(map(str, tokens)) |
| final_text_list = [text_list] |
|
|
| ref_audio_len = audio.shape[-1] // hop_length |
|
|
| if fix_duration is not None: |
| duration = int(fix_duration * target_sample_rate / hop_length) |
| else: |
| def calculate_weighted_length(t): |
| return len(t.encode("utf-8")) + sum(punctuation_weights.get(char, 0) for char in t) |
|
|
| ref_text_len = calculate_weighted_length(ref_text) |
| gen_text_len = calculate_weighted_length(gen_text) |
| duration = max(250, int(ref_audio_len) + int(((ref_audio_len / ref_text_len) * gen_text_len) / speed)) |
|
|
| print(f"Chunk {i + 1}: Duration: {duration} speed {speed}") |
| |
| with torch.inference_mode(): |
| audio = audio.to(ema_model.device) |
| final_text_list = [t.to(ema_model.device) if isinstance(t, torch.Tensor) else t for t in final_text_list] |
| generated, _ = ema_model.sample( |
| cond=audio, |
| text=final_text_list, |
| duration=duration, |
| steps=nfe_step, |
| cfg_strength=cfg_strength, |
| sway_sampling_coef=sway_sampling_coef, |
| ) |
|
|
| generated = generated[:, ref_audio_len:, :] |
| generated_mel_spec = rearrange(generated, "1 n d -> 1 d n") |
| generated_wave = vocos.decode(generated_mel_spec) |
|
|
| if rms < target_rms: |
| generated_wave = generated_wave * rms / target_rms |
|
|
| generated_wave = generated_wave.squeeze().cpu().numpy() |
| generated_waves.append(generated_wave) |
| |
| mel_spec_np = generated_mel_spec[0].to(dtype=torch.float32).cpu().numpy() |
| spectrograms.append(mel_spec_np) |
|
|
| if cross_fade_duration <= 0 or len(generated_waves) == 1: |
| final_wave = np.concatenate(generated_waves) |
| else: |
| final_wave = generated_waves[0] |
| for i in range(1, len(generated_waves)): |
| prev_wave = final_wave |
| next_wave = generated_waves[i] |
|
|
| cross_fade_samples = int(cross_fade_duration * target_sample_rate) |
| cross_fade_samples = min(cross_fade_samples, len(prev_wave), len(next_wave)) |
|
|
| if cross_fade_samples <= 0: |
| final_wave = np.concatenate([prev_wave, next_wave]) |
| continue |
|
|
| prev_overlap = prev_wave[-cross_fade_samples:] |
| next_overlap = next_wave[:cross_fade_samples] |
|
|
| fade_out = np.linspace(1, 0, cross_fade_samples) |
| fade_in = np.linspace(0, 1, cross_fade_samples) |
|
|
| cross_faded_overlap = prev_overlap * fade_out + next_overlap * fade_in |
| final_wave = np.concatenate([ |
| prev_wave[:-cross_fade_samples], |
| cross_faded_overlap, |
| next_wave[cross_fade_samples:] |
| ]) |
|
|
| if remove_silence: |
| with tempfile.NamedTemporaryFile(delete=False, suffix=".wav") as f: |
| final_wave_float32 = final_wave.astype(np.float32) |
| sf.write(f.name, final_wave_float32, target_sample_rate) |
| aseg = AudioSegment.from_file(f.name) |
| non_silent_segs = silence.split_on_silence(aseg, min_silence_len=1000, silence_thresh=-50, keep_silence=500) |
| non_silent_wave = AudioSegment.silent(duration=0) |
| for non_silent_seg in non_silent_segs: |
| non_silent_wave += non_silent_seg |
| aseg = non_silent_wave |
| aseg.export(f.name, format="wav") |
| final_wave, _ = torchaudio.load(f.name) |
| final_wave = final_wave.squeeze().cpu().numpy() |
|
|
| combined_spectrogram = np.concatenate(spectrograms, axis=1) |
| with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as tmp_spectrogram: |
| spectrogram_path = tmp_spectrogram.name |
| save_spectrogram(combined_spectrogram, spectrogram_path) |
|
|
| return (target_sample_rate, final_wave), spectrogram_path |
|
|
| @gpu_decorator |
| def infer(ref_audio_orig, ref_text, gen_text, exp_name, remove_silence, cross_fade_duration=0.15): |
| |
| ref_text = ref_text or "" |
| |
| gr.Info("Converting audio...") |
| with tempfile.NamedTemporaryFile(delete=False, suffix=".wav") as f: |
| aseg = AudioSegment.from_file(ref_audio_orig) |
| aseg = remove_silence_edges(aseg) + AudioSegment.silent(duration=150) |
| non_silent_segs = silence.split_on_silence( |
| aseg, min_silence_len=700, silence_thresh=-50, keep_silence=700 |
| ) |
| non_silent_wave = AudioSegment.silent(duration=0) |
| for non_silent_seg in non_silent_segs: |
| non_silent_wave += non_silent_seg |
| aseg = non_silent_wave |
| |
| audio_duration = len(aseg) |
| if audio_duration > 10000: |
| gr.Warning("Audio is over 10s, clipping to only first 10s.") |
| aseg = aseg[:10000] |
| aseg.export(f.name, format="wav") |
| ref_audio = f.name |
|
|
| if not ref_text.strip(): |
| gr.Info("No reference text provided, transcribing reference audio...") |
| ref_text = pipe( |
| ref_audio, |
| chunk_length_s=15, |
| batch_size=128, |
| generate_kwargs={"task": "transcribe"}, |
| return_timestamps=False, |
| )["text"].strip() |
| gr.Info("Finished transcription") |
| else: |
| gr.Info("Using custom reference text...") |
|
|
| if not ref_text.endswith(". "): |
| if ref_text.endswith("."): |
| ref_text += " " |
| else: |
| ref_text += ". " |
|
|
| audio, sr = torchaudio.load(ref_audio) |
| |
| |
| max_chars = int((len(ref_text.encode('utf-8')) / (audio.shape[-1] / sr) * (25 - audio.shape[-1] / sr ))) |
| gen_text_batches = chunk_text(gen_text, max_chars=max_chars) |
| |
| gr.Info(f"Generating audio using {exp_name} in {len(gen_text_batches)} batches") |
| return infer_batch((audio, sr), ref_text, gen_text_batches, exp_name, remove_silence, cross_fade_duration) |
|
|
| def parse_speechtypes_text(gen_text): |
| pattern = r'\((.*?)\)' |
| tokens = re.split(pattern, gen_text) |
| segments = [] |
| current_emotion = 'Regular' |
|
|
| for i in range(len(tokens)): |
| if i % 2 == 0: |
| text = tokens[i].strip() |
| if text: |
| segments.append({'emotion': current_emotion, 'text': text}) |
| else: |
| emotion = tokens[i].strip() |
| current_emotion = emotion |
| return segments |
|
|
| def update_language(new_language): |
| global language |
| language = new_language |
|
|
| def update_language1(new_ref_language): |
| global ref_language |
| ref_language = new_ref_language |
|
|
| def update_speed(new_speed): |
| global speed |
| speed = new_speed |
| return f"Speed set to: {speed}" |
|
|
| |
| with gr.Blocks() as app_credits: |
| gr.Markdown(""" |
| # Credits |
| * [mrfakename](https://github.com/fakerybakery) for the original [online demo](https://huggingface.co/spaces/mrfakename/E2-F5-TTS) |
| * [RootingInLoad](https://github.com/RootingInLoad) for the podcast generation |
| * [jpgallegoar](https://github.com/jpgallegoar) for multiple speech-type generation |
| """) |
|
|
| LANG_CHOICES = [ |
| |
| "af", "nl", "en-us", "en-gb", "en-029", "en-gb-x-gbclan", "en-gb-x-rp", "en-gb-scotland", "en-gb-x-gbcwmd", "de", "lb", |
| |
| "my", "yue", "hak", "cmn", |
| |
| "an", "ca", "fr-be", "fr-fr", "fr-ch", "ht", "it", "pap", "pt-br", "pt", "ro", "es", "es-419", |
| |
| "as", "bn", "bpy", "gu", "hi", "kok", "mr", "ne", "or", "pa", "sd", "si", "ur", |
| |
| "am", "ar", "he", "mt", |
| |
| "be", "ru", "ru-lv", "uk", |
| |
| "ja", |
| |
| "ko", |
| |
| "id", "mi", "ms", |
| |
| "az", "ba", "cu", "kk", "ky", "nog", "tk", "tt", "tr", "ug", "uz", |
| |
| "vi-vn-x-central", "vi", "vi-vn-x-south", |
| |
| "kn", "ml", "ta", "te", |
| |
| "shn", "th", |
| |
| "cs", "pl", "sk", |
| |
| "da", "is", "nb", "sv", |
| |
| "fa", "fa-latn", "ku", |
| |
| "tn", "sw", |
| |
| "grc", "el", |
| |
| "et", "fi", "hu", "smj", |
| |
| "bs", "bg", "hr", "mk", "sr", "sl", |
| |
| "sq", "hy", "hyw", |
| |
| "om", |
| |
| "ka", |
| |
| "ga", "gd", "cy", |
| |
| "ltg", "lv", "lt", |
| |
| "gn", |
| |
| "quc", "qu", |
| |
| "nci", |
| |
| "kl", |
| |
| "chr", |
| |
| "haw", |
| |
| "la", |
| |
| "eo", "ia", "io", "lfn", "jbo", "py", "qdb", "qya", "piqd", "sjn" |
| ] |
| with gr.Blocks() as app_tts: |
| gr.Markdown("# Batched TTS") |
| ref_audio_input = gr.Audio(label="Reference Audio", type="filepath") |
| gen_text_input = gr.Textbox(label="Text to Generate", lines=10) |
| model_choice = gr.Radio(choices=["Multi"], label="Choose TTS Model", value="Multi") |
| |
| gr.Markdown("#Select Reference Language") |
| language_choice1 = gr.Dropdown(choices=LANG_CHOICES, label="Choose Language", value="de") |
| |
| gr.Markdown("#Select Synthesized Language") |
| language_choice = gr.Dropdown(choices=LANG_CHOICES, label="Choose Language", value="de") |
| |
| generate_btn = gr.Button("Synthesize", variant="primary") |
| with gr.Accordion("Advanced Settings", open=False): |
| ref_text_input = gr.Textbox( |
| label="Reference Text", |
| info="Leave blank to automatically transcribe the reference audio.", |
| lines=2 |
| ) |
| remove_silence = gr.Checkbox(label="Remove Silences", value=False) |
| speed_slider = gr.Slider(label="Speed", minimum=0.3, maximum=2.0, value=1.0, step=0.1) |
| cross_fade_duration_slider = gr.Slider(label="Cross-Fade Duration (s)", minimum=0.0, maximum=1.0, value=0.15, step=0.01) |
| |
| language_status = gr.Textbox(label="Current Language", interactive=False) |
| ref_language_status = gr.Textbox(label="Reference Language", interactive=False) |
| |
| speed_slider.change(update_speed, inputs=speed_slider) |
| language_choice.change(update_language, inputs=language_choice, outputs=language_status) |
| language_choice1.change(update_language1, inputs=language_choice1, outputs=ref_language_status) |
|
|
| audio_output = gr.Audio(label="Synthesized Audio") |
| spectrogram_output = gr.Image(label="Spectrogram") |
|
|
| generate_btn.click( |
| infer, |
| inputs=[ref_audio_input, ref_text_input, gen_text_input, model_choice, remove_silence, cross_fade_duration_slider], |
| outputs=[audio_output, spectrogram_output], |
| ) |
|
|
| with gr.Blocks() as app_emotional: |
| gr.Markdown("# Multiple Speech-Type Generation") |
| with gr.Row(): |
| regular_name = gr.Textbox(value='Regular', label='Speech Type Name', interactive=False) |
| regular_audio = gr.Audio(label='Regular Reference Audio', type='filepath') |
| regular_ref_text = gr.Textbox(label='Reference Text (Regular)', lines=2) |
|
|
| max_speech_types = 10 |
| speech_type_names = [] |
| speech_type_audios = [] |
| speech_type_ref_texts = [] |
| speech_type_delete_btns = [] |
|
|
| for i in range(max_speech_types - 1): |
| with gr.Row(): |
| name_input = gr.Textbox(label='Speech Type Name', visible=False) |
| audio_input = gr.Audio(label='Reference Audio', type='filepath', visible=False) |
| ref_text_input = gr.Textbox(label='Reference Text', lines=2, visible=False) |
| delete_btn = gr.Button("Delete", variant="secondary", visible=False) |
| speech_type_names.append(name_input) |
| speech_type_audios.append(audio_input) |
| speech_type_ref_texts.append(ref_text_input) |
| speech_type_delete_btns.append(delete_btn) |
|
|
| add_speech_type_btn = gr.Button("Add Speech Type") |
| speech_type_count = gr.State(value=0) |
|
|
| def add_speech_type_fn(count): |
| if count < max_speech_types - 1: |
| count += 1 |
| name_updates = [gr.update(visible=True) if i < count else gr.update() for i in range(max_speech_types - 1)] |
| audio_updates = [gr.update(visible=True) if i < count else gr.update() for i in range(max_speech_types - 1)] |
| ref_text_updates = [gr.update(visible=True) if i < count else gr.update() for i in range(max_speech_types - 1)] |
| delete_updates = [gr.update(visible=True) if i < count else gr.update() for i in range(max_speech_types - 1)] |
| return [count] + name_updates + audio_updates + ref_text_updates + delete_updates |
| return [count] + [gr.update() for _ in range((max_speech_types - 1) * 4)] |
|
|
| add_speech_type_btn.click( |
| add_speech_type_fn, |
| inputs=speech_type_count, |
| outputs=[speech_type_count] + speech_type_names + speech_type_audios + speech_type_ref_texts + speech_type_delete_btns |
| ) |
|
|
| gen_text_input_emotional = gr.Textbox(label="Text to Generate", lines=10) |
| model_choice_emotional = gr.Radio(choices=["Multi"], label="Choose TTS Model", value="Multi") |
|
|
| with gr.Accordion("Advanced Settings", open=False): |
| remove_silence_emotional = gr.Radio(choices=["True", "False"], label="Remove Silences", value="False") |
| |
| generate_emotional_btn = gr.Button("Generate Emotional Speech", variant="primary") |
| audio_output_emotional = gr.Audio(label="Synthesized Audio") |
|
|
| @gpu_decorator |
| def generate_emotional_speech(regular_audio, regular_ref_text, gen_text, *args): |
| num_additional = max_speech_types - 1 |
| st_names = args[0:num_additional] |
| st_audios = args[num_additional: 2 * num_additional] |
| st_texts = args[2 * num_additional: 3 * num_additional] |
| m_choice = args[3 * num_additional] |
| rem_silence = args[3 * num_additional + 1] == "True" |
|
|
| speech_types = {'Regular': {'audio': regular_audio, 'ref_text': regular_ref_text}} |
| for n, a, t in zip(st_names, st_audios, st_texts): |
| if n and a: |
| speech_types[n] = {'audio': a, 'ref_text': t or ""} |
|
|
| segments = parse_speechtypes_text(gen_text) |
| generated_audio_segments = [] |
| sr = target_sample_rate |
|
|
| for segment in segments: |
| emotion = segment['emotion'] if segment['emotion'] in speech_types else 'Regular' |
| ref_audio_path = speech_types[emotion]['audio'] |
| ref_txt = speech_types[emotion].get('ref_text', '') or "" |
|
|
| audio, _ = infer(ref_audio_path, ref_txt, segment['text'], m_choice, rem_silence) |
| sr, audio_data = audio |
| generated_audio_segments.append(audio_data) |
|
|
| if generated_audio_segments: |
| return (sr, np.concatenate(generated_audio_segments)) |
| gr.Warning("No audio generated.") |
| return None |
|
|
| input_components = [regular_audio, regular_ref_text, gen_text_input_emotional] + speech_type_names + speech_type_audios + speech_type_ref_texts + [model_choice_emotional, remove_silence_emotional] |
|
|
| generate_emotional_btn.click( |
| generate_emotional_speech, |
| inputs=input_components, |
| outputs=audio_output_emotional |
| ) |
|
|
| with gr.Blocks() as app: |
| gr.Markdown("# F5 TTS Local Interface") |
| gr.TabbedInterface([app_tts, app_emotional, app_credits], ["TTS", "Multi-Style", "Credits"]) |
|
|
| @click.command() |
| @click.option("--port", "-p", default=None, type=int) |
| @click.option("--host", "-H", default=None) |
| @click.option("--share", "-s", default=False, is_flag=True) |
| @click.option("--api", "-a", default=True, is_flag=True) |
| def main(port, host, share, api): |
| global app |
| print("Starting app...") |
| app.queue(api_open=api).launch( |
| server_name=host, server_port=port, share=share |
| ) |
|
|
| if __name__ == "__main__": |
| main() |