Feat: val number and log interval option

This commit is contained in:
litagin02
2024-02-09 15:47:58 +09:00
parent 618d3e1626
commit 5876624f01
8 changed files with 59 additions and 11 deletions

1
.gitignore vendored
View File

@@ -27,3 +27,4 @@ venv/
/mos_results/ /mos_results/
safetensors.ipynb safetensors.ipynb
*.wav

View File

@@ -1,6 +1,6 @@
import enum import enum
LATEST_VERSION: str = "2.1" LATEST_VERSION: str = "2.1.1"
DEFAULT_STYLE: str = "Neutral" DEFAULT_STYLE: str = "Neutral"
DEFAULT_STYLE_WEIGHT: float = 5.0 DEFAULT_STYLE_WEIGHT: float = 5.0

View File

@@ -67,5 +67,5 @@
"use_spectral_norm": false, "use_spectral_norm": false,
"gin_channels": 256 "gin_channels": 256
}, },
"version": "2.1" "version": "2.1.1"
} }

View File

@@ -74,5 +74,5 @@
"initial_channel": 64 "initial_channel": 64
} }
}, },
"version": "2.1-JP-Extra" "version": "2.1.1-JP-Extra"
} }

View File

@@ -16,7 +16,7 @@ preprocess_text:
train_path: "train.list" train_path: "train.list"
val_path: "val.list" val_path: "val.list"
config_path: "config.json" config_path: "config.json"
val_per_lang: 4 val_per_lang: 0
max_val_total: 12 max_val_total: 12
clean: true clean: true

View File

@@ -1,5 +1,16 @@
# Changelog # Changelog
## v2.2
### 変更・機能追加
- bfloat16オプションはデメリットしか無さそうということで、常にオフで学習するよう変更
- 学習の際の検証データ数をデフォルトで0に変更し、また検証データ数を学習用WebUIで指定できるようにした
- Tensorboardのログ間隔を学習用WebUIで指定できるようにした
### バグ修正等
- 「こんにちは!?!?!?!?」等、感嘆符等の記号が連続すると学習・音声合成でエラーになるバグを修正
- `—` (em dash, U+2014) や `―` (quotation dash, U+2015) 等のダッシュやハイフンの各種変種が、種類によって「-」に正規化されたりされていなかったりする処理を、全て「-」に正規化するように修正
## v2.1 (2024-02-07) ## v2.1 (2024-02-07)
### 変更 ### 変更

View File

@@ -600,9 +600,6 @@ if __name__ == "__main__":
"./bert/deberta-v2-large-japanese-char-wwm" "./bert/deberta-v2-large-japanese-char-wwm"
) )
text = "こんにちは、世界。" text = "こんにちは、世界。"
print("Original:")
print("NFKC:")
text = unicodedata.normalize("NFKC", text)
from text.japanese_bert import get_bert_feature from text.japanese_bert import get_bert_feature
text = text_normalize(text) text = text_normalize(text)

View File

