ヌルモデル変更時にloadし直す用途で、推論時にforce_reload_modelを追加
This commit is contained in:
@@ -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(
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user