Merge branch 'dev' of github.com:litagin02/Style-Bert-VITS2
This commit is contained in:
1
.gitignore
vendored
1
.gitignore
vendored
@@ -44,3 +44,4 @@ safetensors.ipynb
|
|||||||
*.dic
|
*.dic
|
||||||
|
|
||||||
playground.ipynb
|
playground.ipynb
|
||||||
|
playgrounds/
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
import datetime
|
import datetime
|
||||||
import json
|
import json
|
||||||
from typing import Optional
|
from typing import Any, Optional, Union
|
||||||
|
|
||||||
import gradio as gr
|
import gradio as gr
|
||||||
|
|
||||||
@@ -186,6 +186,10 @@ style_md = f"""
|
|||||||
- どのくらいに強さがいいかはモデルやスタイルによって異なるようです。
|
- どのくらいに強さがいいかはモデルやスタイルによって異なるようです。
|
||||||
- 音声ファイルを入力する場合は、学習データと似た声音の話者(特に同じ性別)でないとよい効果が出ないかもしれません。
|
- 音声ファイルを入力する場合は、学習データと似た声音の話者(特に同じ性別)でないとよい効果が出ないかもしれません。
|
||||||
"""
|
"""
|
||||||
|
voice_keys = ["dec"]
|
||||||
|
voice_pitch_keys = ["flow"]
|
||||||
|
speech_style_keys = ["enc_p"]
|
||||||
|
tempo_keys = ["sdp", "dp"]
|
||||||
|
|
||||||
|
|
||||||
def make_interactive():
|
def make_interactive():
|
||||||
@@ -203,6 +207,35 @@ def gr_util(item):
|
|||||||
return (gr.update(visible=False), gr.update(visible=True))
|
return (gr.update(visible=False), gr.update(visible=True))
|
||||||
|
|
||||||
|
|
||||||
|
null_models_frame = 0
|
||||||
|
|
||||||
|
|
||||||
|
def change_null_model_row(
|
||||||
|
null_model_index: int,
|
||||||
|
null_model_name: str,
|
||||||
|
null_model_path: str,
|
||||||
|
null_voice_weights: float,
|
||||||
|
null_voice_pitch_weights: float,
|
||||||
|
null_speech_style_weights: float,
|
||||||
|
null_tempo_weights: float,
|
||||||
|
null_models: dict[int, dict[str, Any]],
|
||||||
|
):
|
||||||
|
null_models[null_model_index] = {
|
||||||
|
"name": null_model_name,
|
||||||
|
"path": null_model_path,
|
||||||
|
"weight": null_voice_weights,
|
||||||
|
"pitch": null_voice_pitch_weights,
|
||||||
|
"style": null_speech_style_weights,
|
||||||
|
"tempo": null_tempo_weights,
|
||||||
|
}
|
||||||
|
if len(null_models) > null_models_frame:
|
||||||
|
keys_to_keep = list(range(null_models_frame))
|
||||||
|
result = {k: null_models[k] for k in keys_to_keep}
|
||||||
|
else:
|
||||||
|
result = null_models
|
||||||
|
return result, True
|
||||||
|
|
||||||
|
|
||||||
def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks:
|
def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks:
|
||||||
def tts_fn(
|
def tts_fn(
|
||||||
model_name,
|
model_name,
|
||||||
@@ -226,9 +259,12 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks:
|
|||||||
speaker,
|
speaker,
|
||||||
pitch_scale,
|
pitch_scale,
|
||||||
intonation_scale,
|
intonation_scale,
|
||||||
|
null_models: dict[int, dict[str, Union[str, float]]],
|
||||||
|
force_reload_model: bool,
|
||||||
):
|
):
|
||||||
model_holder.get_model(model_name, model_path)
|
model_holder.get_model(model_name, model_path)
|
||||||
assert model_holder.current_model is not None
|
assert model_holder.current_model is not None
|
||||||
|
logger.debug(f"Null models setting: {null_models}")
|
||||||
|
|
||||||
wrong_tone_message = ""
|
wrong_tone_message = ""
|
||||||
kata_tone: Optional[list[tuple[str, int]]] = None
|
kata_tone: Optional[list[tuple[str, int]]] = None
|
||||||
@@ -283,6 +319,8 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks:
|
|||||||
speaker_id=speaker_id,
|
speaker_id=speaker_id,
|
||||||
pitch_scale=pitch_scale,
|
pitch_scale=pitch_scale,
|
||||||
intonation_scale=intonation_scale,
|
intonation_scale=intonation_scale,
|
||||||
|
null_model_params=null_models,
|
||||||
|
force_reload_model=force_reload_model,
|
||||||
)
|
)
|
||||||
except InvalidToneError as e:
|
except InvalidToneError as e:
|
||||||
logger.error(f"Tone error: {e}")
|
logger.error(f"Tone error: {e}")
|
||||||
@@ -304,7 +342,10 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks:
|
|||||||
message = f"Success, time: {duration} seconds."
|
message = f"Success, time: {duration} seconds."
|
||||||
if wrong_tone_message != "":
|
if wrong_tone_message != "":
|
||||||
message = wrong_tone_message + "\n" + message
|
message = wrong_tone_message + "\n" + message
|
||||||
return message, (sr, audio), kata_tone_json_str
|
return message, (sr, audio), kata_tone_json_str, False
|
||||||
|
|
||||||
|
def get_model_files(model_name: str):
|
||||||
|
return [str(f) for f in model_holder.model_files_dict[model_name]]
|
||||||
|
|
||||||
model_names = model_holder.model_names
|
model_names = model_holder.model_names
|
||||||
if len(model_names) == 0:
|
if len(model_names) == 0:
|
||||||
@@ -317,13 +358,13 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks:
|
|||||||
)
|
)
|
||||||
return app
|
return app
|
||||||
initial_id = 0
|
initial_id = 0
|
||||||
initial_pth_files = [
|
initial_pth_files = get_model_files(model_names[initial_id])
|
||||||
str(f) for f in model_holder.model_files_dict[model_names[initial_id]]
|
|
||||||
]
|
|
||||||
|
|
||||||
with gr.Blocks(theme=GRADIO_THEME) as app:
|
with gr.Blocks(theme=GRADIO_THEME) as app:
|
||||||
gr.Markdown(initial_md)
|
gr.Markdown(initial_md)
|
||||||
gr.Markdown(terms_of_use_md)
|
gr.Markdown(terms_of_use_md)
|
||||||
|
null_models = gr.State({})
|
||||||
|
force_reload_model = gr.State(False)
|
||||||
with gr.Accordion(label="使い方", open=False):
|
with gr.Accordion(label="使い方", open=False):
|
||||||
gr.Markdown(how_to_md)
|
gr.Markdown(how_to_md)
|
||||||
with gr.Row():
|
with gr.Row():
|
||||||
@@ -437,6 +478,166 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks:
|
|||||||
inputs=[use_assist_text],
|
inputs=[use_assist_text],
|
||||||
outputs=[assist_text, assist_text_weight],
|
outputs=[assist_text, assist_text_weight],
|
||||||
)
|
)
|
||||||
|
with gr.Accordion(label="ヌルモデル", open=False):
|
||||||
|
with gr.Row():
|
||||||
|
null_models_count = gr.Number(
|
||||||
|
label="ヌルモデルの数", value=0, step=1
|
||||||
|
)
|
||||||
|
with gr.Column(variant="panel"):
|
||||||
|
|
||||||
|
@gr.render(inputs=[null_models_count])
|
||||||
|
def render_null_models(
|
||||||
|
null_models_count: int,
|
||||||
|
):
|
||||||
|
global null_models_frame
|
||||||
|
null_models_frame = null_models_count
|
||||||
|
for i in range(null_models_count):
|
||||||
|
with gr.Row():
|
||||||
|
null_model_index = gr.Number(
|
||||||
|
value=i,
|
||||||
|
key=f"null_model_index_{i}",
|
||||||
|
visible=False,
|
||||||
|
)
|
||||||
|
null_model_name = gr.Dropdown(
|
||||||
|
label="モデル一覧",
|
||||||
|
choices=model_names,
|
||||||
|
key=f"null_model_name_{i}",
|
||||||
|
value=model_names[initial_id],
|
||||||
|
)
|
||||||
|
null_model_path = gr.Dropdown(
|
||||||
|
label="モデルファイル",
|
||||||
|
key=f"null_model_path_{i}",
|
||||||
|
# FIXME: 再レンダー時に選択肢が消えるのでどうにかしたい
|
||||||
|
# 現在は再レンダーでvalueは保存されるが選択肢は保存されないので選択肢が空になる
|
||||||
|
# そのときに選択肢にない値となるので、それを許す
|
||||||
|
allow_custom_value=True,
|
||||||
|
)
|
||||||
|
null_voice_weights = gr.Slider(
|
||||||
|
minimum=0,
|
||||||
|
maximum=1,
|
||||||
|
value=1,
|
||||||
|
step=0.1,
|
||||||
|
key=f"null_voice_weights_{i}",
|
||||||
|
label="声質",
|
||||||
|
)
|
||||||
|
null_voice_pitch_weights = gr.Slider(
|
||||||
|
minimum=0,
|
||||||
|
maximum=1,
|
||||||
|
value=1,
|
||||||
|
step=0.1,
|
||||||
|
key=f"null_voice_pitch_weights_{i}",
|
||||||
|
label="声の高さ",
|
||||||
|
)
|
||||||
|
null_speech_style_weights = gr.Slider(
|
||||||
|
minimum=0,
|
||||||
|
maximum=1,
|
||||||
|
value=1,
|
||||||
|
step=0.1,
|
||||||
|
key=f"null_speech_style_weights_{i}",
|
||||||
|
label="話し方",
|
||||||
|
)
|
||||||
|
null_tempo_weights = gr.Slider(
|
||||||
|
minimum=0,
|
||||||
|
maximum=1,
|
||||||
|
value=1,
|
||||||
|
step=0.1,
|
||||||
|
key=f"null_tempo_weights_{i}",
|
||||||
|
label="テンポ",
|
||||||
|
)
|
||||||
|
|
||||||
|
null_model_name.change(
|
||||||
|
model_holder.update_model_files_for_gradio,
|
||||||
|
inputs=[null_model_name],
|
||||||
|
outputs=[null_model_path],
|
||||||
|
)
|
||||||
|
null_model_path.change(
|
||||||
|
make_non_interactive, outputs=[tts_button]
|
||||||
|
)
|
||||||
|
# 愚直すぎるのでもう少しなんとかしたい
|
||||||
|
null_model_path.change(
|
||||||
|
change_null_model_row,
|
||||||
|
inputs=[
|
||||||
|
null_model_index,
|
||||||
|
null_model_name,
|
||||||
|
null_model_path,
|
||||||
|
null_voice_weights,
|
||||||
|
null_voice_pitch_weights,
|
||||||
|
null_speech_style_weights,
|
||||||
|
null_tempo_weights,
|
||||||
|
null_models,
|
||||||
|
],
|
||||||
|
outputs=[null_models, force_reload_model],
|
||||||
|
)
|
||||||
|
null_voice_weights.change(
|
||||||
|
change_null_model_row,
|
||||||
|
inputs=[
|
||||||
|
null_model_index,
|
||||||
|
null_model_name,
|
||||||
|
null_model_path,
|
||||||
|
null_voice_weights,
|
||||||
|
null_voice_pitch_weights,
|
||||||
|
null_speech_style_weights,
|
||||||
|
null_tempo_weights,
|
||||||
|
null_models,
|
||||||
|
],
|
||||||
|
outputs=[null_models, force_reload_model],
|
||||||
|
)
|
||||||
|
null_voice_pitch_weights.change(
|
||||||
|
change_null_model_row,
|
||||||
|
inputs=[
|
||||||
|
null_model_index,
|
||||||
|
null_model_name,
|
||||||
|
null_model_path,
|
||||||
|
null_voice_weights,
|
||||||
|
null_voice_pitch_weights,
|
||||||
|
null_speech_style_weights,
|
||||||
|
null_tempo_weights,
|
||||||
|
null_models,
|
||||||
|
],
|
||||||
|
outputs=[null_models, force_reload_model],
|
||||||
|
)
|
||||||
|
null_speech_style_weights.change(
|
||||||
|
change_null_model_row,
|
||||||
|
inputs=[
|
||||||
|
null_model_index,
|
||||||
|
null_model_name,
|
||||||
|
null_model_path,
|
||||||
|
null_voice_weights,
|
||||||
|
null_voice_pitch_weights,
|
||||||
|
null_speech_style_weights,
|
||||||
|
null_tempo_weights,
|
||||||
|
null_models,
|
||||||
|
],
|
||||||
|
outputs=[null_models, force_reload_model],
|
||||||
|
)
|
||||||
|
null_tempo_weights.change(
|
||||||
|
change_null_model_row,
|
||||||
|
inputs=[
|
||||||
|
null_model_index,
|
||||||
|
null_model_name,
|
||||||
|
null_model_path,
|
||||||
|
null_voice_weights,
|
||||||
|
null_voice_pitch_weights,
|
||||||
|
null_speech_style_weights,
|
||||||
|
null_tempo_weights,
|
||||||
|
null_models,
|
||||||
|
],
|
||||||
|
outputs=[null_models, force_reload_model],
|
||||||
|
)
|
||||||
|
|
||||||
|
add_btn = gr.Button("ヌルモデルを増やす")
|
||||||
|
del_btn = gr.Button("ヌルモデルを減らす")
|
||||||
|
add_btn.click(
|
||||||
|
lambda x: x + 1,
|
||||||
|
inputs=[null_models_count],
|
||||||
|
outputs=[null_models_count],
|
||||||
|
)
|
||||||
|
del_btn.click(
|
||||||
|
lambda x: x - 1 if x > 0 else 0,
|
||||||
|
inputs=[null_models_count],
|
||||||
|
outputs=[null_models_count],
|
||||||
|
)
|
||||||
|
|
||||||
with gr.Column():
|
with gr.Column():
|
||||||
with gr.Accordion("スタイルについて詳細", open=False):
|
with gr.Accordion("スタイルについて詳細", open=False):
|
||||||
gr.Markdown(style_md)
|
gr.Markdown(style_md)
|
||||||
@@ -494,8 +695,10 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks:
|
|||||||
speaker,
|
speaker,
|
||||||
pitch_scale,
|
pitch_scale,
|
||||||
intonation_scale,
|
intonation_scale,
|
||||||
|
null_models,
|
||||||
|
force_reload_model,
|
||||||
],
|
],
|
||||||
outputs=[text_output, audio_output, tone],
|
outputs=[text_output, audio_output, tone, force_reload_model],
|
||||||
)
|
)
|
||||||
|
|
||||||
model_name.change(
|
model_name.change(
|
||||||
|
|||||||
@@ -105,6 +105,8 @@ def merge_style_usual(
|
|||||||
new_config = config_a.copy()
|
new_config = config_a.copy()
|
||||||
new_config["data"]["num_styles"] = len(new_style2id)
|
new_config["data"]["num_styles"] = len(new_style2id)
|
||||||
new_config["data"]["style2id"] = new_style2id
|
new_config["data"]["style2id"] = new_style2id
|
||||||
|
if new_config["data"]["n_speakers"] == 1:
|
||||||
|
new_config["data"]["spk2id"] = { output_name : 0}
|
||||||
new_config["model_name"] = output_name
|
new_config["model_name"] = output_name
|
||||||
save_config(new_config, output_name)
|
save_config(new_config, output_name)
|
||||||
|
|
||||||
@@ -160,6 +162,8 @@ def merge_style_add_diff(
|
|||||||
new_config = config_a.copy()
|
new_config = config_a.copy()
|
||||||
new_config["data"]["num_styles"] = len(new_style2id)
|
new_config["data"]["num_styles"] = len(new_style2id)
|
||||||
new_config["data"]["style2id"] = new_style2id
|
new_config["data"]["style2id"] = new_style2id
|
||||||
|
if new_config["data"]["n_speakers"] == 1:
|
||||||
|
new_config["data"]["spk2id"] = { output_name : 0}
|
||||||
new_config["model_name"] = output_name
|
new_config["model_name"] = output_name
|
||||||
save_config(new_config, output_name)
|
save_config(new_config, output_name)
|
||||||
|
|
||||||
@@ -219,6 +223,8 @@ def merge_style_weighted_sum(
|
|||||||
new_config = config_a.copy()
|
new_config = config_a.copy()
|
||||||
new_config["data"]["num_styles"] = len(new_style2id)
|
new_config["data"]["num_styles"] = len(new_style2id)
|
||||||
new_config["data"]["style2id"] = new_style2id
|
new_config["data"]["style2id"] = new_style2id
|
||||||
|
if new_config["data"]["n_speakers"] == 1:
|
||||||
|
new_config["data"]["spk2id"] = { output_name : 0}
|
||||||
new_config["model_name"] = output_name
|
new_config["model_name"] = output_name
|
||||||
save_config(new_config, output_name)
|
save_config(new_config, output_name)
|
||||||
|
|
||||||
@@ -268,6 +274,8 @@ def merge_style_add_null(
|
|||||||
new_config = config_a.copy()
|
new_config = config_a.copy()
|
||||||
new_config["data"]["num_styles"] = len(new_style2id)
|
new_config["data"]["num_styles"] = len(new_style2id)
|
||||||
new_config["data"]["style2id"] = new_style2id
|
new_config["data"]["style2id"] = new_style2id
|
||||||
|
if new_config["data"]["n_speakers"] == 1:
|
||||||
|
new_config["data"]["spk2id"] = { output_name : 0}
|
||||||
new_config["model_name"] = output_name
|
new_config["model_name"] = output_name
|
||||||
save_config(new_config, output_name)
|
save_config(new_config, output_name)
|
||||||
|
|
||||||
@@ -362,6 +370,8 @@ def merge_models_usual(
|
|||||||
new_config["model_name"] = output_name
|
new_config["model_name"] = output_name
|
||||||
new_config["data"]["num_styles"] = 1
|
new_config["data"]["num_styles"] = 1
|
||||||
new_config["data"]["style2id"] = {DEFAULT_STYLE: 0}
|
new_config["data"]["style2id"] = {DEFAULT_STYLE: 0}
|
||||||
|
if new_config["data"]["n_speakers"] == 1:
|
||||||
|
new_config["data"]["spk2id"] = { output_name : 0}
|
||||||
save_config(new_config, output_name)
|
save_config(new_config, output_name)
|
||||||
|
|
||||||
neutral_vector_a = style_vectors_a[0]
|
neutral_vector_a = style_vectors_a[0]
|
||||||
@@ -444,6 +454,8 @@ def merge_models_add_diff(
|
|||||||
new_config["model_name"] = output_name
|
new_config["model_name"] = output_name
|
||||||
new_config["data"]["num_styles"] = 1
|
new_config["data"]["num_styles"] = 1
|
||||||
new_config["data"]["style2id"] = {DEFAULT_STYLE: 0}
|
new_config["data"]["style2id"] = {DEFAULT_STYLE: 0}
|
||||||
|
if new_config["data"]["n_speakers"] == 1:
|
||||||
|
new_config["data"]["spk2id"] = { output_name : 0}
|
||||||
with open(assets_root / output_name / "config.json", "w", encoding="utf-8") as f:
|
with open(assets_root / output_name / "config.json", "w", encoding="utf-8") as f:
|
||||||
json.dump(new_config, f, indent=2, ensure_ascii=False)
|
json.dump(new_config, f, indent=2, ensure_ascii=False)
|
||||||
|
|
||||||
@@ -519,6 +531,8 @@ def merge_models_weighted_sum(
|
|||||||
new_config["model_name"] = output_name
|
new_config["model_name"] = output_name
|
||||||
new_config["data"]["num_styles"] = 1
|
new_config["data"]["num_styles"] = 1
|
||||||
new_config["data"]["style2id"] = {DEFAULT_STYLE: 0}
|
new_config["data"]["style2id"] = {DEFAULT_STYLE: 0}
|
||||||
|
if new_config["data"]["n_speakers"] == 1:
|
||||||
|
new_config["data"]["spk2id"] = { output_name : 0}
|
||||||
with open(assets_root / output_name / "config.json", "w", encoding="utf-8") as f:
|
with open(assets_root / output_name / "config.json", "w", encoding="utf-8") as f:
|
||||||
json.dump(new_config, f, indent=2, ensure_ascii=False)
|
json.dump(new_config, f, indent=2, ensure_ascii=False)
|
||||||
|
|
||||||
@@ -595,6 +609,8 @@ def merge_models_add_null(
|
|||||||
new_config["model_name"] = output_name
|
new_config["model_name"] = output_name
|
||||||
new_config["data"]["num_styles"] = 1
|
new_config["data"]["num_styles"] = 1
|
||||||
new_config["data"]["style2id"] = {DEFAULT_STYLE: 0}
|
new_config["data"]["style2id"] = {DEFAULT_STYLE: 0}
|
||||||
|
if new_config["data"]["n_speakers"] == 1:
|
||||||
|
new_config["data"]["spk2id"] = { output_name : 0}
|
||||||
with open(assets_root / output_name / "config.json", "w", encoding="utf-8") as f:
|
with open(assets_root / output_name / "config.json", "w", encoding="utf-8") as f:
|
||||||
json.dump(new_config, f, indent=2, ensure_ascii=False)
|
json.dump(new_config, f, indent=2, ensure_ascii=False)
|
||||||
|
|
||||||
|
|||||||
@@ -108,8 +108,9 @@ class TTSModel:
|
|||||||
)
|
)
|
||||||
self.style_vector_inference: Optional[Any] = None
|
self.style_vector_inference: Optional[Any] = None
|
||||||
|
|
||||||
# net_g は PyTorch 推論時のみ遅延初期化される
|
# net_g / null_model_params は PyTorch 推論時のみ遅延初期化される
|
||||||
self.net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra, None] = None
|
self.net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra, None] = None
|
||||||
|
self.null_model_params: dict[int, dict[str, Union[float, str]]] = {}
|
||||||
|
|
||||||
# onnx_session は ONNX 推論時のみ遅延初期化される
|
# onnx_session は ONNX 推論時のみ遅延初期化される
|
||||||
self.onnx_session: Optional[onnxruntime.InferenceSession] = None
|
self.onnx_session: Optional[onnxruntime.InferenceSession] = None
|
||||||
@@ -135,6 +136,44 @@ class TTSModel:
|
|||||||
f'Model loaded successfully from {self.model_path} to "{self.device}" device ({time.time() - start_time:.2f}s)'
|
f'Model loaded successfully from {self.model_path} to "{self.device}" device ({time.time() - start_time:.2f}s)'
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# ここからはヌルモデルのロード用パラメータが指定されている場合のみ
|
||||||
|
if len(self.null_model_params.keys()) == 0:
|
||||||
|
return
|
||||||
|
|
||||||
|
# 推論対象のモデルの重みとヌルモデルの重みをマージ
|
||||||
|
for null_model_info in self.null_model_params.values():
|
||||||
|
logger.info(f"Adding null model: {null_model_info['path']}...")
|
||||||
|
null_model_add = get_net_g(
|
||||||
|
model_path=str(null_model_info["path"]),
|
||||||
|
version=self.hyper_parameters.version,
|
||||||
|
device=self.device,
|
||||||
|
hps=self.hyper_parameters,
|
||||||
|
)
|
||||||
|
# 愚直。もっと上手い方法ありそう
|
||||||
|
params = zip(self.net_g.dec.parameters(), null_model_add.dec.parameters())
|
||||||
|
for v in params:
|
||||||
|
v[0].data.add_(v[1].data, alpha=float(null_model_info["weight"]))
|
||||||
|
params = zip(
|
||||||
|
self.net_g.flow.parameters(), null_model_add.flow.parameters()
|
||||||
|
)
|
||||||
|
for v in params:
|
||||||
|
v[0].data.add_(v[1].data, alpha=float(null_model_info["pitch"]))
|
||||||
|
|
||||||
|
params = zip(
|
||||||
|
self.net_g.enc_p.parameters(), null_model_add.enc_p.parameters()
|
||||||
|
)
|
||||||
|
for v in params:
|
||||||
|
v[0].data.add_(v[1].data, alpha=float(null_model_info["style"]))
|
||||||
|
# テンポは sdp と dp 二つあるからとりあえずどっちも足す
|
||||||
|
params = zip(self.net_g.sdp.parameters(), null_model_add.sdp.parameters())
|
||||||
|
for v in params:
|
||||||
|
v[0].data.add_(v[1].data, alpha=float(null_model_info["tempo"]))
|
||||||
|
params = zip(self.net_g.dp.parameters(), null_model_add.dp.parameters())
|
||||||
|
for v in params:
|
||||||
|
v[0].data.add_(v[1].data, alpha=float(null_model_info["tempo"]))
|
||||||
|
|
||||||
|
logger.info(f"Null models merged successfully ({time.time() - start_time:.2f}s)")
|
||||||
|
|
||||||
# ONNX 推論時
|
# ONNX 推論時
|
||||||
else:
|
else:
|
||||||
sess_options = onnxruntime.SessionOptions()
|
sess_options = onnxruntime.SessionOptions()
|
||||||
@@ -289,6 +328,8 @@ class TTSModel:
|
|||||||
given_tone: Optional[list[int]] = None,
|
given_tone: Optional[list[int]] = None,
|
||||||
pitch_scale: float = 1.0,
|
pitch_scale: float = 1.0,
|
||||||
intonation_scale: float = 1.0,
|
intonation_scale: float = 1.0,
|
||||||
|
null_model_params: dict[int, dict[str, Union[str, float]]] = {},
|
||||||
|
force_reload_model: bool = False,
|
||||||
) -> tuple[int, NDArray[Any]]:
|
) -> tuple[int, NDArray[Any]]:
|
||||||
"""
|
"""
|
||||||
テキストから音声を合成する。
|
テキストから音声を合成する。
|
||||||
@@ -313,7 +354,8 @@ class TTSModel:
|
|||||||
given_tone (Optional[list[int]], optional): アクセントのトーンのリスト. Defaults to None.
|
given_tone (Optional[list[int]], optional): アクセントのトーンのリスト. Defaults to None.
|
||||||
pitch_scale (float, optional): ピッチの高さ (1.0 から変更すると若干音質が低下する). Defaults to 1.0.
|
pitch_scale (float, optional): ピッチの高さ (1.0 から変更すると若干音質が低下する). Defaults to 1.0.
|
||||||
intonation_scale (float, optional): 抑揚の平均からの変化幅 (1.0 から変更すると若干音質が低下する). Defaults to 1.0.
|
intonation_scale (float, optional): 抑揚の平均からの変化幅 (1.0 から変更すると若干音質が低下する). Defaults to 1.0.
|
||||||
|
null_model_params (dict[int, dict[str, Union[str, float]]], optional): 推論時に使用するヌルモデルの名前、適用割合の dict が入った dict 。ONNX 推論では無視される。
|
||||||
|
force_reload_model (bool, optional): モデルを強制的に再ロードするかどうか. Defaults to False.
|
||||||
Returns:
|
Returns:
|
||||||
tuple[int, NDArray[Any]]: サンプリングレートと音声データ (16bit PCM)
|
tuple[int, NDArray[Any]]: サンプリングレートと音声データ (16bit PCM)
|
||||||
"""
|
"""
|
||||||
@@ -344,6 +386,15 @@ class TTSModel:
|
|||||||
|
|
||||||
from style_bert_vits2.models.infer import infer
|
from style_bert_vits2.models.infer import infer
|
||||||
|
|
||||||
|
if null_model_params is not {}:
|
||||||
|
self.null_model_params = null_model_params
|
||||||
|
else:
|
||||||
|
self.null_model_params = {}
|
||||||
|
|
||||||
|
# force_reload_model が True のとき、メモリ上に保持されているモデルを破棄する
|
||||||
|
if force_reload_model is True:
|
||||||
|
self.net_g = None
|
||||||
|
|
||||||
# モデルがロードされていない場合はロードする
|
# モデルがロードされていない場合はロードする
|
||||||
if self.net_g is None:
|
if self.net_g is None:
|
||||||
self.load()
|
self.load()
|
||||||
@@ -402,6 +453,10 @@ class TTSModel:
|
|||||||
else:
|
else:
|
||||||
from style_bert_vits2.models.infer_onnx import infer_onnx
|
from style_bert_vits2.models.infer_onnx import infer_onnx
|
||||||
|
|
||||||
|
# force_reload_model が True のとき、メモリ上に保持されているモデルを破棄する
|
||||||
|
if force_reload_model is True:
|
||||||
|
self.onnx_session = None
|
||||||
|
|
||||||
# モデルがロードされていない場合はロードする
|
# モデルがロードされていない場合はロードする
|
||||||
if self.onnx_session is None:
|
if self.onnx_session is None:
|
||||||
self.load()
|
self.load()
|
||||||
|
|||||||
13
train_ms.py
13
train_ms.py
@@ -13,6 +13,7 @@ from torch.nn.parallel import DistributedDataParallel as DDP
|
|||||||
from torch.utils.data import DataLoader
|
from torch.utils.data import DataLoader
|
||||||
from torch.utils.tensorboard import SummaryWriter
|
from torch.utils.tensorboard import SummaryWriter
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
from transformers.trainer_pt_utils import DistributedLengthGroupedSampler
|
||||||
|
|
||||||
# logging.getLogger("numba").setLevel(logging.WARNING)
|
# logging.getLogger("numba").setLevel(logging.WARNING)
|
||||||
import default_style
|
import default_style
|
||||||
@@ -242,15 +243,23 @@ def run():
|
|||||||
# prefetch_factor=6,
|
# prefetch_factor=6,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
train_sampler = DistributedLengthGroupedSampler(
|
||||||
|
dataset=train_dataset,
|
||||||
|
batch_size=hps.train.batch_size,
|
||||||
|
num_replicas=n_gpus,
|
||||||
|
rank=rank,
|
||||||
|
lengths=train_dataset.lengths,
|
||||||
|
drop_last=True,
|
||||||
|
)
|
||||||
train_loader = DataLoader(
|
train_loader = DataLoader(
|
||||||
train_dataset,
|
train_dataset,
|
||||||
# メモリ消費量を減らそうとnum_workersを1にしてみる
|
# メモリ消費量を減らそうとnum_workersを1にしてみる
|
||||||
# num_workers=min(config.train_ms_config.num_workers, os.cpu_count() // 2),
|
# num_workers=min(config.train_ms_config.num_workers, os.cpu_count() // 2),
|
||||||
num_workers=1,
|
num_workers=1,
|
||||||
shuffle=True,
|
# shuffle=True,
|
||||||
pin_memory=True,
|
pin_memory=True,
|
||||||
collate_fn=collate_fn,
|
collate_fn=collate_fn,
|
||||||
# batch_sampler=train_sampler,
|
sampler=train_sampler,
|
||||||
batch_size=hps.train.batch_size,
|
batch_size=hps.train.batch_size,
|
||||||
persistent_workers=True,
|
persistent_workers=True,
|
||||||
# これもメモリ消費量を減らそうとしてコメントアウト
|
# これもメモリ消費量を減らそうとしてコメントアウト
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ from torch.nn.parallel import DistributedDataParallel as DDP
|
|||||||
from torch.utils.data import DataLoader
|
from torch.utils.data import DataLoader
|
||||||
from torch.utils.tensorboard import SummaryWriter
|
from torch.utils.tensorboard import SummaryWriter
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
|
from transformers.trainer_pt_utils import DistributedLengthGroupedSampler
|
||||||
|
|
||||||
# logging.getLogger("numba").setLevel(logging.WARNING)
|
# logging.getLogger("numba").setLevel(logging.WARNING)
|
||||||
import default_style
|
import default_style
|
||||||
@@ -243,20 +244,32 @@ def run():
|
|||||||
# prefetch_factor=6,
|
# prefetch_factor=6,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
train_sampler = DistributedLengthGroupedSampler(
|
||||||
|
dataset=train_dataset,
|
||||||
|
batch_size=hps.train.batch_size,
|
||||||
|
num_replicas=n_gpus,
|
||||||
|
rank=rank,
|
||||||
|
lengths=train_dataset.lengths,
|
||||||
|
drop_last=True,
|
||||||
|
)
|
||||||
train_loader = DataLoader(
|
train_loader = DataLoader(
|
||||||
train_dataset,
|
train_dataset,
|
||||||
# メモリ消費量を減らそうとnum_workersを1にしてみる
|
# メモリ消費量を減らそうとnum_workersを1にしてみる
|
||||||
# num_workers=min(config.train_ms_config.num_workers, os.cpu_count() // 2),
|
# num_workers=min(config.train_ms_config.num_workers, os.cpu_count() // 2),
|
||||||
num_workers=1,
|
num_workers=1,
|
||||||
shuffle=True,
|
# shuffle=True,
|
||||||
pin_memory=True,
|
pin_memory=True,
|
||||||
collate_fn=collate_fn,
|
collate_fn=collate_fn,
|
||||||
# batch_sampler=train_sampler,
|
sampler=train_sampler,
|
||||||
batch_size=hps.train.batch_size,
|
batch_size=hps.train.batch_size,
|
||||||
persistent_workers=True,
|
persistent_workers=True,
|
||||||
# これもメモリ消費量を減らそうとしてコメントアウト
|
# これもメモリ消費量を減らそうとしてコメントアウト
|
||||||
# prefetch_factor=6,
|
# prefetch_factor=6,
|
||||||
)
|
)
|
||||||
|
logger.info("Using DistributedLengthGroupedSampler for training.")
|
||||||
|
logger.debug(f"len(train_dataset): {len(train_dataset)}")
|
||||||
|
logger.debug(f"len(train_loader): {len(train_loader)}")
|
||||||
|
|
||||||
eval_dataset = None
|
eval_dataset = None
|
||||||
eval_loader = None
|
eval_loader = None
|
||||||
if rank == 0 and not args.speedup:
|
if rank == 0 and not args.speedup:
|
||||||
|
|||||||
Reference in New Issue
Block a user