Refactor a little

This commit is contained in:
litagin02
2024-06-22 17:10:44 +09:00
parent 402346e493
commit c1381b2c49
2 changed files with 37 additions and 82 deletions

View File

@@ -219,21 +219,19 @@ def change_null_model_row(
null_tempo_weights: float, null_tempo_weights: float,
null_models: dict[int, dict[str, Any]], null_models: dict[int, dict[str, Any]],
): ):
# logger.debug("change_null_model_row:sta"+str(null_models)) null_models[null_model_index] = {
mid_result = {} "name": null_model_name,
mid_result["name"] = null_model_name "path": null_model_path,
mid_result["path"] = null_model_path "weight": null_voice_weights,
mid_result["weight"] = null_voice_weights "pitch": null_voice_pitch_weights,
mid_result["pitch"] = null_voice_pitch_weights "style": null_speech_style_weights,
mid_result["style"] = null_speech_style_weights "tempo": null_tempo_weights,
mid_result["tempo"] = null_tempo_weights }
null_models[null_model_index] = mid_result if len(null_models) > null_models_frame:
# logger.debug("decreasing:"+str(null_models_frame)+":"+str(len(null_models.keys()))) keys_to_keep = list(range(null_models_frame))
if null_models_frame < len(null_models.keys()): result = {k: null_models[k] for k in keys_to_keep}
for i in range(null_models_frame, len(null_models.keys())): else:
_ = null_models.pop(i, None)
result = null_models result = null_models
# logger.debug("change_null_model_row:res"+str(null_models))
return result, True return result, True
@@ -265,6 +263,7 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks:
): ):
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
@@ -344,6 +343,9 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks:
message = wrong_tone_message + "\n" + message message = wrong_tone_message + "\n" + message
return message, (sr, audio), kata_tone_json_str, False 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:
logger.error( logger.error(
@@ -355,9 +357,7 @@ 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)
@@ -478,23 +478,19 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks:
outputs=[assist_text, assist_text_weight], outputs=[assist_text, assist_text_weight],
) )
with gr.Accordion(label="ヌルモデル", open=False): with gr.Accordion(label="ヌルモデル", open=False):
with gr.Row() as null_row: with gr.Row():
null_models_count = gr.Number( null_models_count = gr.Number(
label="ヌルモデルの数", value=0, step=1 label="ヌルモデルの数", value=0, step=1
) )
with gr.Column(variant="panel") as null_column: with gr.Column(variant="panel"):
@gr.render( @gr.render(inputs=[null_models_count])
inputs=[ def render_null_models(
null_models_count,
]
)
def render_style(
null_models_count: int, null_models_count: int,
): ):
global null_models_frame global null_models_frame
null_models_frame = null_models_count null_models_frame = null_models_count
for i in range(0, null_models_count): for i in range(null_models_count):
with gr.Row(): with gr.Row():
null_model_index = gr.Number( null_model_index = gr.Number(
value=i, value=i,
@@ -506,27 +502,15 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks:
choices=model_names, choices=model_names,
key=f"null_model_name_{i}", key=f"null_model_name_{i}",
value=model_names[initial_id], value=model_names[initial_id],
interactive=True,
) )
if i in null_models.value:
logger.debug(
f"null model parameter exists in index {i}"
)
null_model_name.value = null_models.value[i][
"name"
]
null_model_path = gr.Dropdown( null_model_path = gr.Dropdown(
label="モデルファイル", label="モデルファイル",
choices=initial_pth_files,
key=f"null_model_path_{i}", key=f"null_model_path_{i}",
value=initial_pth_files[0], # FIXME: 再レンダー時に選択肢が消えるのでどうにかしたい
interactive=True, # 現在は再レンダーでvalueは保存されるが選択肢は保存されないので選択肢が空になる
# そのときに選択肢にない値となるので、それを許す
allow_custom_value=True,
) )
if i in null_models.value:
# null_model_path.choices = #ToDo
null_model_path.value = null_models.value[i][
"path"
]
null_voice_weights = gr.Slider( null_voice_weights = gr.Slider(
minimum=0, minimum=0,
maximum=1, maximum=1,
@@ -534,12 +518,7 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks:
step=0.1, step=0.1,
key=f"null_voice_weights_{i}", key=f"null_voice_weights_{i}",
label="声質", label="声質",
interactive=True,
) )
if i in null_models.value:
null_voice_weights.value = null_models.value[i][
"weight"
]
null_voice_pitch_weights = gr.Slider( null_voice_pitch_weights = gr.Slider(
minimum=0, minimum=0,
maximum=1, maximum=1,
@@ -547,11 +526,6 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks:
step=0.1, step=0.1,
key=f"null_voice_pitch_weights_{i}", key=f"null_voice_pitch_weights_{i}",
label="声の高さ", label="声の高さ",
interactive=True,
)
if i in null_models.value:
null_voice_pitch_weights.value = (
null_models.value[i]["pitch"]
) )
null_speech_style_weights = gr.Slider( null_speech_style_weights = gr.Slider(
minimum=0, minimum=0,
@@ -560,11 +534,6 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks:
step=0.1, step=0.1,
key=f"null_speech_style_weights_{i}", key=f"null_speech_style_weights_{i}",
label="話し方", label="話し方",
interactive=True,
)
if i in null_models.value:
null_speech_style_weights.value = (
null_models.value[i]["style"]
) )
null_tempo_weights = gr.Slider( null_tempo_weights = gr.Slider(
minimum=0, minimum=0,
@@ -573,18 +542,13 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks:
step=0.1, step=0.1,
key=f"null_tempo_weights_{i}", key=f"null_tempo_weights_{i}",
label="テンポ", label="テンポ",
interactive=True,
) )
if i in null_models.value:
null_tempo_weights.value = null_models.value[i][
"tempo"
]
null_model_name.change( null_model_name.change(
model_holder.update_model_files_for_gradio, model_holder.update_model_files_for_gradio,
inputs=[null_model_name], inputs=[null_model_name],
outputs=[null_model_path], outputs=[null_model_path],
) )
# null_model_name.change(model_holder.refresh, outputs=[])
null_model_path.change( null_model_path.change(
make_non_interactive, outputs=[tts_button] make_non_interactive, outputs=[tts_button]
) )

