From ffdbb8e0d18cc0c6ab84b4d36592e7abe3eb082b Mon Sep 17 00:00:00 2001 From: litagin02 Date: Fri, 23 Feb 2024 20:57:54 +0900 Subject: [PATCH] Feat: freezing decoder option --- colab.ipynb | 1 + configs/config.json | 3 +- configs/configs_jp_extra.json | 3 +- docs/CLI.md | 3 +- preprocess_all.py | 6 +++ scripts/Update-to-Dict-Editor.bat | 66 +++++++++++++++++++++++++++++++ train_ms.py | 5 +++ train_ms_jp_extra.py | 5 +++ webui_train.py | 57 ++++++++++++++++++-------- 9 files changed, 130 insertions(+), 19 deletions(-) create mode 100644 scripts/Update-to-Dict-Editor.bat diff --git a/colab.ipynb b/colab.ipynb index 37b708f..f4e932b 100644 --- a/colab.ipynb +++ b/colab.ipynb @@ -265,6 +265,7 @@ " freeze_JP_bert=False,\n", " freeze_ZH_bert=False,\n", " freeze_style=False,\n", + " freeze_decoder=False, # ここをTrueにするともしかしたら違う結果になるかもしれません。\n", " use_jp_extra=use_jp_extra,\n", " val_per_lang=0,\n", " log_interval=200,\n", diff --git a/configs/config.json b/configs/config.json index 5196633..25e86db 100644 --- a/configs/config.json +++ b/configs/config.json @@ -20,7 +20,8 @@ "freeze_ZH_bert": false, "freeze_JP_bert": false, "freeze_EN_bert": false, - "freeze_style": false + "freeze_style": false, + "freeze_encoder": false }, "data": { "training_files": "Data/your_model_name/filelists/train.list", diff --git a/configs/configs_jp_extra.json b/configs/configs_jp_extra.json index 490f4d6..616d1d3 100644 --- a/configs/configs_jp_extra.json +++ b/configs/configs_jp_extra.json @@ -22,7 +22,8 @@ "freeze_JP_bert": false, "freeze_EN_bert": false, "freeze_emo": false, - "freeze_style": false + "freeze_style": false, + "freeze_decoder": false }, "data": { "use_jp_extra": true, diff --git a/docs/CLI.md b/docs/CLI.md index 88ad5fd..08e2fd0 100644 --- a/docs/CLI.md +++ b/docs/CLI.md @@ -55,7 +55,7 @@ Optional ## 2. Preprocess ```bash -python preprocess_all.py -m [--use_jp_extra] [-b ] [-e ] [-s ] [--num_processes ] [--normalize] [--trim] [--val_per_lang ] [--log_interval ] [--freeze_EN_bert] [--freeze_JP_bert] [--freeze_ZH_bert] [--freeze_style] +python preprocess_all.py -m [--use_jp_extra] [-b ] [-e ] [-s ] [--num_processes ] [--normalize] [--trim] [--val_per_lang ] [--log_interval ] [--freeze_EN_bert] [--freeze_JP_bert] [--freeze_ZH_bert] [--freeze_style] [--freeze_decoder] ``` Required: @@ -72,6 +72,7 @@ Optional: - `--freeze_JP_bert`: Freeze Japanese BERT. - `--freeze_ZH_bert`: Freeze Chinese BERT. - `--freeze_style`: Freeze style vector. +- `--freeze_decoder`: Freeze decoder. - `--use_jp_extra`: Use JP-Extra model. - `--val_per_lang`: Validation data per language (default: 0). - `--log_interval`: Log interval (default: 200). diff --git a/preprocess_all.py b/preprocess_all.py index 679abb2..82a3c20 100644 --- a/preprocess_all.py +++ b/preprocess_all.py @@ -52,6 +52,11 @@ if __name__ == "__main__": action="store_true", help="Freeze style vector", ) + parser.add_argument( + "--freeze_decoder", + action="store_true", + help="Freeze decoder", + ) parser.add_argument( "--use_jp_extra", action="store_true", @@ -84,6 +89,7 @@ if __name__ == "__main__": freeze_JP_bert=args.freeze_JP_bert, freeze_ZH_bert=args.freeze_ZH_bert, freeze_style=args.freeze_style, + freeze_decoder=args.freeze_decoder, use_jp_extra=args.use_jp_extra, val_per_lang=args.val_per_lang, log_interval=args.log_interval, diff --git a/scripts/Update-to-Dict-Editor.bat b/scripts/Update-to-Dict-Editor.bat new file mode 100644 index 0000000..254b060 --- /dev/null +++ b/scripts/Update-to-Dict-Editor.bat @@ -0,0 +1,66 @@ +chcp 65001 > NUL +@echo off + +pushd %~dp0 +set PS_CMD=PowerShell -Version 5.1 -ExecutionPolicy Bypass + +set CURL_CMD=C:\Windows\System32\curl.exe +if not exist %CURL_CMD% ( + echo [ERROR] %CURL_CMD% が見つかりません。 + pause & popd & exit /b 1 +) + +@REM Style-Bert-VITS2.zip をGitHubのmasterの最新のものをダウンロード +%CURL_CMD% -Lo Style-Bert-VITS2.zip^ + https://github.com/litagin02/Style-Bert-VITS2/archive/refs/heads/master.zip +if %errorlevel% neq 0 ( pause & popd & exit /b %errorlevel% ) + +@REM Style-Bert-VITS2.zip を解凍(フォルダ名前がBert-VITS2-masterになる) +%PS_CMD% Expand-Archive -Path Style-Bert-VITS2.zip -DestinationPath . -Force +if %errorlevel% neq 0 ( pause & popd & exit /b %errorlevel% ) + +@REM 元のzipを削除 +del Style-Bert-VITS2.zip +if %errorlevel% neq 0 ( pause & popd & exit /b %errorlevel% ) + +@REM Bert-VITS2-masterの中身をStyle-Bert-VITS2に上書き移動 +xcopy /QSY .\Style-Bert-VITS2-master\ .\Style-Bert-VITS2\ +rmdir /s /q Style-Bert-VITS2-master +if %errorlevel% neq 0 ( pause & popd & exit /b %errorlevel% ) + +@REM 仮想環境のpip requirements.txtを更新 + +echo call .\Style-Bert-VITS2\scripts\activate.bat +call .\Style-Bert-VITS2\venv\Scripts\activate.bat +if %errorlevel% neq 0 ( pause & popd & exit /b %errorlevel% ) + +@REM pyopenjtalk-prebuiltやpyopenjtalkが入っていたら削除 +echo pip uninstall -y pyopenjtalk-prebuilt pyopenjtalk +pip uninstall -y pyopenjtalk-prebuilt pyopenjtalk +if %errorlevel% neq 0 ( pause & popd & exit /b %errorlevel% ) + +@REM pyopenjtalk-dictをインストール +echo pip install -U pyopenjtalk-dict +pip install -U pyopenjtalk-dict + +@REM その他のrequirements.txtも一応更新 +pip install -U -r Style-Bert-VITS2\requirements.txt +if %errorlevel% neq 0 ( pause & popd & exit /b %errorlevel% ) + +pushd Style-Bert-VITS2 + +echo Update completed. Running Style-Bert-VITS2 Editor... + +@REM Style-Bert-VITS2 Editorを起動 +python server_editor.py + +pause + +popd + + +pause + +popd + +popd \ No newline at end of file diff --git a/train_ms.py b/train_ms.py index fb7fa9c..ea4b65f 100644 --- a/train_ms.py +++ b/train_ms.py @@ -306,6 +306,11 @@ def run(): for param in net_g.enc_p.style_proj.parameters(): param.requires_grad = False + if getattr(hps.train, "freeze_decoder", False): + logger.info("Freezing decoder !!!") + for param in net_g.dec.parameters(): + param.requires_grad = False + net_d = MultiPeriodDiscriminator(hps.model.use_spectral_norm).cuda(local_rank) optim_g = torch.optim.AdamW( filter(lambda p: p.requires_grad, net_g.parameters()), diff --git a/train_ms_jp_extra.py b/train_ms_jp_extra.py index 0cc9d2a..bae922d 100644 --- a/train_ms_jp_extra.py +++ b/train_ms_jp_extra.py @@ -310,6 +310,11 @@ def run(): for param in net_g.enc_p.style_proj.parameters(): param.requires_grad = False + if getattr(hps.train, "freeze_decoder", False): + logger.info("Freezing decoder !!!") + for param in net_g.dec.parameters(): + param.requires_grad = False + net_d = MultiPeriodDiscriminator(hps.model.use_spectral_norm).cuda(local_rank) optim_g = torch.optim.AdamW( filter(lambda p: p.requires_grad, net_g.parameters()), diff --git a/webui_train.py b/webui_train.py index c3a876e..eccfeaf 100644 --- a/webui_train.py +++ b/webui_train.py @@ -47,6 +47,7 @@ def initialize( freeze_JP_bert, freeze_ZH_bert, freeze_style, + freeze_decoder, use_jp_extra, log_interval, ): @@ -61,7 +62,7 @@ def initialize( logger_handler = logger.add(os.path.join(dataset_path, file_name)) logger.info( - f"Step 1: start initialization...\nmodel_name: {model_name}, batch_size: {batch_size}, epochs: {epochs}, save_every_steps: {save_every_steps}, freeze_ZH_bert: {freeze_ZH_bert}, freeze_JP_bert: {freeze_JP_bert}, freeze_EN_bert: {freeze_EN_bert}, freeze_style: {freeze_style}, use_jp_extra: {use_jp_extra}" + f"Step 1: start initialization...\nmodel_name: {model_name}, batch_size: {batch_size}, epochs: {epochs}, save_every_steps: {save_every_steps}, freeze_ZH_bert: {freeze_ZH_bert}, freeze_JP_bert: {freeze_JP_bert}, freeze_EN_bert: {freeze_EN_bert}, freeze_style: {freeze_style}, freeze_decoder: {freeze_decoder}, use_jp_extra: {use_jp_extra}" ) default_config_path = ( @@ -82,12 +83,15 @@ def initialize( config["train"]["freeze_JP_bert"] = freeze_JP_bert config["train"]["freeze_ZH_bert"] = freeze_ZH_bert config["train"]["freeze_style"] = freeze_style + config["train"]["freeze_decoder"] = freeze_decoder config["train"]["bf16_run"] = False # デフォルトでFalseのはずだが念のため model_path = os.path.join(dataset_path, "models") if os.path.exists(model_path): - logger.warning(f"Step 1: {model_path} already exists, so copy it to backup to {model_path}_backup") + logger.warning( + f"Step 1: {model_path} already exists, so copy it to backup to {model_path}_backup" + ) shutil.copytree( src=model_path, dst=os.path.join(dataset_path, "models_backup"), @@ -262,6 +266,7 @@ def preprocess_all( freeze_JP_bert, freeze_ZH_bert, freeze_style, + freeze_decoder, use_jp_extra, val_per_lang, log_interval, @@ -269,29 +274,39 @@ def preprocess_all( if model_name == "": return False, "Error: モデル名を入力してください" success, message = initialize( - model_name, - batch_size, - epochs, - save_every_steps, - freeze_EN_bert, - freeze_JP_bert, - freeze_ZH_bert, - freeze_style, - use_jp_extra, - log_interval, + model_name=model_name, + batch_size=batch_size, + epochs=epochs, + save_every_steps=save_every_steps, + freeze_EN_bert=freeze_EN_bert, + freeze_JP_bert=freeze_JP_bert, + freeze_ZH_bert=freeze_ZH_bert, + freeze_style=freeze_style, + freeze_decoder=freeze_decoder, + use_jp_extra=use_jp_extra, + log_interval=log_interval, ) if not success: return False, message - success, message = resample(model_name, normalize, trim, num_processes) + success, message = resample( + model_name=model_name, + normalize=normalize, + trim=trim, + num_processes=num_processes, + ) if not success: return False, message - success, message = preprocess_text(model_name, use_jp_extra, val_per_lang) + success, message = preprocess_text( + model_name=model_name, use_jp_extra=use_jp_extra, val_per_lang=val_per_lang + ) if not success: return False, message - success, message = bert_gen(model_name) # bert_genは重いのでプロセス数いじらない + success, message = bert_gen( + model_name=model_name + ) # bert_genは重いのでプロセス数いじらない if not success: return False, message - success, message = style_gen(model_name, num_processes) + success, message = style_gen(model_name=model_name, num_processes=num_processes) if not success: return False, message logger.success("Success: All preprocess finished!") @@ -507,6 +522,10 @@ if __name__ == "__main__": label="スタイル部分を凍結", value=False, ) + freeze_decoder = gr.Checkbox( + label="デコーダ部分を凍結", + value=False, + ) with gr.Column(): preprocess_button = gr.Button( @@ -565,6 +584,10 @@ if __name__ == "__main__": label="スタイル部分を凍結", value=False, ) + freeze_decoder_manual = gr.Checkbox( + label="デコーダ部分を凍結", + value=False, + ) with gr.Column(): generate_config_btn = gr.Button(value="実行", variant="primary") info_init = gr.Textbox(label="状況") @@ -658,6 +681,7 @@ if __name__ == "__main__": freeze_JP_bert, freeze_ZH_bert, freeze_style, + freeze_decoder, use_jp_extra, val_per_lang, log_interval, @@ -677,6 +701,7 @@ if __name__ == "__main__": freeze_JP_bert_manual, freeze_ZH_bert_manual, freeze_style_manual, + freeze_decoder_manual, use_jp_extra_manual, log_interval_manual, ],