Feat: freezing decoder option

This commit is contained in:
litagin02
2024-02-23 20:57:54 +09:00
parent 20b3326b79
commit ffdbb8e0d1
9 changed files with 130 additions and 19 deletions

View File

@@ -265,6 +265,7 @@
" freeze_JP_bert=False,\n", " freeze_JP_bert=False,\n",
" freeze_ZH_bert=False,\n", " freeze_ZH_bert=False,\n",
" freeze_style=False,\n", " freeze_style=False,\n",
" freeze_decoder=False, # ここをTrueにするともしかしたら違う結果になるかもしれません。\n",
" use_jp_extra=use_jp_extra,\n", " use_jp_extra=use_jp_extra,\n",
" val_per_lang=0,\n", " val_per_lang=0,\n",
" log_interval=200,\n", " log_interval=200,\n",

View File

@@ -20,7 +20,8 @@
"freeze_ZH_bert": false, "freeze_ZH_bert": false,
"freeze_JP_bert": false, "freeze_JP_bert": false,
"freeze_EN_bert": false, "freeze_EN_bert": false,
"freeze_style": false "freeze_style": false,
"freeze_encoder": false
}, },
"data": { "data": {
"training_files": "Data/your_model_name/filelists/train.list", "training_files": "Data/your_model_name/filelists/train.list",

View File

@@ -22,7 +22,8 @@
"freeze_JP_bert": false, "freeze_JP_bert": false,
"freeze_EN_bert": false, "freeze_EN_bert": false,
"freeze_emo": false, "freeze_emo": false,
"freeze_style": false "freeze_style": false,
"freeze_decoder": false
}, },
"data": { "data": {
"use_jp_extra": true, "use_jp_extra": true,

View File

@@ -55,7 +55,7 @@ Optional
## 2. Preprocess ## 2. Preprocess
```bash ```bash
python preprocess_all.py -m <model_name> [--use_jp_extra] [-b <batch_size>] [-e <epochs>] [-s <save_every_steps>] [--num_processes <num_processes>] [--normalize] [--trim] [--val_per_lang <val_per_lang>] [--log_interval <log_interval>] [--freeze_EN_bert] [--freeze_JP_bert] [--freeze_ZH_bert] [--freeze_style] python preprocess_all.py -m <model_name> [--use_jp_extra] [-b <batch_size>] [-e <epochs>] [-s <save_every_steps>] [--num_processes <num_processes>] [--normalize] [--trim] [--val_per_lang <val_per_lang>] [--log_interval <log_interval>] [--freeze_EN_bert] [--freeze_JP_bert] [--freeze_ZH_bert] [--freeze_style] [--freeze_decoder]
``` ```
Required: Required:
@@ -72,6 +72,7 @@ Optional:
- `--freeze_JP_bert`: Freeze Japanese BERT. - `--freeze_JP_bert`: Freeze Japanese BERT.
- `--freeze_ZH_bert`: Freeze Chinese BERT. - `--freeze_ZH_bert`: Freeze Chinese BERT.
- `--freeze_style`: Freeze style vector. - `--freeze_style`: Freeze style vector.
- `--freeze_decoder`: Freeze decoder.
- `--use_jp_extra`: Use JP-Extra model. - `--use_jp_extra`: Use JP-Extra model.
- `--val_per_lang`: Validation data per language (default: 0). - `--val_per_lang`: Validation data per language (default: 0).
- `--log_interval`: Log interval (default: 200). - `--log_interval`: Log interval (default: 200).

View File

@@ -52,6 +52,11 @@ if __name__ == "__main__":
action="store_true", action="store_true",
help="Freeze style vector", help="Freeze style vector",
) )
parser.add_argument(
"--freeze_decoder",
action="store_true",
help="Freeze decoder",
)
parser.add_argument( parser.add_argument(
"--use_jp_extra", "--use_jp_extra",
action="store_true", action="store_true",
@@ -84,6 +89,7 @@ if __name__ == "__main__":
freeze_JP_bert=args.freeze_JP_bert, freeze_JP_bert=args.freeze_JP_bert,
freeze_ZH_bert=args.freeze_ZH_bert, freeze_ZH_bert=args.freeze_ZH_bert,
freeze_style=args.freeze_style, freeze_style=args.freeze_style,
freeze_decoder=args.freeze_decoder,
use_jp_extra=args.use_jp_extra, use_jp_extra=args.use_jp_extra,
val_per_lang=args.val_per_lang, val_per_lang=args.val_per_lang,
log_interval=args.log_interval, log_interval=args.log_interval,

View File

@@ -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

View File

@@ -306,6 +306,11 @@ def run():
for param in net_g.enc_p.style_proj.parameters(): for param in net_g.enc_p.style_proj.parameters():
param.requires_grad = False 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) net_d = MultiPeriodDiscriminator(hps.model.use_spectral_norm).cuda(local_rank)
optim_g = torch.optim.AdamW( optim_g = torch.optim.AdamW(
filter(lambda p: p.requires_grad, net_g.parameters()), filter(lambda p: p.requires_grad, net_g.parameters()),

View File

@@ -310,6 +310,11 @@ def run():
for param in net_g.enc_p.style_proj.parameters(): for param in net_g.enc_p.style_proj.parameters():
param.requires_grad = False 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) net_d = MultiPeriodDiscriminator(hps.model.use_spectral_norm).cuda(local_rank)
optim_g = torch.optim.AdamW( optim_g = torch.optim.AdamW(
filter(lambda p: p.requires_grad, net_g.parameters()), filter(lambda p: p.requires_grad, net_g.parameters()),

View File

@@ -47,6 +47,7 @@ def initialize(
freeze_JP_bert, freeze_JP_bert,
freeze_ZH_bert, freeze_ZH_bert,
freeze_style, freeze_style,
freeze_decoder,
use_jp_extra, use_jp_extra,
log_interval, log_interval,
): ):
@@ -61,7 +62,7 @@ def initialize(
logger_handler = logger.add(os.path.join(dataset_path, file_name)) logger_handler = logger.add(os.path.join(dataset_path, file_name))
logger.info( 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 = ( default_config_path = (
@@ -82,12 +83,15 @@ def initialize(
config["train"]["freeze_JP_bert"] = freeze_JP_bert config["train"]["freeze_JP_bert"] = freeze_JP_bert
config["train"]["freeze_ZH_bert"] = freeze_ZH_bert config["train"]["freeze_ZH_bert"] = freeze_ZH_bert
config["train"]["freeze_style"] = freeze_style config["train"]["freeze_style"] = freeze_style
config["train"]["freeze_decoder"] = freeze_decoder
config["train"]["bf16_run"] = False # デフォルトでFalseのはずだが念のため config["train"]["bf16_run"] = False # デフォルトでFalseのはずだが念のため
model_path = os.path.join(dataset_path, "models") model_path = os.path.join(dataset_path, "models")
if os.path.exists(model_path): 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( shutil.copytree(
src=model_path, src=model_path,
dst=os.path.join(dataset_path, "models_backup"), dst=os.path.join(dataset_path, "models_backup"),
@@ -262,6 +266,7 @@ def preprocess_all(
freeze_JP_bert, freeze_JP_bert,
freeze_ZH_bert, freeze_ZH_bert,
freeze_style, freeze_style,
freeze_decoder,
use_jp_extra, use_jp_extra,
val_per_lang, val_per_lang,
log_interval, log_interval,
@@ -269,29 +274,39 @@ def preprocess_all(
if model_name == "": if model_name == "":
return False, "Error: モデル名を入力してください" return False, "Error: モデル名を入力してください"
success, message = initialize( success, message = initialize(
model_name, model_name=model_name,
batch_size, batch_size=batch_size,
epochs, epochs=epochs,
save_every_steps, save_every_steps=save_every_steps,
freeze_EN_bert, freeze_EN_bert=freeze_EN_bert,
freeze_JP_bert, freeze_JP_bert=freeze_JP_bert,
freeze_ZH_bert, freeze_ZH_bert=freeze_ZH_bert,
freeze_style, freeze_style=freeze_style,
use_jp_extra, freeze_decoder=freeze_decoder,
log_interval, use_jp_extra=use_jp_extra,
log_interval=log_interval,
) )
if not success: if not success:
return False, message 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: if not success:
return False, message 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: if not success:
return False, message return False, message
success, message = bert_gen(model_name) # bert_genは重いのでプロセス数いじらない success, message = bert_gen(
model_name=model_name
) # bert_genは重いのでプロセス数いじらない
if not success: if not success:
return False, message 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: if not success:
return False, message return False, message
logger.success("Success: All preprocess finished!") logger.success("Success: All preprocess finished!")
@@ -507,6 +522,10 @@ if __name__ == "__main__":
label="スタイル部分を凍結", label="スタイル部分を凍結",
value=False, value=False,
) )
freeze_decoder = gr.Checkbox(
label="デコーダ部分を凍結",
value=False,
)
with gr.Column(): with gr.Column():
preprocess_button = gr.Button( preprocess_button = gr.Button(
@@ -565,6 +584,10 @@ if __name__ == "__main__":
label="スタイル部分を凍結", label="スタイル部分を凍結",
value=False, value=False,
) )
freeze_decoder_manual = gr.Checkbox(
label="デコーダ部分を凍結",
value=False,
)
with gr.Column(): with gr.Column():
generate_config_btn = gr.Button(value="実行", variant="primary") generate_config_btn = gr.Button(value="実行", variant="primary")
info_init = gr.Textbox(label="状況") info_init = gr.Textbox(label="状況")
@@ -658,6 +681,7 @@ if __name__ == "__main__":
freeze_JP_bert, freeze_JP_bert,
freeze_ZH_bert, freeze_ZH_bert,
freeze_style, freeze_style,
freeze_decoder,
use_jp_extra, use_jp_extra,
val_per_lang, val_per_lang,
log_interval, log_interval,
@@ -677,6 +701,7 @@ if __name__ == "__main__":
freeze_JP_bert_manual, freeze_JP_bert_manual,
freeze_ZH_bert_manual, freeze_ZH_bert_manual,
freeze_style_manual, freeze_style_manual,
freeze_decoder_manual,
use_jp_extra_manual, use_jp_extra_manual,
log_interval_manual, log_interval_manual,
], ],