Download app.py from formospeech/taiwanese-hakka-f5-tts: direct link, hf CLI and curl.
- Browser
- Download file 12.3 kB
-
https://huggingface.co/spaces/formospeech/taiwanese-hakka-f5-tts/resolve/main/app.py
- Command line
-
hf download hf://spaces/formospeech/taiwanese-hakka-f5-tts/app.py
-
curl -L -o app.py https://huggingface.co/spaces/formospeech/taiwanese-hakka-f5-tts/resolve/main/app.py
12.3 kB
| import re | |
| import tempfile | |
| from importlib.resources import files | |
| import gradio as gr | |
| import soundfile as sf | |
| import torch | |
| import torchcodec | |
| from cached_path import cached_path | |
| from formog2p.hakka import g2p | |
| from omegaconf import OmegaConf | |
| from torchcodec.decoders import AudioDecoder | |
| try: | |
| import spaces | |
| USING_SPACES = True | |
| except ImportError: | |
| USING_SPACES = False | |
| from f5_tts.infer.utils_infer import ( | |
| device, | |
| hop_length, | |
| infer_process, | |
| load_checkpoint, | |
| load_vocoder, | |
| mel_spec_type, | |
| n_fft, | |
| n_mel_channels, | |
| ode_method, | |
| preprocess_ref_audio_text, | |
| remove_silence_for_generated_wav, | |
| save_spectrogram, | |
| target_sample_rate, | |
| win_length, | |
| ) | |
| from f5_tts.model import CFM, DiT | |
| from f5_tts.model.utils import get_tokenizer | |
| DIALECT_MAP = { | |
| "四縣腔": "客語_四縣", | |
| "海陸腔": "客語_海陸", | |
| "南四縣腔": "客語_南四縣", | |
| "大埔腔": "客語_大埔", | |
| "饒平腔": "客語_饒平", | |
| "詔安腔": "客語_詔安", | |
| } | |
| def to_pinyin(text, dialect): | |
| # Split text by English parts, convert non-English to pinyin, keep English as-is | |
| # Pattern matches English words (including spaces between them) | |
| pattern = r"([a-zA-Z]+(?:\s+[a-zA-Z]+)*)" | |
| parts = re.split(pattern, text) | |
| result_parts = [] | |
| for part in parts: | |
| if not part: | |
| continue | |
| # Check if part is English | |
| if re.match(r"^[a-zA-Z]+(?:\s+[a-zA-Z]+)*$", part): | |
| # Keep English as-is | |
| result_parts.append(part) | |
| else: | |
| # Convert non-English to pinyin | |
| result = g2p(part, dialect, pronunciation_type="pinyin") | |
| if len(result.unknown_words) > 0: | |
| raise gr.Error( | |
| f"參考句子中的[{','.join(result.unknown_words)}]目前無法轉成 pinyin。請嘗試其他句子。" | |
| ) | |
| pinyin = " ".join(result.pronunciations) | |
| result_parts.append(pinyin) | |
| combined = " ".join(result_parts) | |
| combined = re.sub(r"[!?]", "。", combined) | |
| combined = re.sub(r"(\s?)([,。])(\s?)", r"\2", combined) | |
| return combined | |
| def gpu_decorator(func): | |
| if USING_SPACES: | |
| return spaces.GPU(func) | |
| else: | |
| return func | |
| vocoder = load_vocoder() | |
| def load_model( | |
| model_cls, | |
| model_cfg, | |
| ckpt_path, | |
| mel_spec_type=mel_spec_type, | |
| vocab_file="", | |
| ode_method=ode_method, | |
| use_ema=True, | |
| device=device, | |
| fp16=False, | |
| ): | |
| if vocab_file == "": | |
| vocab_file = str(files("f5_tts").joinpath("infer/examples/vocab.txt")) | |
| tokenizer = "custom" | |
| print("\nvocab : ", vocab_file) | |
| print("token : ", tokenizer) | |
| print("model : ", ckpt_path, "\n") | |
| vocab_char_map, vocab_size = get_tokenizer(vocab_file, tokenizer) | |
| model = CFM( | |
| transformer=model_cls( | |
| **model_cfg, text_num_embeds=vocab_size, mel_dim=n_mel_channels | |
| ), | |
| mel_spec_kwargs=dict( | |
| n_fft=n_fft, | |
| hop_length=hop_length, | |
| win_length=win_length, | |
| n_mel_channels=n_mel_channels, | |
| target_sample_rate=target_sample_rate, | |
| mel_spec_type=mel_spec_type, | |
| ), | |
| odeint_kwargs=dict( | |
| method=ode_method, | |
| ), | |
| vocab_char_map=vocab_char_map, | |
| ).to(device) | |
| dtype = torch.float32 if mel_spec_type == "bigvgan" or not fp16 else None | |
| model = load_checkpoint(model, ckpt_path, device, dtype=dtype, use_ema=use_ema) | |
| return model | |
| def load_f5tts(ckpt_path, vocab_path, fp16=False): | |
| ckpt_path = str(cached_path(ckpt_path)) | |
| F5TTS_model_cfg = dict( | |
| dim=1024, depth=22, heads=16, ff_mult=2, text_dim=512, conv_layers=4 | |
| ) | |
| vocab_path = str(cached_path(vocab_path)) | |
| return load_model( | |
| DiT, | |
| F5TTS_model_cfg, | |
| ckpt_path, | |
| vocab_file=vocab_path, | |
| use_ema=False, | |
| fp16=fp16, | |
| ) | |
| OmegaConf.register_new_resolver("load_f5tts", load_f5tts) | |
| models_config = OmegaConf.to_object(OmegaConf.load("configs/models.yaml")) | |
| DEFAULT_MODEL_ID = list(models_config.keys())[0] | |
| def infer( | |
| ref_audio_orig, | |
| ref_text, | |
| gen_text, | |
| model, | |
| remove_silence=False, | |
| cross_fade_duration=0.15, | |
| nfe_step=32, | |
| speed=1, | |
| show_info=gr.Info, | |
| ): | |
| if not ref_audio_orig: | |
| gr.Warning("Please provide reference audio.") | |
| return gr.update(), gr.update(), ref_text | |
| if not gen_text.strip(): | |
| gr.Warning("Please enter text to generate.") | |
| return gr.update(), gr.update(), ref_text | |
| ref_audio, ref_text = preprocess_ref_audio_text( | |
| ref_audio_orig, ref_text, show_info=show_info | |
| ) | |
| ref_duration = AudioDecoder(ref_audio).metadata.duration_seconds_from_header | |
| ref_text_len = len(re.split(r"[, ]", ref_text)) + ref_text.count(",") | |
| gen_text_len = len(re.split(r"[, ]", gen_text)) + gen_text.count(",") | |
| target_duration = ref_duration + ref_duration * gen_text_len / ref_text_len / speed | |
| print(target_duration) | |
| final_wave, final_sample_rate, combined_spectrogram = infer_process( | |
| ref_audio, | |
| ref_text, | |
| gen_text, | |
| model, | |
| vocoder, | |
| cross_fade_duration=cross_fade_duration, | |
| nfe_step=nfe_step, | |
| speed=speed, | |
| show_info=show_info, | |
| progress=gr.Progress(), | |
| fix_duration=target_duration, | |
| ) | |
| # Remove silence | |
| if remove_silence: | |
| with tempfile.NamedTemporaryFile(delete=False, suffix=".wav") as f: | |
| sf.write(f.name, final_wave, final_sample_rate) | |
| remove_silence_for_generated_wav(f.name) | |
| final_wave = torchcodec.decoders.AudioDecoder(f.name).get_all_samples().data | |
| final_wave = final_wave.squeeze().cpu().numpy() | |
| # Save the spectrogram | |
| with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as tmp_spectrogram: | |
| spectrogram_path = tmp_spectrogram.name | |
| save_spectrogram(combined_spectrogram, spectrogram_path) | |
| return (final_sample_rate, final_wave), spectrogram_path | |
| def get_title(): | |
| with open("DEMO.md", encoding="utf-8") as tong: | |
| return tong.readline().strip("# ") | |
| with gr.Blocks( | |
| title="臺灣客語語音生成系統", | |
| ) as demo: | |
| gr.Markdown( | |
| """ | |
| # 臺灣客語語音合成系統 | |
| ### Taiwanese Hakka Text-to-Speech System | |
| ### 研發團隊 | |
| - **[李鴻欣 Hung-Shin Lee](mailto:hungshinlee@gmail.com)** | |
| - **[陳力瑋 Li-Wei Chen](mailto:wayne900619@gmail.com)** | |
| ### 合作單位 | |
| - **[國立聯合大學智慧客家實驗室](https://www.gohakka.org)** | |
| """ | |
| ) | |
| with gr.Row(): | |
| with gr.Column(): | |
| model_drop_down = gr.Dropdown( | |
| models_config.keys(), | |
| value=DEFAULT_MODEL_ID, | |
| label="模型", | |
| ) | |
| ref_dialect_radio = gr.Radio( | |
| choices=DIALECT_MAP.items(), | |
| value="客語_四縣", | |
| label="ref 腔調", | |
| ) | |
| ref_text_input = gr.Textbox( | |
| value="", | |
| label="Reference Text", | |
| ) | |
| ref_audio_input = gr.Audio( | |
| type="filepath", | |
| waveform_options=gr.WaveformOptions( | |
| sample_rate=24000, | |
| ), | |
| label="Reference Audio", | |
| ) | |
| gen_dialect_radio = gr.Radio( | |
| choices=DIALECT_MAP.items(), | |
| value="客語_四縣", | |
| label="gen 腔調", | |
| ) | |
| gen_text_input = gr.Textbox( | |
| label="Text to Generate", | |
| value="", | |
| ) | |
| generate_btn = gr.Button("Synthesize", variant="primary") | |
| with gr.Accordion("Advanced Settings", open=False): | |
| remove_silence = gr.Checkbox( | |
| label="Remove Silences", | |
| info="The model tends to produce silences, especially on longer audio. We can manually remove silences if needed. Note that this is an experimental feature and may produce strange results. This will also increase generation time.", | |
| value=False, | |
| ) | |
| speed_slider = gr.Slider( | |
| label="Speed", | |
| minimum=0.3, | |
| maximum=2.0, | |
| value=1.0, | |
| step=0.1, | |
| info="語速(越小越慢)", | |
| ) | |
| nfe_slider = gr.Slider( | |
| label="NFE Steps", | |
| minimum=4, | |
| maximum=64, | |
| value=32, | |
| step=2, | |
| info="Set the number of denoising steps.", | |
| ) | |
| cross_fade_duration_slider = gr.Slider( | |
| label="Cross-Fade Duration (s)", | |
| minimum=0.0, | |
| maximum=1.0, | |
| value=0.15, | |
| step=0.01, | |
| info="Set the duration of the cross-fade between audio clips.", | |
| ) | |
| with gr.Column(): | |
| audio_output = gr.Audio(label="Synthesized Audio") | |
| spectrogram_output = gr.Image(label="Spectrogram") | |
| def basic_tts( | |
| model_drop_down: str, | |
| ref_dialect_radio: str, | |
| ref_audio_input: str, | |
| ref_text_input: str, | |
| gen_dialect_radio: str, | |
| gen_text_input: str, | |
| remove_silence: bool, | |
| cross_fade_duration_slider: float, | |
| nfe_slider: int, | |
| speed_slider: float, | |
| ): | |
| ref_text_input = ref_text_input.strip() | |
| if len(ref_text_input) == 0: | |
| raise gr.Error("請勿輸入空字串。") | |
| gen_text_input = gen_text_input.strip() | |
| if len(gen_text_input) == 0: | |
| raise gr.Error("請勿輸入空字串。") | |
| # let text | |
| ref_text_input = to_pinyin(ref_text_input, ref_dialect_radio) | |
| print(ref_text_input) | |
| gen_text_input = to_pinyin(gen_text_input, gen_dialect_radio) | |
| print(gen_text_input) | |
| audio_out, spectrogram_path = infer( | |
| ref_audio_input, | |
| ref_text_input, | |
| gen_text_input, | |
| models_config[model_drop_down], | |
| remove_silence, | |
| cross_fade_duration=cross_fade_duration_slider, | |
| nfe_step=nfe_slider, | |
| speed=speed_slider, | |
| ) | |
| return audio_out, spectrogram_path | |
| generate_btn.click( | |
| basic_tts, | |
| inputs=[ | |
| model_drop_down, | |
| ref_dialect_radio, | |
| ref_audio_input, | |
| ref_text_input, | |
| gen_dialect_radio, | |
| gen_text_input, | |
| remove_silence, | |
| cross_fade_duration_slider, | |
| nfe_slider, | |
| speed_slider, | |
| ], | |
| outputs=[audio_output, spectrogram_output], | |
| ) | |
| gr.Examples( | |
| [ | |
| [ | |
| "./ref_wav/0000001_0.15-0.93.wav", | |
| "恁早。", | |
| "客語_四縣", | |
| "食飯愛正經食,正毋會食到半出半入。", | |
| "客語_四縣", | |
| ], | |
| [ | |
| "./ref_wav/0000002_0.15-2.73.wav", | |
| "你今晡日著到恁派頭。", | |
| "客語_四縣", | |
| "食飯愛正經食,正毋會食到半出半入。", | |
| "客語_四縣", | |
| ], | |
| [ | |
| "./ref_wav/0000002_0.15-2.73.wav", | |
| "你今晡日著到恁派頭。", | |
| "客語_四縣", | |
| "歸條路吊等長長个花燈,祈求風調雨順,歸屋下人个心願,親像花燈下燒暖个光華。", | |
| "客語_四縣", | |
| ], | |
| ], | |
| label="範例", | |
| inputs=[ | |
| ref_audio_input, | |
| ref_text_input, | |
| ref_dialect_radio, | |
| gen_text_input, | |
| gen_dialect_radio, | |
| ], | |
| ) | |
| demo.launch( | |
| css="@import url(https://tauhu.tw/tauhu-oo.css);", | |
| theme=gr.themes.Default( | |
| font=( | |
| "tauhu-oo", | |
| gr.themes.GoogleFont("Source Sans Pro"), | |
| "ui-sans-serif", | |
| "system-ui", | |
| "sans-serif", | |
| ) | |
| ), | |
| ) | |