ヌルモデル変更時にloadし直す用途で、推論時にforce_reload_modelを追加

This commit is contained in:
liruk
2024-06-21 09:19:46 +09:00
parent 7f2a4a8f3e
commit b645de0750
2 changed files with 19 additions and 14 deletions

View File

@@ -223,7 +223,7 @@ def change_null_model_row(null_model_index:int, null_model_name:str, null_model_
_ = null_models.pop(i, None) _ = null_models.pop(i, None)
result = null_models result = null_models
#logger.debug("change_null_model_row:res"+str(null_models)) #logger.debug("change_null_model_row:res"+str(null_models))
return result 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(
@@ -248,7 +248,8 @@ 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]]] 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
@@ -306,7 +307,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 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}")
@@ -328,7 +330,7 @@ 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
model_names = model_holder.model_names model_names = model_holder.model_names
if len(model_names) == 0: if len(model_names) == 0:
@@ -349,6 +351,7 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks:
gr.Markdown(initial_md) gr.Markdown(initial_md)
gr.Markdown(terms_of_use_md) gr.Markdown(terms_of_use_md)
null_models = gr.State({}) 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():
@@ -559,31 +562,31 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks:
inputs=[null_model_index, null_model_name, null_model_path,null_voice_weights, 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_voice_pitch_weights, null_speech_style_weights,null_tempo_weights,
null_models], null_models],
outputs=[null_models] outputs=[null_models,force_reload_model]
) )
null_voice_weights.change(change_null_model_row, null_voice_weights.change(change_null_model_row,
inputs=[null_model_index, null_model_name, null_model_path,null_voice_weights, 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_voice_pitch_weights, null_speech_style_weights,null_tempo_weights,
null_models], null_models],
outputs=[null_models] outputs=[null_models,force_reload_model]
) )
null_voice_pitch_weights.change(change_null_model_row, null_voice_pitch_weights.change(change_null_model_row,
inputs=[null_model_index, null_model_name, null_model_path,null_voice_weights, 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_voice_pitch_weights, null_speech_style_weights,null_tempo_weights,
null_models], null_models],
outputs=[null_models] outputs=[null_models,force_reload_model]
) )
null_speech_style_weights.change(change_null_model_row, null_speech_style_weights.change(change_null_model_row,
inputs=[null_model_index, null_model_name, null_model_path,null_voice_weights, 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_voice_pitch_weights, null_speech_style_weights,null_tempo_weights,
null_models], null_models],
outputs=[null_models] outputs=[null_models,force_reload_model]
) )
null_tempo_weights.change(change_null_model_row, null_tempo_weights.change(change_null_model_row,
inputs=[null_model_index, null_model_name, null_model_path,null_voice_weights, 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_voice_pitch_weights, null_speech_style_weights,null_tempo_weights,
null_models], null_models],
outputs=[null_models] outputs=[null_models,force_reload_model]
) )
add_btn = gr.Button("ヌルモデルを増やす") add_btn = gr.Button("ヌルモデルを増やす")
del_btn = gr.Button("ヌルモデルを減らす") del_btn = gr.Button("ヌルモデルを減らす")
@@ -655,9 +658,10 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks:
speaker, speaker,
pitch_scale, pitch_scale,
intonation_scale, intonation_scale,
null_models 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(

View File

@@ -258,7 +258,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]]] = {} null_model_params: dict[int,dict[str,Union[str, float]]] = {},
force_reload_model:bool = False
) -> tuple[int, NDArray[Any]]: ) -> tuple[int, NDArray[Any]]:
""" """
テキストから音声を合成する。 テキストから音声を合成する。
@@ -298,11 +299,11 @@ class TTSModel:
if assist_text == "" or not use_assist_text: if assist_text == "" or not use_assist_text:
assist_text = None assist_text = None
if null_model_params is not {}: if null_model_params is not {}:
#ヌルモデルがあるときは常時ロードしなおすけどもっといい手段ありそう
self.__net_g = None
self.null_model_params = null_model_params self.null_model_params = null_model_params
else: else:
self.null_model_params = {} self.null_model_params = {}
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()
assert self.__net_g is not None assert self.__net_g is not None