View File

@@ -117,46 +117,36 @@ class TTSModel:
if len(self.null_model_params.keys()) == 0: if len(self.null_model_params.keys()) == 0:
return return
for index, null_model in enumerate(self.null_model_params.keys()): 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( null_model_add = get_net_g(
model_path=str(self.null_model_params[index]["path"]), model_path=str(null_model_info["path"]),
version=self.hyper_parameters.version, version=self.hyper_parameters.version,
device=self.device, device=self.device,
hps=self.hyper_parameters, hps=self.hyper_parameters,
) )
# 愚直。もっと上手い方法ありそう # 愚直。もっと上手い方法ありそう
print(str(self.null_model_params[index]["weight"]))
params = zip(self.__net_g.dec.parameters(), null_model_add.dec.parameters()) params = zip(self.__net_g.dec.parameters(), null_model_add.dec.parameters())
for v in params: for v in params:
v[0].data.add_( v[0].data.add_(v[1].data, alpha=float(null_model_info["weight"]))
v[1].data, alpha=float(self.null_model_params[index]["weight"])
)
params = zip( params = zip(
self.__net_g.flow.parameters(), null_model_add.flow.parameters() self.__net_g.flow.parameters(), null_model_add.flow.parameters()
) )
for v in params: for v in params:
v[0].data.add_( v[0].data.add_(v[1].data, alpha=float(null_model_info["pitch"]))
v[1].data, alpha=float(self.null_model_params[index]["pitch"])
)
params = zip( params = zip(
self.__net_g.enc_p.parameters(), null_model_add.enc_p.parameters() self.__net_g.enc_p.parameters(), null_model_add.enc_p.parameters()
) )
for v in params: for v in params:
v[0].data.add_( v[0].data.add_(v[1].data, alpha=float(null_model_info["style"]))
v[1].data, alpha=float(self.null_model_params[index]["style"])
)
# テンポはsdpとdp二つあるからとりあえずどっちも足す # テンポはsdpとdp二つあるからとりあえずどっちも足す
params = zip(self.__net_g.sdp.parameters(), null_model_add.sdp.parameters()) params = zip(self.__net_g.sdp.parameters(), null_model_add.sdp.parameters())
for v in params: for v in params:
v[0].data.add_( v[0].data.add_(v[1].data, alpha=float(null_model_info["tempo"]))
v[1].data, alpha=float(self.null_model_params[index]["tempo"])
)
params = zip(self.__net_g.dp.parameters(), null_model_add.dp.parameters()) params = zip(self.__net_g.dp.parameters(), null_model_add.dp.parameters())
for v in params: for v in params:
v[0].data.add_( v[0].data.add_(v[1].data, alpha=float(null_model_info["tempo"]))
v[1].data, alpha=float(self.null_model_params[index]["tempo"])
)
def __get_style_vector(self, style_id: int, weight: float = 1.0) -> NDArray[Any]: def __get_style_vector(self, style_id: int, weight: float = 1.0) -> NDArray[Any]:
""" """
@@ -298,6 +288,7 @@ class TTSModel:
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。 null_model_params (dict[int, dict[str, Union[str, float]]], optional): 推論時に使用するヌルモデルの名前、適用割合のdictが入ったdict。
force_reload_model (bool, optional): モデルを強制的に再ロードするかどうか. Defaults to False.
Returns: Returns:
tuple[int, NDArray[Any]]: サンプリングレートと音声データ (16bit PCM) tuple[int, NDArray[Any]]: サンプリングレートと音声データ (16bit PCM)
""" """