From 86c2a1b08724643097cef728fb702380e998bc82 Mon Sep 17 00:00:00 2001 From: liruk <34234523+liruk@users.noreply.github.com> Date: Tue, 18 Jun 2024 19:15:20 +0900 Subject: [PATCH 01/13] =?UTF-8?q?=E6=8E=A8=E8=AB=96=E6=99=82=E3=83=8C?= =?UTF-8?q?=E3=83=AB=E3=83=A2=E3=83=87=E3=83=AB=E3=83=9E=E3=83=BC=E3=82=B8?= =?UTF-8?q?(=E3=81=BE=E3=81=A0=E5=8B=95=E3=81=8B=E3=81=AA=E3=81=84?= =?UTF-8?q?=E3=82=88)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- gradio_tabs/inference.py | 122 +++++++++++++++++++++++++++++++++- style_bert_vits2/tts_model.py | 68 ++++++++++++++++++- 2 files changed, 187 insertions(+), 3 deletions(-) diff --git a/gradio_tabs/inference.py b/gradio_tabs/inference.py index 53393be..c7a8d5e 100644 --- a/gradio_tabs/inference.py +++ b/gradio_tabs/inference.py @@ -185,7 +185,12 @@ style_md = f""" - どのくらいに強さがいいかはモデルやスタイルによって異なるようです。 - 音声ファイルを入力する場合は、学習データと似た声音の話者(特に同じ性別)でないとよい効果が出ないかもしれません。 """ +voice_keys = ["dec"] +voice_pitch_keys = ["flow"] +speech_style_keys = ["enc_p"] +tempo_keys = ["sdp", "dp"] +null_models = {}#グローバルに置いてるけどもっと良い置き場所ありそう def make_interactive(): return gr.update(interactive=True, value="音声合成") @@ -228,6 +233,14 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: ): model_holder.get_model(model_name, model_path) assert model_holder.current_model is not None + #null_model_paths, + #null_voice_weights, + #null_voice_pitch_weights, + #null_speech_style_weights, + #null_tempo_weights + if len(null_models) > 0 and len(null_models["names"].keys()) > 0: + model_holder.get_null_models(null_models) + assert len(model_holder.current_null_models) > 0 wrong_tone_message = "" kata_tone: Optional[list[tuple[str, int]]] = None @@ -282,6 +295,7 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: speaker_id=speaker_id, pitch_scale=pitch_scale, intonation_scale=intonation_scale, + null_model_params = null_models ) except InvalidToneError as e: logger.error(f"Tone error: {e}") @@ -436,6 +450,113 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: inputs=[use_assist_text], outputs=[assist_text, assist_text_weight], ) + with gr.Accordion(label="ヌルモデル", open=False): + #null_voice_weights, + #null_voice_pitch_weights, + #null_speech_style_weights, + #null_tempo_weights + with gr.Row(): + style_count = gr.Number(label="作るスタイルの数", value=0, step=1) + with gr.Column(variant="panel"): + @gr.render( + inputs=[ + style_count, + ] + ) + def render_style( + style_count + ): + name_components = {} + path_components = {} + weight_components = {} + pitch_components = {} + style_components = {} + tempo_components = {} + for i in range(style_count): + with gr.Row(): + null_model_name = gr.Dropdown( + label="モデル一覧", + choices=model_names, + key=f"null_model_name_{i}", + value=model_names[initial_id], + interactive=True + ) + null_model_path = gr.Dropdown( + label="モデルファイル", + choices=initial_pth_files, + key=f"null_model_path_{i}", + value=initial_pth_files[0], + interactive=True + ) + null_voice_weights = gr.Slider( + minimum=0, + maximum=1, + value=1, + step=0.1, + key=f"null_voice_weights_{i}", + label="声質", + interactive=True + ) + null_voice_pitch_weights = gr.Slider( + minimum=0, + maximum=1, + value=1, + step=0.1, + key=f"null_voice_pitch_weights_{i}", + label="声の高さ", + interactive=True + ) + null_speech_style_weights = gr.Slider( + minimum=0, + maximum=1, + value=1, + step=0.1, + key=f"null_speech_style_weights_{i}", + label="話し方", + interactive=True + ) + null_tempo_weights = gr.Slider( + minimum=0, + maximum=1, + value=1, + step=0.1, + key=f"null_tempo_weights_{i}", + label="テンポ", + interactive=True + ) + 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]) + #もっといい方法ありそう + name_components[str(i)]=(null_model_name) + path_components[str(i)]=(null_model_path) + weight_components[str(i)]=(null_voice_weights) + pitch_components[str(i)]=(null_voice_pitch_weights) + style_components[str(i)]=(null_speech_style_weights) + tempo_components[str(i)]=(null_tempo_weights) + null_models["names"]=(name_components) + null_models["paths"]=(path_components) + null_models["weights"]=(weight_components) + null_models["pitchs"]=(pitch_components) + null_models["styles"]=(style_components) + null_models["tempos"]=(tempo_components) + + add_btn = gr.Button("ヌルモデルを増やす") + del_btn = gr.Button("ヌルモデルを減らす") + add_btn.click( + lambda x: x + 1, + inputs=[style_count], + outputs=[style_count], + ) + del_btn.click( + lambda x: x - 1 if x > 0 else 0, + inputs=[style_count], + outputs=[style_count], + ) + with gr.Column(): with gr.Accordion("スタイルについて詳細", open=False): gr.Markdown(style_md) @@ -524,7 +645,6 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: return app - if __name__ == "__main__": from config import get_path_config import torch diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index 6df8394..bb38b3a 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -61,6 +61,7 @@ class TTSModel: self.model_path: Path = model_path self.device: str = device + self.null_model_params: dict[str, dict[str,Any]] = {} # ハイパーパラメータの Pydantic モデルが直接指定された if isinstance(config_path, HyperParameters): @@ -113,6 +114,36 @@ class TTSModel: device=self.device, hps=self.hyper_parameters, ) + if(len(self.null_model_params.keys())==0): + return + + for index, null_model in enumerate(self.null_model_params["names"].keys()): + null_model_add = get_net_g( + model_path=str(self.null_model_params["paths"][str(index)].value), + 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=self.null_model_params["weights"][str(index)].value) + + params = zip(self.__net_g.flow.parameters(), null_model_add.flow.parameters()) + for v in params: + v[0].data.add(v[1].data,alpha=self.null_model_params["pitchs"][str(index)].value) + + 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=self.null_model_params["styles"][str(index)].value) + #テンポは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=self.null_model_params["tempos"][str(index)].value) + params = zip(self.__net_g.dp.parameters(), null_model_add.dp.parameters()) + for v in params: + v[0].data.add(v[1].data,alpha=self.null_model_params["tempos"][str(index)].value) + def __get_style_vector(self, style_id: int, weight: float = 1.0) -> NDArray[Any]: """ @@ -227,6 +258,7 @@ class TTSModel: given_tone: Optional[list[int]] = None, pitch_scale: float = 1.0, intonation_scale: float = 1.0, + null_model_params: dict[str,dict[str,Any]] = {} ) -> tuple[int, NDArray[Any]]: """ テキストから音声を合成する。 @@ -251,7 +283,7 @@ class TTSModel: given_tone (Optional[list[int]], optional): アクセントのトーンのリスト. Defaults to None. pitch_scale (float, optional): ピッチの高さ (1.0 から変更すると若干音質が低下する). Defaults to 1.0. intonation_scale (float, optional): 抑揚の平均からの変化幅 (1.0 から変更すると若干音質が低下する). Defaults to 1.0. - + null_model_params(dict[str,dict[str,gr.Component],optional):推論時に使用するヌルモデルの名前、重みのdictが入ったdict。 Returns: tuple[int, NDArray[Any]]: サンプリングレートと音声データ (16bit PCM) """ @@ -265,7 +297,10 @@ class TTSModel: reference_audio_path = None if assist_text == "" or not use_assist_text: assist_text = None - + if null_model_params is not {}: + self.null_model_params = null_model_params + else: + self.null_model_params = {} if self.__net_g is None: self.load() assert self.__net_g is not None @@ -372,6 +407,8 @@ class TTSModelHolder: self.device: str = device self.model_files_dict: dict[str, list[Path]] = {} self.current_model: Optional[TTSModel] = None + self.current_null_models: dict[str,TTSModel] = {} + self.null_models_params: dict[str,dict[str,Any]] = {} self.model_names: list[str] = [] self.models_info: list[TTSModelInfo] = [] self.refresh() @@ -384,6 +421,8 @@ class TTSModelHolder: self.model_files_dict = {} self.model_names = [] self.current_model = None + self.current_null_models = {} + self.null_models_params = {} self.models_info = [] model_dirs = [d for d in self.root_dir.iterdir() if d.is_dir()] @@ -445,6 +484,29 @@ class TTSModelHolder: ) return self.current_model + def get_null_models(self, null_model_index:dict[str,dict[str,Any]]) -> dict[str, TTSModel]: + """ + get_modelをヌルモデル用に改変。複数まとめて使うかも知れないのでdictで戻す + + """ + self.current_null_models = {} + self.null_models_params = null_model_index + for index, value in enumerate(null_model_index["names"]): + model_path = Path(null_model_index["paths"][str(index)].value) + model_name = null_model_index["names"][str(index)].value + if model_name not in self.model_files_dict: + raise ValueError(f"Model `{model_name}` is not found") + if model_path not in self.model_files_dict[model_name]: + raise ValueError(f"Model file `{model_path}` is not found") + + if len(self.current_null_models) == 0 or model_path not in self.current_null_models.keys(): + self.current_null_models[model_name] = TTSModel( + model_path=model_path, + config_path=self.root_dir / model_name / "config.json", + style_vec_path=self.root_dir / model_name / "style_vectors.npy", + device=self.device, + ) + return self.current_null_models def get_model_for_gradio(self, model_name: str, model_path_str: str): import gradio as gr @@ -474,6 +536,8 @@ class TTSModelHolder: ) speakers = list(self.current_model.spk2id.keys()) styles = list(self.current_model.style2id.keys()) + #if len(null_models.keys())!=0: + # self.get_null_models(null_models) return ( gr.Dropdown(choices=styles, value=styles[0]), # type: ignore gr.Button(interactive=True, value="音声合成"), From fb99a776ce17ef97a736575647961afbf5d0549e Mon Sep 17 00:00:00 2001 From: liruk <34234523+liruk@users.noreply.github.com> Date: Wed, 19 Jun 2024 18:31:27 +0900 Subject: [PATCH 02/13] =?UTF-8?q?=E6=8E=A8=E8=AB=96=E6=99=82=E3=83=8C?= =?UTF-8?q?=E3=83=AB=E3=83=A2=E3=83=87=E3=83=AB=E3=83=9E=E3=83=BC=E3=82=B8?= =?UTF-8?q?(=E3=81=A7=E3=81=8D=E3=82=8B=E3=81=91=E3=81=A9=E3=83=90?= =?UTF-8?q?=E3=82=B0=E3=83=90=E3=82=B0)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- gradio_tabs/inference.py | 110 +++++++++++++++++++++------------- style_bert_vits2/tts_model.py | 48 +++++---------- 2 files changed, 81 insertions(+), 77 deletions(-) diff --git a/gradio_tabs/inference.py b/gradio_tabs/inference.py index c7a8d5e..38c2af6 100644 --- a/gradio_tabs/inference.py +++ b/gradio_tabs/inference.py @@ -1,6 +1,6 @@ import datetime import json -from typing import Optional +from typing import Optional, Any, Union import gradio as gr @@ -190,8 +190,6 @@ voice_pitch_keys = ["flow"] speech_style_keys = ["enc_p"] tempo_keys = ["sdp", "dp"] -null_models = {}#グローバルに置いてるけどもっと良い置き場所ありそう - def make_interactive(): return gr.update(interactive=True, value="音声合成") @@ -205,7 +203,19 @@ def gr_util(item): return (gr.update(visible=True), gr.Audio(visible=False, value=None)) else: return (gr.update(visible=False), gr.update(visible=True)) - +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]]): + mid_result={} + mid_result["name"]=null_model_name + mid_result["path"]=null_model_path + mid_result["weight"]=null_tempo_weights + mid_result["pitch"]=null_voice_pitch_weights + mid_result["style"]=null_speech_style_weights + mid_result["tempo"]=null_tempo_weights + null_models[null_model_index] = mid_result + result = null_models + return result def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: def tts_fn( @@ -230,17 +240,14 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: speaker, pitch_scale, intonation_scale, + null_models:dict[int, dict[str, Union[str, float]]] ): model_holder.get_model(model_name, model_path) assert model_holder.current_model is not None - #null_model_paths, - #null_voice_weights, - #null_voice_pitch_weights, - #null_speech_style_weights, - #null_tempo_weights - if len(null_models) > 0 and len(null_models["names"].keys()) > 0: + + if len(null_models.keys()) > 0: model_holder.get_null_models(null_models) - assert len(model_holder.current_null_models) > 0 + assert len(model_holder.null_models_params.keys()) > 0 wrong_tone_message = "" kata_tone: Optional[list[tuple[str, int]]] = None @@ -337,6 +344,7 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: with gr.Blocks(theme=GRADIO_THEME) as app: gr.Markdown(initial_md) gr.Markdown(terms_of_use_md) + null_models = gr.State({}) with gr.Accordion(label="使い方", open=False): gr.Markdown(how_to_md) with gr.Row(): @@ -451,29 +459,24 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: outputs=[assist_text, assist_text_weight], ) with gr.Accordion(label="ヌルモデル", open=False): - #null_voice_weights, - #null_voice_pitch_weights, - #null_speech_style_weights, - #null_tempo_weights with gr.Row(): - style_count = gr.Number(label="作るスタイルの数", value=0, step=1) + null_models_count = gr.Number(label="ヌルモデルの数", value=0, step=1) with gr.Column(variant="panel"): @gr.render( inputs=[ - style_count, + null_models_count, ] ) def render_style( - style_count + null_models_count:int, ): - name_components = {} - path_components = {} - weight_components = {} - pitch_components = {} - style_components = {} - tempo_components = {} - for i in range(style_count): + for i in range(0, 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, @@ -530,31 +533,51 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: outputs=[null_model_path], ) null_model_path.change(make_non_interactive, outputs=[tts_button]) - #もっといい方法ありそう - name_components[str(i)]=(null_model_name) - path_components[str(i)]=(null_model_path) - weight_components[str(i)]=(null_voice_weights) - pitch_components[str(i)]=(null_voice_pitch_weights) - style_components[str(i)]=(null_speech_style_weights) - tempo_components[str(i)]=(null_tempo_weights) - null_models["names"]=(name_components) - null_models["paths"]=(path_components) - null_models["weights"]=(weight_components) - null_models["pitchs"]=(pitch_components) - null_models["styles"]=(style_components) - null_models["tempos"]=(tempo_components) - + #愚直すぎるのでもう少しなんとかしたい + 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] + ) + 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] + ) + 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] + ) + 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] + ) + 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] + ) + if null_models_count < len(null_models.value.keys()): + for i in range(null_models_count ,len(null_models.value.keys())): + _ = null_models.value.pop(i, None) add_btn = gr.Button("ヌルモデルを増やす") del_btn = gr.Button("ヌルモデルを減らす") add_btn.click( lambda x: x + 1, - inputs=[style_count], - outputs=[style_count], + inputs=[null_models_count], + outputs=[null_models_count], ) del_btn.click( lambda x: x - 1 if x > 0 else 0, - inputs=[style_count], - outputs=[style_count], + inputs=[null_models_count], + outputs=[null_models_count], ) with gr.Column(): @@ -614,6 +637,7 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: speaker, pitch_scale, intonation_scale, + null_models ], outputs=[text_output, audio_output, tone], ) diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index bb38b3a..f8a387b 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -61,7 +61,7 @@ class TTSModel: self.model_path: Path = model_path self.device: str = device - self.null_model_params: dict[str, dict[str,Any]] = {} + self.null_model_params: dict[int, dict[str,Union[float, str]]] = {} # ハイパーパラメータの Pydantic モデルが直接指定された if isinstance(config_path, HyperParameters): @@ -117,32 +117,32 @@ class TTSModel: if(len(self.null_model_params.keys())==0): return - for index, null_model in enumerate(self.null_model_params["names"].keys()): + for index, null_model in enumerate(self.null_model_params.keys()): null_model_add = get_net_g( - model_path=str(self.null_model_params["paths"][str(index)].value), + model_path=str(self.null_model_params[index]["path"]), version=self.hyper_parameters.version, device=self.device, hps=self.hyper_parameters, ) #愚直。もっと上手い方法ありそう + print(str(self.null_model_params[index]["weight"])) params = zip(self.__net_g.dec.parameters(), null_model_add.dec.parameters()) for v in params: - v[0].data.add(v[1].data,alpha=self.null_model_params["weights"][str(index)].value) - + v[0].data.add_(v[1].data,alpha=float(self.null_model_params[index]["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=self.null_model_params["pitchs"][str(index)].value) + v[0].data.add_(v[1].data,alpha=float(self.null_model_params[index]["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=self.null_model_params["styles"][str(index)].value) + v[0].data.add_(v[1].data,alpha=float(self.null_model_params[index]["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=self.null_model_params["tempos"][str(index)].value) + v[0].data.add_(v[1].data,alpha=float(self.null_model_params[index]["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=self.null_model_params["tempos"][str(index)].value) + v[0].data.add_(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]: @@ -258,7 +258,7 @@ class TTSModel: given_tone: Optional[list[int]] = None, pitch_scale: float = 1.0, intonation_scale: float = 1.0, - null_model_params: dict[str,dict[str,Any]] = {} + null_model_params: dict[int,dict[str,Union[str, float]]] = {} ) -> tuple[int, NDArray[Any]]: """ テキストから音声を合成する。 @@ -283,7 +283,7 @@ class TTSModel: given_tone (Optional[list[int]], optional): アクセントのトーンのリスト. Defaults to None. pitch_scale (float, optional): ピッチの高さ (1.0 から変更すると若干音質が低下する). Defaults to 1.0. intonation_scale (float, optional): 抑揚の平均からの変化幅 (1.0 から変更すると若干音質が低下する). Defaults to 1.0. - null_model_params(dict[str,dict[str,gr.Component],optional):推論時に使用するヌルモデルの名前、重みのdictが入ったdict。 + null_model_params(dict[int,dict[str,Union[str,float]],optional):推論時に使用するヌルモデルの名前、適用割合のdictが入ったdict。 Returns: tuple[int, NDArray[Any]]: サンプリングレートと音声データ (16bit PCM) """ @@ -407,8 +407,7 @@ class TTSModelHolder: self.device: str = device self.model_files_dict: dict[str, list[Path]] = {} self.current_model: Optional[TTSModel] = None - self.current_null_models: dict[str,TTSModel] = {} - self.null_models_params: dict[str,dict[str,Any]] = {} + self.null_models_params: dict[int,dict[str,Union[str, float]]] = {} self.model_names: list[str] = [] self.models_info: list[TTSModelInfo] = [] self.refresh() @@ -421,7 +420,6 @@ class TTSModelHolder: self.model_files_dict = {} self.model_names = [] self.current_model = None - self.current_null_models = {} self.null_models_params = {} self.models_info = [] @@ -484,29 +482,13 @@ class TTSModelHolder: ) return self.current_model - def get_null_models(self, null_model_index:dict[str,dict[str,Any]]) -> dict[str, TTSModel]: + 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.current_null_models = {} self.null_models_params = null_model_index - for index, value in enumerate(null_model_index["names"]): - model_path = Path(null_model_index["paths"][str(index)].value) - model_name = null_model_index["names"][str(index)].value - if model_name not in self.model_files_dict: - raise ValueError(f"Model `{model_name}` is not found") - if model_path not in self.model_files_dict[model_name]: - raise ValueError(f"Model file `{model_path}` is not found") - - if len(self.current_null_models) == 0 or model_path not in self.current_null_models.keys(): - self.current_null_models[model_name] = TTSModel( - model_path=model_path, - config_path=self.root_dir / model_name / "config.json", - style_vec_path=self.root_dir / model_name / "style_vectors.npy", - device=self.device, - ) - return self.current_null_models + return self.null_models_params def get_model_for_gradio(self, model_name: str, model_path_str: str): import gradio as gr @@ -536,8 +518,6 @@ class TTSModelHolder: ) speakers = list(self.current_model.spk2id.keys()) styles = list(self.current_model.style2id.keys()) - #if len(null_models.keys())!=0: - # self.get_null_models(null_models) return ( gr.Dropdown(choices=styles, value=styles[0]), # type: ignore gr.Button(interactive=True, value="音声合成"), From c957cab1e55b0fb2f946d1b072f5b38485c2e8d3 Mon Sep 17 00:00:00 2001 From: liruk <34234523+liruk@users.noreply.github.com> Date: Thu, 20 Jun 2024 19:42:16 +0900 Subject: [PATCH 03/13] =?UTF-8?q?nullmodel=E6=95=B0=E3=82=92=E5=A2=97?= =?UTF-8?q?=E6=B8=9B=E3=81=95=E3=81=9B=E3=81=9F=E3=81=A8=E3=81=8D=E3=81=AE?= =?UTF-8?q?=E5=95=8F=E9=A1=8C=E3=82=92=E4=B8=80=E9=83=A8=E4=BF=AE=E6=AD=A3?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 1→0にした時の問題は未修正 --- gradio_tabs/inference.py | 32 +++++++++++++++++++++++++++----- 1 file changed, 27 insertions(+), 5 deletions(-) diff --git a/gradio_tabs/inference.py b/gradio_tabs/inference.py index 38c2af6..324b72c 100644 --- a/gradio_tabs/inference.py +++ b/gradio_tabs/inference.py @@ -203,9 +203,12 @@ def gr_util(item): return (gr.update(visible=True), gr.Audio(visible=False, value=None)) else: 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]]): + logger.debug("change_null_model_row:sta"+str(null_models)) mid_result={} mid_result["name"]=null_model_name mid_result["path"]=null_model_path @@ -214,7 +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["tempo"]=null_tempo_weights null_models[null_model_index] = mid_result + logger.debug("decreasing:"+str(null_models_frame)+":"+str(len(null_models.keys()))) + if null_models_frame < len(null_models.keys()): + for i in range(null_models_frame ,len(null_models.keys())): + _ = null_models.pop(i, None) result = null_models + logger.debug("change_null_model_row:res"+str(null_models)) return result def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: @@ -459,9 +467,9 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: outputs=[assist_text, assist_text_weight], ) with gr.Accordion(label="ヌルモデル", open=False): - with gr.Row(): + with gr.Row() as null_row: null_models_count = gr.Number(label="ヌルモデルの数", value=0, step=1) - with gr.Column(variant="panel"): + with gr.Column(variant="panel") as null_column: @gr.render( inputs=[ null_models_count, @@ -470,6 +478,8 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: def render_style( null_models_count:int, ): + global null_models_frame + null_models_frame = null_models_count for i in range(0, null_models_count): with gr.Row(): null_model_index = gr.Number( @@ -484,6 +494,9 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: 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( label="モデルファイル", choices=initial_pth_files, @@ -491,6 +504,9 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: value=initial_pth_files[0], interactive=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( minimum=0, maximum=1, @@ -500,6 +516,8 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: label="声質", interactive=True ) + if i in null_models.value: + null_voice_weights.value=null_models.value[i]["weight"] null_voice_pitch_weights = gr.Slider( minimum=0, maximum=1, @@ -509,6 +527,8 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: 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( minimum=0, maximum=1, @@ -518,6 +538,8 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: label="話し方", interactive=True ) + if i in null_models.value: + null_speech_style_weights.value=null_models.value[i]["style"] null_tempo_weights = gr.Slider( minimum=0, maximum=1, @@ -527,11 +549,14 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: label="テンポ", interactive=True ) + if i in null_models.value: + null_tempo_weights.value=null_models.value[i]["tempo"] null_model_name.change( model_holder.update_model_files_for_gradio, inputs=[null_model_name], outputs=[null_model_path], ) + #null_model_name.change(model_holder.refresh, outputs=[]) null_model_path.change(make_non_interactive, outputs=[tts_button]) #愚直すぎるのでもう少しなんとかしたい null_model_path.change(change_null_model_row, @@ -564,9 +589,6 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: null_models], outputs=[null_models] ) - if null_models_count < len(null_models.value.keys()): - for i in range(null_models_count ,len(null_models.value.keys())): - _ = null_models.value.pop(i, None) add_btn = gr.Button("ヌルモデルを増やす") del_btn = gr.Button("ヌルモデルを減らす") add_btn.click( From 7f2a4a8f3e4692825b55c50210c45d99d3a9e202 Mon Sep 17 00:00:00 2001 From: liruk <34234523+liruk@users.noreply.github.com> Date: Fri, 21 Jun 2024 09:07:09 +0900 Subject: [PATCH 04/13] =?UTF-8?q?=E3=82=B3=E3=83=BC=E3=83=89=E6=95=B4?= =?UTF-8?q?=E7=90=86+=E6=9A=AB=E5=AE=9A=E3=81=A7=E3=83=8C=E3=83=AB?= =?UTF-8?q?=E3=83=A2=E3=83=87=E3=83=AB=E4=BD=BF=E7=94=A8=E6=99=82=E3=81=AF?= =?UTF-8?q?=E5=B8=B8=E6=99=82load()=E3=81=99=E3=82=8B=E3=82=88=E3=81=86?= =?UTF-8?q?=E3=81=AB?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit TODO:ヌルモデルの変更を検知できればそれがよさそう。 --- gradio_tabs/inference.py | 10 +++------- style_bert_vits2/tts_model.py | 11 ++--------- 2 files changed, 5 insertions(+), 16 deletions(-) diff --git a/gradio_tabs/inference.py b/gradio_tabs/inference.py index 324b72c..570fce6 100644 --- a/gradio_tabs/inference.py +++ b/gradio_tabs/inference.py @@ -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, null_voice_pitch_weights:float, null_speech_style_weights:float,null_tempo_weights:float, 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["name"]=null_model_name 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["tempo"]=null_tempo_weights 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()): for i in range(null_models_frame ,len(null_models.keys())): _ = null_models.pop(i, None) 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 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) 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 = "" kata_tone: Optional[list[tuple[str, int]]] = None if use_tone and kata_tone_json_str != "": diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index f8a387b..4e1cdac 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -298,6 +298,8 @@ class TTSModel: if assist_text == "" or not use_assist_text: assist_text = None if null_model_params is not {}: + #ヌルモデルがあるときは常時ロードしなおすけどもっといい手段ありそう + self.__net_g = None self.null_model_params = null_model_params else: self.null_model_params = {} @@ -407,7 +409,6 @@ class TTSModelHolder: self.device: str = device self.model_files_dict: dict[str, list[Path]] = {} self.current_model: Optional[TTSModel] = None - self.null_models_params: dict[int,dict[str,Union[str, float]]] = {} self.model_names: list[str] = [] self.models_info: list[TTSModelInfo] = [] self.refresh() @@ -420,7 +421,6 @@ class TTSModelHolder: self.model_files_dict = {} self.model_names = [] self.current_model = None - self.null_models_params = {} self.models_info = [] model_dirs = [d for d in self.root_dir.iterdir() if d.is_dir()] @@ -482,13 +482,6 @@ class TTSModelHolder: ) 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): import gradio as gr From b645de0750691b7f0dd884a424e4bbb6fac17c7e Mon Sep 17 00:00:00 2001 From: liruk <34234523+liruk@users.noreply.github.com> Date: Fri, 21 Jun 2024 09:19:46 +0900 Subject: [PATCH 05/13] =?UTF-8?q?=E3=83=8C=E3=83=AB=E3=83=A2=E3=83=87?= =?UTF-8?q?=E3=83=AB=E5=A4=89=E6=9B=B4=E6=99=82=E3=81=ABload=E3=81=97?= =?UTF-8?q?=E7=9B=B4=E3=81=99=E7=94=A8=E9=80=94=E3=81=A7=E3=80=81=E6=8E=A8?= =?UTF-8?q?=E8=AB=96=E6=99=82=E3=81=ABforce=5Freload=5Fmodel=E3=82=92?= =?UTF-8?q?=E8=BF=BD=E5=8A=A0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- gradio_tabs/inference.py | 26 +++++++++++++++----------- style_bert_vits2/tts_model.py | 7 ++++--- 2 files changed, 19 insertions(+), 14 deletions(-) diff --git a/gradio_tabs/inference.py b/gradio_tabs/inference.py index 570fce6..ab8bb1e 100644 --- a/gradio_tabs/inference.py +++ b/gradio_tabs/inference.py @@ -223,7 +223,7 @@ def change_null_model_row(null_model_index:int, null_model_name:str, null_model_ _ = null_models.pop(i, None) result = 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 tts_fn( @@ -248,7 +248,8 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: speaker, pitch_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) 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, pitch_scale=pitch_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: 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." if wrong_tone_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 if len(model_names) == 0: @@ -349,6 +351,7 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: gr.Markdown(initial_md) gr.Markdown(terms_of_use_md) null_models = gr.State({}) + force_reload_model = gr.State(False) with gr.Accordion(label="使い方", open=False): gr.Markdown(how_to_md) 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, null_voice_pitch_weights, null_speech_style_weights,null_tempo_weights, null_models], - outputs=[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] + 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] + 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] + 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] + outputs=[null_models,force_reload_model] ) add_btn = gr.Button("ヌルモデルを増やす") del_btn = gr.Button("ヌルモデルを減らす") @@ -655,9 +658,10 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: speaker, pitch_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( diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index 4e1cdac..77db3dd 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -258,7 +258,8 @@ class TTSModel: given_tone: Optional[list[int]] = None, pitch_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]]: """ テキストから音声を合成する。 @@ -298,11 +299,11 @@ class TTSModel: if assist_text == "" or not use_assist_text: assist_text = None if null_model_params is not {}: - #ヌルモデルがあるときは常時ロードしなおすけどもっといい手段ありそう - self.__net_g = None self.null_model_params = null_model_params else: self.null_model_params = {} + if force_reload_model is True: + self.__net_g = None if self.__net_g is None: self.load() assert self.__net_g is not None From 9eab977138c9abf690a1e893fb73d51b027829cb Mon Sep 17 00:00:00 2001 From: litagin02 Date: Sat, 22 Jun 2024 14:54:47 +0900 Subject: [PATCH 06/13] Style --- gradio_tabs/inference.py | 206 +++++++++++++++++++++++----------- style_bert_vits2/tts_model.py | 43 ++++--- 2 files changed, 168 insertions(+), 81 deletions(-) diff --git a/gradio_tabs/inference.py b/gradio_tabs/inference.py index ab8bb1e..2673151 100644 --- a/gradio_tabs/inference.py +++ b/gradio_tabs/inference.py @@ -1,6 +1,6 @@ import datetime import json -from typing import Optional, Any, Union +from typing import Any, Optional, Union import gradio as gr @@ -190,6 +190,7 @@ voice_pitch_keys = ["flow"] speech_style_keys = ["enc_p"] tempo_keys = ["sdp", "dp"] + def make_interactive(): return gr.update(interactive=True, value="音声合成") @@ -204,27 +205,38 @@ def gr_util(item): else: 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]]): - #logger.debug("change_null_model_row:sta"+str(null_models)) - mid_result={} - mid_result["name"]=null_model_name - mid_result["path"]=null_model_path - mid_result["weight"]=null_tempo_weights - mid_result["pitch"]=null_voice_pitch_weights - mid_result["style"]=null_speech_style_weights - mid_result["tempo"]=null_tempo_weights + + +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]], +): + # logger.debug("change_null_model_row:sta"+str(null_models)) + mid_result = {} + mid_result["name"] = null_model_name + mid_result["path"] = null_model_path + mid_result["weight"] = null_tempo_weights + mid_result["pitch"] = null_voice_pitch_weights + mid_result["style"] = null_speech_style_weights + mid_result["tempo"] = null_tempo_weights 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()): - 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) 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, True + def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: def tts_fn( model_name, @@ -248,8 +260,8 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: speaker, pitch_scale, intonation_scale, - null_models:dict[int, dict[str, Union[str, float]]], - force_reload_model:bool + null_models: dict[int, dict[str, Union[str, float]]], + force_reload_model: bool, ): model_holder.get_model(model_name, model_path) assert model_holder.current_model is not None @@ -307,8 +319,8 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: speaker_id=speaker_id, pitch_scale=pitch_scale, intonation_scale=intonation_scale, - null_model_params = null_models, - force_reload_model = force_reload_model + null_model_params=null_models, + force_reload_model=force_reload_model, ) except InvalidToneError as e: logger.error(f"Tone error: {e}") @@ -467,15 +479,18 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: ) with gr.Accordion(label="ヌルモデル", open=False): with gr.Row() as null_row: - null_models_count = gr.Number(label="ヌルモデルの数", value=0, step=1) + null_models_count = gr.Number( + label="ヌルモデルの数", value=0, step=1 + ) with gr.Column(variant="panel") as null_column: + @gr.render( inputs=[ null_models_count, ] ) def render_style( - null_models_count:int, + null_models_count: int, ): global null_models_frame null_models_frame = null_models_count @@ -484,28 +499,34 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: null_model_index = gr.Number( value=i, key=f"null_model_index_{i}", - visible=False + visible=False, ) null_model_name = gr.Dropdown( label="モデル一覧", choices=model_names, key=f"null_model_name_{i}", value=model_names[initial_id], - interactive=True + 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"] + logger.debug( + f"null model parameter exists in index {i}" + ) + null_model_name.value = null_models.value[i][ + "name" + ] null_model_path = gr.Dropdown( label="モデルファイル", choices=initial_pth_files, key=f"null_model_path_{i}", value=initial_pth_files[0], - interactive=True + interactive=True, ) if i in null_models.value: - #null_model_path.choices = #ToDo - null_model_path.value=null_models.value[i]["path"] + # null_model_path.choices = #ToDo + null_model_path.value = null_models.value[i][ + "path" + ] null_voice_weights = gr.Slider( minimum=0, maximum=1, @@ -513,10 +534,12 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: step=0.1, key=f"null_voice_weights_{i}", label="声質", - interactive=True + interactive=True, ) if i in null_models.value: - null_voice_weights.value=null_models.value[i]["weight"] + null_voice_weights.value = null_models.value[i][ + "weight" + ] null_voice_pitch_weights = gr.Slider( minimum=0, maximum=1, @@ -524,10 +547,12 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: step=0.1, key=f"null_voice_pitch_weights_{i}", label="声の高さ", - interactive=True + interactive=True, ) if i in null_models.value: - null_voice_pitch_weights.value=null_models.value[i]["pitch"] + null_voice_pitch_weights.value = ( + null_models.value[i]["pitch"] + ) null_speech_style_weights = gr.Slider( minimum=0, maximum=1, @@ -535,10 +560,12 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: step=0.1, key=f"null_speech_style_weights_{i}", label="話し方", - interactive=True + interactive=True, ) if i in null_models.value: - null_speech_style_weights.value=null_models.value[i]["style"] + null_speech_style_weights.value = ( + null_models.value[i]["style"] + ) null_tempo_weights = gr.Slider( minimum=0, maximum=1, @@ -546,48 +573,93 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: step=0.1, key=f"null_tempo_weights_{i}", label="テンポ", - interactive=True + interactive=True, ) if i in null_models.value: - null_tempo_weights.value=null_models.value[i]["tempo"] + null_tempo_weights.value = null_models.value[i][ + "tempo" + ] null_model_name.change( model_holder.update_model_files_for_gradio, inputs=[null_model_name], outputs=[null_model_path], ) - #null_model_name.change(model_holder.refresh, outputs=[]) - 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_model_name.change(model_holder.refresh, outputs=[]) + null_model_path.change( + make_non_interactive, outputs=[tts_button] ) - 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_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_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_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_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_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_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] + 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( @@ -659,7 +731,7 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: pitch_scale, intonation_scale, null_models, - force_reload_model + force_reload_model, ], outputs=[text_output, audio_output, tone, force_reload_model], ) @@ -691,10 +763,12 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: return app + if __name__ == "__main__": - from config import get_path_config import torch + from config import get_path_config + path_config = get_path_config() assets_root = path_config.assets_root device = "cuda" if torch.cuda.is_available() else "cpu" diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index 77db3dd..96c68cc 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -61,7 +61,7 @@ class TTSModel: self.model_path: Path = model_path self.device: str = device - self.null_model_params: dict[int, dict[str,Union[float, str]]] = {} + self.null_model_params: dict[int, dict[str, Union[float, str]]] = {} # ハイパーパラメータの Pydantic モデルが直接指定された if isinstance(config_path, HyperParameters): @@ -114,9 +114,9 @@ class TTSModel: device=self.device, hps=self.hyper_parameters, ) - if(len(self.null_model_params.keys())==0): + if len(self.null_model_params.keys()) == 0: return - + for index, null_model in enumerate(self.null_model_params.keys()): null_model_add = get_net_g( model_path=str(self.null_model_params[index]["path"]), @@ -124,26 +124,39 @@ class TTSModel: device=self.device, hps=self.hyper_parameters, ) - #愚直。もっと上手い方法ありそう + # 愚直。もっと上手い方法ありそう print(str(self.null_model_params[index]["weight"])) 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(self.null_model_params[index]["weight"])) - params = zip(self.__net_g.flow.parameters(), null_model_add.flow.parameters()) + v[0].data.add_( + v[1].data, alpha=float(self.null_model_params[index]["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(self.null_model_params[index]["pitch"])) + v[0].data.add_( + v[1].data, alpha=float(self.null_model_params[index]["pitch"]) + ) - params = zip(self.__net_g.enc_p.parameters(), null_model_add.enc_p.parameters()) + 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(self.null_model_params[index]["style"])) - #テンポはsdpとdp二つあるからとりあえずどっちも足す + v[0].data.add_( + v[1].data, alpha=float(self.null_model_params[index]["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(self.null_model_params[index]["tempo"])) + v[0].data.add_( + v[1].data, alpha=float(self.null_model_params[index]["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(self.null_model_params[index]["tempo"])) - + v[0].data.add_( + 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]: """ @@ -258,8 +271,8 @@ class TTSModel: given_tone: Optional[list[int]] = None, pitch_scale: float = 1.0, intonation_scale: float = 1.0, - null_model_params: dict[int,dict[str,Union[str, float]]] = {}, - force_reload_model:bool = False + null_model_params: dict[int, dict[str, Union[str, float]]] = {}, + force_reload_model: bool = False, ) -> tuple[int, NDArray[Any]]: """ テキストから音声を合成する。 From 402346e4932743797e7ab3f0fbf53789c6787e66 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Sat, 22 Jun 2024 15:18:57 +0900 Subject: [PATCH 07/13] Fix mid_result assignment --- gradio_tabs/inference.py | 2 +- style_bert_vits2/tts_model.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/gradio_tabs/inference.py b/gradio_tabs/inference.py index 2673151..3f93a15 100644 --- a/gradio_tabs/inference.py +++ b/gradio_tabs/inference.py @@ -223,7 +223,7 @@ def change_null_model_row( mid_result = {} mid_result["name"] = null_model_name mid_result["path"] = null_model_path - mid_result["weight"] = null_tempo_weights + mid_result["weight"] = null_voice_weights mid_result["pitch"] = null_voice_pitch_weights mid_result["style"] = null_speech_style_weights mid_result["tempo"] = null_tempo_weights diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index 96c68cc..2ad6f4b 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -297,7 +297,7 @@ class TTSModel: given_tone (Optional[list[int]], optional): アクセントのトーンのリスト. Defaults to None. pitch_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。 Returns: tuple[int, NDArray[Any]]: サンプリングレートと音声データ (16bit PCM) """ From c1381b2c495999212d859e2c3465be4e9bfb4d95 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Sat, 22 Jun 2024 17:10:44 +0900 Subject: [PATCH 08/13] Refactor a little --- gradio_tabs/inference.py | 92 +++++++++++------------------------ style_bert_vits2/tts_model.py | 27 ++++------ 2 files changed, 37 insertions(+), 82 deletions(-) diff --git a/gradio_tabs/inference.py b/gradio_tabs/inference.py index 3f93a15..27dafc5 100644 --- a/gradio_tabs/inference.py +++ b/gradio_tabs/inference.py @@ -219,21 +219,19 @@ def change_null_model_row( null_tempo_weights: float, null_models: dict[int, dict[str, Any]], ): - # logger.debug("change_null_model_row:sta"+str(null_models)) - mid_result = {} - mid_result["name"] = null_model_name - mid_result["path"] = null_model_path - mid_result["weight"] = null_voice_weights - mid_result["pitch"] = null_voice_pitch_weights - mid_result["style"] = null_speech_style_weights - mid_result["tempo"] = null_tempo_weights - null_models[null_model_index] = mid_result - # logger.debug("decreasing:"+str(null_models_frame)+":"+str(len(null_models.keys()))) - if null_models_frame < len(null_models.keys()): - for i in range(null_models_frame, len(null_models.keys())): - _ = null_models.pop(i, None) - result = null_models - # logger.debug("change_null_model_row:res"+str(null_models)) + 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 @@ -265,6 +263,7 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: ): model_holder.get_model(model_name, model_path) assert model_holder.current_model is not None + logger.debug(f"Null models setting: {null_models}") wrong_tone_message = "" 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 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 if len(model_names) == 0: logger.error( @@ -355,9 +357,7 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: ) return app initial_id = 0 - initial_pth_files = [ - str(f) for f in model_holder.model_files_dict[model_names[initial_id]] - ] + initial_pth_files = get_model_files(model_names[initial_id]) with gr.Blocks(theme=GRADIO_THEME) as app: gr.Markdown(initial_md) @@ -478,23 +478,19 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: outputs=[assist_text, assist_text_weight], ) with gr.Accordion(label="ヌルモデル", open=False): - with gr.Row() as null_row: + with gr.Row(): null_models_count = gr.Number( label="ヌルモデルの数", value=0, step=1 ) - with gr.Column(variant="panel") as null_column: + with gr.Column(variant="panel"): - @gr.render( - inputs=[ - null_models_count, - ] - ) - def render_style( + @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(0, null_models_count): + for i in range(null_models_count): with gr.Row(): null_model_index = gr.Number( value=i, @@ -506,27 +502,15 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: choices=model_names, key=f"null_model_name_{i}", 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( label="モデルファイル", - choices=initial_pth_files, key=f"null_model_path_{i}", - value=initial_pth_files[0], - interactive=True, + # FIXME: 再レンダー時に選択肢が消えるのでどうにかしたい + # 現在は再レンダーで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( minimum=0, maximum=1, @@ -534,12 +518,7 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: step=0.1, key=f"null_voice_weights_{i}", label="声質", - interactive=True, ) - if i in null_models.value: - null_voice_weights.value = null_models.value[i][ - "weight" - ] null_voice_pitch_weights = gr.Slider( minimum=0, maximum=1, @@ -547,12 +526,7 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: step=0.1, key=f"null_voice_pitch_weights_{i}", 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( minimum=0, maximum=1, @@ -560,12 +534,7 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: step=0.1, key=f"null_speech_style_weights_{i}", label="話し方", - interactive=True, ) - if i in null_models.value: - null_speech_style_weights.value = ( - null_models.value[i]["style"] - ) null_tempo_weights = gr.Slider( minimum=0, maximum=1, @@ -573,18 +542,13 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks: step=0.1, key=f"null_tempo_weights_{i}", label="テンポ", - interactive=True, ) - if i in null_models.value: - null_tempo_weights.value = null_models.value[i][ - "tempo" - ] + null_model_name.change( model_holder.update_model_files_for_gradio, inputs=[null_model_name], outputs=[null_model_path], ) - # null_model_name.change(model_holder.refresh, outputs=[]) null_model_path.change( make_non_interactive, outputs=[tts_button] ) diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index 2ad6f4b..5ba2586 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -117,46 +117,36 @@ class TTSModel: if len(self.null_model_params.keys()) == 0: 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( - model_path=str(self.null_model_params[index]["path"]), + model_path=str(null_model_info["path"]), version=self.hyper_parameters.version, device=self.device, hps=self.hyper_parameters, ) # 愚直。もっと上手い方法ありそう - print(str(self.null_model_params[index]["weight"])) 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(self.null_model_params[index]["weight"]) - ) + 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(self.null_model_params[index]["pitch"]) - ) + 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(self.null_model_params[index]["style"]) - ) + 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(self.null_model_params[index]["tempo"]) - ) + 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(self.null_model_params[index]["tempo"]) - ) + v[0].data.add_(v[1].data, alpha=float(null_model_info["tempo"])) 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. intonation_scale (float, optional): 抑揚の平均からの変化幅 (1.0 から変更すると若干音質が低下する). Defaults to 1.0. null_model_params (dict[int, dict[str, Union[str, float]]], optional): 推論時に使用するヌルモデルの名前、適用割合のdictが入ったdict。 + force_reload_model (bool, optional): モデルを強制的に再ロードするかどうか. Defaults to False. Returns: tuple[int, NDArray[Any]]: サンプリングレートと音声データ (16bit PCM) """ From e78d2a6f7e21c574408653af143c60feeeae6e07 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Sat, 22 Jun 2024 17:21:58 +0900 Subject: [PATCH 09/13] Feat: use DistributedLengthGroupedSampler (experimental, to be checked) --- train_ms.py | 14 +++++++++++--- train_ms_jp_extra.py | 13 +++++++++++-- 2 files changed, 22 insertions(+), 5 deletions(-) diff --git a/train_ms.py b/train_ms.py index 1adee4e..b1cc3a4 100644 --- a/train_ms.py +++ b/train_ms.py @@ -13,6 +13,7 @@ from torch.nn.parallel import DistributedDataParallel as DDP from torch.utils.data import DataLoader from torch.utils.tensorboard import SummaryWriter from tqdm import tqdm +from transformers.trainer_pt_utils import DistributedLengthGroupedSampler # logging.getLogger("numba").setLevel(logging.WARNING) import default_style @@ -35,7 +36,6 @@ from style_bert_vits2.models.models import ( from style_bert_vits2.nlp.symbols import SYMBOLS from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT - torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = ( True # If encontered training problem,please try to disable TF32. @@ -242,15 +242,23 @@ def run(): # prefetch_factor=6, ) 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_dataset, # メモリ消費量を減らそうとnum_workersを1にしてみる # num_workers=min(config.train_ms_config.num_workers, os.cpu_count() // 2), num_workers=1, - shuffle=True, + # shuffle=True, pin_memory=True, collate_fn=collate_fn, - # batch_sampler=train_sampler, + batch_sampler=train_sampler, batch_size=hps.train.batch_size, persistent_workers=True, # これもメモリ消費量を減らそうとしてコメントアウト diff --git a/train_ms_jp_extra.py b/train_ms_jp_extra.py index e5c5bd1..c2539ca 100644 --- a/train_ms_jp_extra.py +++ b/train_ms_jp_extra.py @@ -13,6 +13,7 @@ from torch.nn.parallel import DistributedDataParallel as DDP from torch.utils.data import DataLoader from torch.utils.tensorboard import SummaryWriter from tqdm import tqdm +from transformers.trainer_pt_utils import DistributedLengthGroupedSampler # logging.getLogger("numba").setLevel(logging.WARNING) import default_style @@ -243,15 +244,23 @@ def run(): # prefetch_factor=6, ) 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_dataset, # メモリ消費量を減らそうとnum_workersを1にしてみる # num_workers=min(config.train_ms_config.num_workers, os.cpu_count() // 2), num_workers=1, - shuffle=True, + # shuffle=True, pin_memory=True, collate_fn=collate_fn, - # batch_sampler=train_sampler, + batch_sampler=train_sampler, batch_size=hps.train.batch_size, persistent_workers=True, # これもメモリ消費量を減らそうとしてコメントアウト From 1671967068d6a972c07e98cbd8bed56d5ed3a908 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Sat, 22 Jun 2024 17:24:21 +0900 Subject: [PATCH 10/13] Style --- train_ms.py | 1 + 1 file changed, 1 insertion(+) diff --git a/train_ms.py b/train_ms.py index b1cc3a4..9f8ab5f 100644 --- a/train_ms.py +++ b/train_ms.py @@ -36,6 +36,7 @@ from style_bert_vits2.models.models import ( from style_bert_vits2.nlp.symbols import SYMBOLS from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT + torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = ( True # If encontered training problem,please try to disable TF32. From 8832cfb4c922f8daa95976a390284a1f21125ad6 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Fri, 28 Jun 2024 09:19:28 +0900 Subject: [PATCH 11/13] Add playgrounds ignore --- .gitignore | 1 + 1 file changed, 1 insertion(+) diff --git a/.gitignore b/.gitignore index 6ca8d26..099d525 100644 --- a/.gitignore +++ b/.gitignore @@ -41,3 +41,4 @@ safetensors.ipynb *.dic playground.ipynb +playgrounds/ From b6d09bad445de51a517018c060b33e243f58f65f Mon Sep 17 00:00:00 2001 From: litagin02 Date: Fri, 28 Jun 2024 09:34:12 +0900 Subject: [PATCH 12/13] Fix: use sampler instead of batch_sampler for non-custom --- train_ms.py | 2 +- train_ms_jp_extra.py | 6 +++++- 2 files changed, 6 insertions(+), 2 deletions(-) diff --git a/train_ms.py b/train_ms.py index 9f8ab5f..7244bb8 100644 --- a/train_ms.py +++ b/train_ms.py @@ -259,7 +259,7 @@ def run(): # shuffle=True, pin_memory=True, collate_fn=collate_fn, - batch_sampler=train_sampler, + sampler=train_sampler, batch_size=hps.train.batch_size, persistent_workers=True, # これもメモリ消費量を減らそうとしてコメントアウト diff --git a/train_ms_jp_extra.py b/train_ms_jp_extra.py index c2539ca..8c33858 100644 --- a/train_ms_jp_extra.py +++ b/train_ms_jp_extra.py @@ -260,12 +260,16 @@ def run(): # shuffle=True, pin_memory=True, collate_fn=collate_fn, - batch_sampler=train_sampler, + sampler=train_sampler, batch_size=hps.train.batch_size, persistent_workers=True, # これもメモリ消費量を減らそうとしてコメントアウト # 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_loader = None if rank == 0 and not args.speedup: From ffa393a3d0e50b7e3b8729b934d5d737e628f2f5 Mon Sep 17 00:00:00 2001 From: Risenafis Date: Sat, 6 Jul 2024 03:47:41 +0900 Subject: [PATCH 13/13] set spk2id when merge --- gradio_tabs/merge.py | 16 ++++++++++++++++ 1 file changed, 16 insertions(+) diff --git a/gradio_tabs/merge.py b/gradio_tabs/merge.py index b4c9204..c3e7b17 100644 --- a/gradio_tabs/merge.py +++ b/gradio_tabs/merge.py @@ -104,6 +104,8 @@ def merge_style_usual( new_config = config_a.copy() new_config["data"]["num_styles"] = len(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 save_config(new_config, output_name) @@ -159,6 +161,8 @@ def merge_style_add_diff( new_config = config_a.copy() new_config["data"]["num_styles"] = len(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 save_config(new_config, output_name) @@ -218,6 +222,8 @@ def merge_style_weighted_sum( new_config = config_a.copy() new_config["data"]["num_styles"] = len(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 save_config(new_config, output_name) @@ -267,6 +273,8 @@ def merge_style_add_null( new_config = config_a.copy() new_config["data"]["num_styles"] = len(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 save_config(new_config, output_name) @@ -361,6 +369,8 @@ def merge_models_usual( new_config["model_name"] = output_name new_config["data"]["num_styles"] = 1 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) neutral_vector_a = style_vectors_a[0] @@ -443,6 +453,8 @@ def merge_models_add_diff( new_config["model_name"] = output_name new_config["data"]["num_styles"] = 1 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: json.dump(new_config, f, indent=2, ensure_ascii=False) @@ -518,6 +530,8 @@ def merge_models_weighted_sum( new_config["model_name"] = output_name new_config["data"]["num_styles"] = 1 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: json.dump(new_config, f, indent=2, ensure_ascii=False) @@ -594,6 +608,8 @@ def merge_models_add_null( new_config["model_name"] = output_name new_config["data"]["num_styles"] = 1 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: json.dump(new_config, f, indent=2, ensure_ascii=False)