Feat: give tone for sythesize (WIP for better UX)
This commit is contained in:
9
app.py
9
app.py
@@ -46,9 +46,13 @@ def tts_fn(
|
|||||||
use_assist_text,
|
use_assist_text,
|
||||||
style,
|
style,
|
||||||
style_weight,
|
style_weight,
|
||||||
|
given_tone,
|
||||||
):
|
):
|
||||||
assert model_holder.current_model is not None
|
assert model_holder.current_model is not None
|
||||||
|
if given_tone == "":
|
||||||
|
given_tone = None
|
||||||
|
else:
|
||||||
|
given_tone = [int(i) for i in given_tone]
|
||||||
start_time = datetime.datetime.now()
|
start_time = datetime.datetime.now()
|
||||||
|
|
||||||
sr, audio = model_holder.current_model.infer(
|
sr, audio = model_holder.current_model.infer(
|
||||||
@@ -66,6 +70,7 @@ def tts_fn(
|
|||||||
use_assist_text=use_assist_text,
|
use_assist_text=use_assist_text,
|
||||||
style=style,
|
style=style,
|
||||||
style_weight=style_weight,
|
style_weight=style_weight,
|
||||||
|
given_tone=given_tone,
|
||||||
)
|
)
|
||||||
|
|
||||||
end_time = datetime.datetime.now()
|
end_time = datetime.datetime.now()
|
||||||
@@ -238,6 +243,7 @@ if __name__ == "__main__":
|
|||||||
step=0.1,
|
step=0.1,
|
||||||
label="分けた場合に挟む無音の長さ(秒)",
|
label="分けた場合に挟む無音の長さ(秒)",
|
||||||
)
|
)
|
||||||
|
given_tone = gr.Textbox("トーン、0と1の数値列")
|
||||||
language = gr.Dropdown(choices=languages, value="JP", label="Language")
|
language = gr.Dropdown(choices=languages, value="JP", label="Language")
|
||||||
with gr.Accordion(label="詳細設定", open=False):
|
with gr.Accordion(label="詳細設定", open=False):
|
||||||
sdp_ratio = gr.Slider(
|
sdp_ratio = gr.Slider(
|
||||||
@@ -334,6 +340,7 @@ if __name__ == "__main__":
|
|||||||
use_assist_text,
|
use_assist_text,
|
||||||
style,
|
style,
|
||||||
style_weight,
|
style_weight,
|
||||||
|
given_tone,
|
||||||
],
|
],
|
||||||
outputs=[text_output, audio_output],
|
outputs=[text_output, audio_output],
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -95,6 +95,7 @@ class Model:
|
|||||||
use_assist_text: bool = False,
|
use_assist_text: bool = False,
|
||||||
style: str = DEFAULT_STYLE,
|
style: str = DEFAULT_STYLE,
|
||||||
style_weight: float = DEFAULT_STYLE_WEIGHT,
|
style_weight: float = DEFAULT_STYLE_WEIGHT,
|
||||||
|
given_tone: Optional[list[int]] = None,
|
||||||
) -> tuple[int, np.ndarray]:
|
) -> tuple[int, np.ndarray]:
|
||||||
logger.info(f"Start generating audio data from text:\n{text}")
|
logger.info(f"Start generating audio data from text:\n{text}")
|
||||||
if reference_audio_path == "":
|
if reference_audio_path == "":
|
||||||
@@ -127,6 +128,7 @@ class Model:
|
|||||||
assist_text=assist_text,
|
assist_text=assist_text,
|
||||||
assist_text_weight=assist_text_weight,
|
assist_text_weight=assist_text_weight,
|
||||||
style_vec=style_vector,
|
style_vec=style_vector,
|
||||||
|
given_tone=given_tone,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
texts = text.split("\n")
|
texts = text.split("\n")
|
||||||
|
|||||||
@@ -17,9 +17,10 @@
|
|||||||
"c_mel": 45,
|
"c_mel": 45,
|
||||||
"c_kl": 1.0,
|
"c_kl": 1.0,
|
||||||
"skip_optimizer": false,
|
"skip_optimizer": false,
|
||||||
"freeze_ZH_bert": false,
|
"freeze_ZH_bert": true,
|
||||||
"freeze_JP_bert": false,
|
"freeze_JP_bert": true,
|
||||||
"freeze_EN_bert": false
|
"freeze_EN_bert": true,
|
||||||
|
"freeze_style": true
|
||||||
},
|
},
|
||||||
"data": {
|
"data": {
|
||||||
"training_files": "Data/your_model_name/filelists/train.list",
|
"training_files": "Data/your_model_name/filelists/train.list",
|
||||||
|
|||||||
17
infer.py
17
infer.py
@@ -6,6 +6,7 @@ from models import SynthesizerTrn
|
|||||||
from text import cleaned_text_to_sequence, get_bert
|
from text import cleaned_text_to_sequence, get_bert
|
||||||
from text.cleaner import clean_text
|
from text.cleaner import clean_text
|
||||||
from text.symbols import symbols
|
from text.symbols import symbols
|
||||||
|
from common.log import logger
|
||||||
|
|
||||||
# latest_version = "1.0"
|
# latest_version = "1.0"
|
||||||
|
|
||||||
@@ -29,9 +30,21 @@ def get_net_g(model_path: str, version: str, device: str, hps):
|
|||||||
return net_g
|
return net_g
|
||||||
|
|
||||||
|
|
||||||
def get_text(text, language_str, hps, device, assist_text=None, assist_text_weight=0.7):
|
def get_text(
|
||||||
|
text,
|
||||||
|
language_str,
|
||||||
|
hps,
|
||||||
|
device,
|
||||||
|
assist_text=None,
|
||||||
|
assist_text_weight=0.7,
|
||||||
|
given_tone=None,
|
||||||
|
):
|
||||||
# 在此处实现当前版本的get_text
|
# 在此处实现当前版本的get_text
|
||||||
norm_text, phone, tone, word2ph = clean_text(text, language_str)
|
norm_text, phone, tone, word2ph = clean_text(text, language_str)
|
||||||
|
logger.info(f"Original tone: {''.join(str(num) for num in tone)}")
|
||||||
|
if given_tone is not None:
|
||||||
|
logger.debug(f"Tone given: {given_tone}")
|
||||||
|
tone = given_tone
|
||||||
phone, tone, language = cleaned_text_to_sequence(phone, tone, language_str)
|
phone, tone, language = cleaned_text_to_sequence(phone, tone, language_str)
|
||||||
|
|
||||||
if hps.data.add_blank:
|
if hps.data.add_blank:
|
||||||
@@ -88,6 +101,7 @@ def infer(
|
|||||||
skip_end=False,
|
skip_end=False,
|
||||||
assist_text=None,
|
assist_text=None,
|
||||||
assist_text_weight=0.7,
|
assist_text_weight=0.7,
|
||||||
|
given_tone=None,
|
||||||
):
|
):
|
||||||
bert, ja_bert, en_bert, phones, tones, lang_ids = get_text(
|
bert, ja_bert, en_bert, phones, tones, lang_ids = get_text(
|
||||||
text,
|
text,
|
||||||
@@ -96,6 +110,7 @@ def infer(
|
|||||||
device,
|
device,
|
||||||
assist_text=assist_text,
|
assist_text=assist_text,
|
||||||
assist_text_weight=assist_text_weight,
|
assist_text_weight=assist_text_weight,
|
||||||
|
given_tone=given_tone,
|
||||||
)
|
)
|
||||||
if skip_start:
|
if skip_start:
|
||||||
phones = phones[3:]
|
phones = phones[3:]
|
||||||
|
|||||||
@@ -202,6 +202,7 @@ def g2phone_tone_list(text: str) -> list[tuple[str, int]]:
|
|||||||
[('k', 0), ('o', 0), ('n', 1), ('n', 1), ('i', 1), ('ch', 1), ('i', 1), ('w', 1), ('a', 1), ('s', 1), ('e', 1), ('k', 0), ('a', 0), ('i', 0), ('i', 0), ('g', 1), ('e', 1), ('n', 0), ('k', 0), ('i', 0)]
|
[('k', 0), ('o', 0), ('n', 1), ('n', 1), ('i', 1), ('ch', 1), ('i', 1), ('w', 1), ('a', 1), ('s', 1), ('e', 1), ('k', 0), ('a', 0), ('i', 0), ('i', 0), ('g', 1), ('e', 1), ('n', 0), ('k', 0), ('i', 0)]
|
||||||
"""
|
"""
|
||||||
prosodies = pyopenjtalk_g2p_prosody(text, drop_unvoiced_vowels=True)
|
prosodies = pyopenjtalk_g2p_prosody(text, drop_unvoiced_vowels=True)
|
||||||
|
logger.debug(f"prosodies: {prosodies}")
|
||||||
result: list[tuple[str, int]] = []
|
result: list[tuple[str, int]] = []
|
||||||
current_phrase: list[tuple[str, int]] = []
|
current_phrase: list[tuple[str, int]] = []
|
||||||
current_tone = 0
|
current_tone = 0
|
||||||
|
|||||||
@@ -258,6 +258,10 @@ def run():
|
|||||||
logger.info("Freezing JP bert encoder !!!")
|
logger.info("Freezing JP bert encoder !!!")
|
||||||
for param in net_g.enc_p.ja_bert_proj.parameters():
|
for param in net_g.enc_p.ja_bert_proj.parameters():
|
||||||
param.requires_grad = False
|
param.requires_grad = False
|
||||||
|
if getattr(hps.train, "freeze_style", False):
|
||||||
|
logger.info("Freezing style encoder !!!")
|
||||||
|
for param in net_g.enc_p.style_proj.parameters():
|
||||||
|
param.requires_grad = False
|
||||||
|
|
||||||
net_d = MultiPeriodDiscriminator(hps.model.use_spectral_norm).cuda(local_rank)
|
net_d = MultiPeriodDiscriminator(hps.model.use_spectral_norm).cuda(local_rank)
|
||||||
optim_g = torch.optim.AdamW(
|
optim_g = torch.optim.AdamW(
|
||||||
|
|||||||
Reference in New Issue
Block a user