txya900619's picture
feat: use new g2p tool to replace old pinyin conversion tool
e7dce34
Raw History Blame Contribute Delete
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]
@gpu_decorator
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")
@gpu_decorator
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",
)
),
)