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_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",

View File

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

View File

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

View File

@@ -55,7 +55,7 @@ Optional
## 2. Preprocess
```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:
@@ -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).

View File

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

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():
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()),

View File

@@ -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()),

View File

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