コード整理+暫定でヌルモデル使用時は常時load()するように
TODO:ヌルモデルの変更を検知できればそれがよさそう。
This commit is contained in:
@@ -208,7 +208,7 @@ 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,
|
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_voice_pitch_weights:float, null_speech_style_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))
|
#logger.debug("change_null_model_row:sta"+str(null_models))
|
||||||
mid_result={}
|
mid_result={}
|
||||||
mid_result["name"]=null_model_name
|
mid_result["name"]=null_model_name
|
||||||
mid_result["path"]=null_model_path
|
mid_result["path"]=null_model_path
|
||||||
@@ -217,12 +217,12 @@ def change_null_model_row(null_model_index:int, null_model_name:str, null_model_
|
|||||||
mid_result["style"]=null_speech_style_weights
|
mid_result["style"]=null_speech_style_weights
|
||||||
mid_result["tempo"]=null_tempo_weights
|
mid_result["tempo"]=null_tempo_weights
|
||||||
null_models[null_model_index] = mid_result
|
null_models[null_model_index] = mid_result
|
||||||
logger.debug("decreasing:"+str(null_models_frame)+":"+str(len(null_models.keys())))
|
#logger.debug("decreasing:"+str(null_models_frame)+":"+str(len(null_models.keys())))
|
||||||
if null_models_frame < len(null_models.keys()):
|
if null_models_frame < len(null_models.keys()):
|
||||||
for i in range(null_models_frame ,len(null_models.keys())):
|
for i in range(null_models_frame ,len(null_models.keys())):
|
||||||
_ = 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
|
||||||
|
|
||||||
def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks:
|
def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks:
|
||||||
@@ -253,10 +253,6 @@ 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
|
||||||
|
|
||||||
if len(null_models.keys()) > 0:
|
|
||||||
model_holder.get_null_models(null_models)
|
|
||||||
assert len(model_holder.null_models_params.keys()) > 0
|
|
||||||
|
|
||||||
wrong_tone_message = ""
|
wrong_tone_message = ""
|
||||||
kata_tone: Optional[list[tuple[str, int]]] = None
|
kata_tone: Optional[list[tuple[str, int]]] = None
|
||||||
if use_tone and kata_tone_json_str != "":
|
if use_tone and kata_tone_json_str != "":
|
||||||
|
|||||||
@@ -298,6 +298,8 @@ 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 = {}
|
||||||
@@ -407,7 +409,6 @@ class TTSModelHolder:
|
|||||||
self.device: str = device
|
self.device: str = device
|
||||||
self.model_files_dict: dict[str, list[Path]] = {}
|
self.model_files_dict: dict[str, list[Path]] = {}
|
||||||
self.current_model: Optional[TTSModel] = None
|
self.current_model: Optional[TTSModel] = None
|
||||||
self.null_models_params: dict[int,dict[str,Union[str, float]]] = {}
|
|
||||||
self.model_names: list[str] = []
|
self.model_names: list[str] = []
|
||||||
self.models_info: list[TTSModelInfo] = []
|
self.models_info: list[TTSModelInfo] = []
|
||||||
self.refresh()
|
self.refresh()
|
||||||
@@ -420,7 +421,6 @@ class TTSModelHolder:
|
|||||||
self.model_files_dict = {}
|
self.model_files_dict = {}
|
||||||
self.model_names = []
|
self.model_names = []
|
||||||
self.current_model = None
|
self.current_model = None
|
||||||
self.null_models_params = {}
|
|
||||||
self.models_info = []
|
self.models_info = []
|
||||||
|
|
||||||
model_dirs = [d for d in self.root_dir.iterdir() if d.is_dir()]
|
model_dirs = [d for d in self.root_dir.iterdir() if d.is_dir()]
|
||||||
@@ -482,13 +482,6 @@ class TTSModelHolder:
|
|||||||
)
|
)
|
||||||
|
|
||||||
return self.current_model
|
return self.current_model
|
||||||
def get_null_models(self, null_model_index:dict[int,dict[str,Union[str, float]]]) -> dict[int, dict[str, Union[str, float]]]:
|
|
||||||
"""
|
|
||||||
get_modelをヌルモデル用に改変。複数まとめて使うかも知れないのでdictで戻す
|
|
||||||
|
|
||||||
"""
|
|
||||||
self.null_models_params = null_model_index
|
|
||||||
return self.null_models_params
|
|
||||||
|
|
||||||
def get_model_for_gradio(self, model_name: str, model_path_str: str):
|
def get_model_for_gradio(self, model_name: str, model_path_str: str):
|
||||||
import gradio as gr
|
import gradio as gr
|
||||||
|
|||||||
Reference in New Issue
Block a user