@@ -48,6 +48,7 @@ def initialize(
freeze_ZH_bert, freeze_ZH_bert,
freeze_style, freeze_style,
use_jp_extra, use_jp_extra,
log_interval,
): ):
global logger_handler global logger_handler
dataset_path, _, train_path, val_path, config_path = get_path(model_name) dataset_path, _, train_path, val_path, config_path = get_path(model_name)
@@ -75,12 +76,15 @@ def initialize(
config["train"]["batch_size"] = batch_size config["train"]["batch_size"] = batch_size
config["train"]["epochs"] = epochs config["train"]["epochs"] = epochs
config["train"]["eval_interval"] = save_every_steps config["train"]["eval_interval"] = save_every_steps
config["train"]["log_interval"] = log_interval
config["train"]["freeze_EN_bert"] = freeze_EN_bert config["train"]["freeze_EN_bert"] = freeze_EN_bert
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"]["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.") logger.warning(f"Step 1: {model_path} already exists, so copy it to backup.")
@@ -146,7 +150,7 @@ def resample(model_name, normalize, trim, num_processes):
return True, "Step 2, Success: 音声ファイルの前処理が完了しました" return True, "Step 2, Success: 音声ファイルの前処理が完了しました"
def preprocess_text(model_name, use_jp_extra): def preprocess_text(model_name, use_jp_extra, val_per_lang):
logger.info("Step 3: start preprocessing text...") logger.info("Step 3: start preprocessing text...")
dataset_path, lbl_path, train_path, val_path, config_path = get_path(model_name) dataset_path, lbl_path, train_path, val_path, config_path = get_path(model_name)
try: try:
@@ -171,6 +175,8 @@ def preprocess_text(model_name, use_jp_extra):
train_path, train_path,
"--val-path", "--val-path",
val_path, val_path,
"--val-per-lang",
val_per_lang,
] ]
if use_jp_extra: if use_jp_extra:
cmd.append("--use_jp_extra") cmd.append("--use_jp_extra")
@@ -257,6 +263,8 @@ def preprocess_all(
freeze_ZH_bert, freeze_ZH_bert,
freeze_style, freeze_style,
use_jp_extra, use_jp_extra,
val_per_lang,
log_interval,
): ):
if model_name == "": if model_name == "":
return False, "Error: モデル名を入力してください" return False, "Error: モデル名を入力してください"
@@ -270,13 +278,14 @@ def preprocess_all(
freeze_ZH_bert, freeze_ZH_bert,
freeze_style, freeze_style,
use_jp_extra, use_jp_extra,
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, normalize, trim, num_processes)
if not success: if not success:
return False, message return False, message
success, message = preprocess_text(model_name, use_jp_extra) success, message = preprocess_text(model_name, use_jp_extra, 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) # bert_genは重いのでプロセス数いじらない
@@ -372,7 +381,7 @@ initial_md = f"""
- 途中から学習を再開する場合は、モデル名を入力してから「学習を開始する」を押せばよいです。 - 途中から学習を再開する場合は、モデル名を入力してから「学習を開始する」を押せばよいです。
注意: 音声合成で使うには、スタイルベクトルファイル`style_vectors.npy`を作る必要があります。これは、`Style.bat`を実行してそこで作成してください。 注意: 標準スタイル以外のスタイルを音声合成で使うには、スタイルベクトルファイル`style_vectors.npy`を作る必要があります。これは、`Style.bat`を実行してそこで作成してください。
動作は軽いはずなので、学習中でも実行でき、何度でも繰り返して試せます。 動作は軽いはずなので、学習中でも実行でき、何度でも繰り返して試せます。
## JP-Extra版について ## JP-Extra版について
@@ -464,6 +473,19 @@ if __name__ == "__main__":
maximum=cpu_count(), maximum=cpu_count(),
step=1, step=1,
) )
val_per_lang = gr.Textbox(
label="検証データ数",
info="学習には使われず、tensorboard等で確認するためのもの",
value="0",
)
log_interval = gr.Slider(
label="Tensorboardのログ出力間隔",
info="Tensorboardで詳しく見たい人は小さめにしてください",
value=200,
minimum=10,
maximum=1000,
step=10,
)
gr.Markdown("学習時に特定の部分を凍結させるかどうか") gr.Markdown("学習時に特定の部分を凍結させるかどうか")
freeze_EN_bert = gr.Checkbox( freeze_EN_bert = gr.Checkbox(
label="英語bert部分を凍結", label="英語bert部分を凍結",
@@ -516,6 +538,13 @@ if __name__ == "__main__":
maximum=10000, maximum=10000,
step=100, step=100,
) )
log_interval_manual = gr.Slider(
label="Tensorboardのログ出力間隔",
value=200,
minimum=10,
maximum=1000,
step=10,
)
freeze_EN_bert_manual = gr.Checkbox( freeze_EN_bert_manual = gr.Checkbox(
label="英語bert部分を凍結", label="英語bert部分を凍結",
value=False, value=False,
@@ -559,6 +588,10 @@ if __name__ == "__main__":
with gr.Row(variant="panel"): with gr.Row(variant="panel"):
with gr.Column(): with gr.Column():
gr.Markdown(value="#### Step 3: 書き起こしファイルの前処理") gr.Markdown(value="#### Step 3: 書き起こしファイルの前処理")
val_per_lang_manual = gr.Textbox(
label="検証データ数",
value="0",
)
with gr.Column(): with gr.Column():
preprocess_text_btn = gr.Button(value="実行", variant="primary") preprocess_text_btn = gr.Button(value="実行", variant="primary")
info_preprocess_text = gr.Textbox(label="状況") info_preprocess_text = gr.Textbox(label="状況")
@@ -599,6 +632,9 @@ if __name__ == "__main__":
) )
train_btn = gr.Button(value="学習を開始する", variant="primary") train_btn = gr.Button(value="学習を開始する", variant="primary")
tensorboard_btn = gr.Button(value="Tensorboardを開く") tensorboard_btn = gr.Button(value="Tensorboardを開く")
gr.Markdown(
"進捗はターミナルで確認してください。随時結果は指定したステップごとに保存されており、また学習を途中から再開もできます。学習を終了するには単にターミナルを終了してください。"
)
info_train = gr.Textbox(label="状況") info_train = gr.Textbox(label="状況")
preprocess_button.click( preprocess_button.click(
@@ -616,6 +652,8 @@ if __name__ == "__main__":
freeze_ZH_bert, freeze_ZH_bert,
freeze_style, freeze_style,
use_jp_extra, use_jp_extra,
val_per_lang,
log_interval,
], ],
outputs=[info_all], outputs=[info_all],
) )
@@ -633,6 +671,7 @@ if __name__ == "__main__":
freeze_ZH_bert_manual, freeze_ZH_bert_manual,
freeze_style_manual, freeze_style_manual,
use_jp_extra_manual, use_jp_extra_manual,
log_interval_manual,
], ],
outputs=[info_init], outputs=[info_init],
) )
@@ -648,7 +687,7 @@ if __name__ == "__main__":
) )
preprocess_text_btn.click( preprocess_text_btn.click(
second_elem_of(preprocess_text), second_elem_of(preprocess_text),
inputs=[model_name, use_jp_extra_manual], inputs=[model_name, use_jp_extra_manual, val_per_lang_manual],
outputs=[info_preprocess_text], outputs=[info_preprocess_text],
) )
bert_gen_btn.click( bert_gen_btn.click(