Fix colab, improve yomi_error doc, CLI yomi_error option
This commit is contained in:
13
colab.ipynb
13
colab.ipynb
@@ -4,7 +4,7 @@
|
|||||||
"cell_type": "markdown",
|
"cell_type": "markdown",
|
||||||
"metadata": {},
|
"metadata": {},
|
||||||
"source": [
|
"source": [
|
||||||
"# Style-Bert-VITS2 (ver 2.3) のGoogle Colabでの学習\n",
|
"# Style-Bert-VITS2 (ver 2.3.1) のGoogle Colabでの学習\n",
|
||||||
"\n",
|
"\n",
|
||||||
"Google Colab上でStyle-Bert-VITS2の学習を行うことができます。\n",
|
"Google Colab上でStyle-Bert-VITS2の学習を行うことができます。\n",
|
||||||
"\n",
|
"\n",
|
||||||
@@ -118,8 +118,8 @@
|
|||||||
"# こういうふうに書き起こして欲しいという例文(句読点の入れ方・笑い方や固有名詞等)\n",
|
"# こういうふうに書き起こして欲しいという例文(句読点の入れ方・笑い方や固有名詞等)\n",
|
||||||
"initial_prompt = \"こんにちは。元気、ですかー?ふふっ、私は……ちゃんと元気だよ!\"\n",
|
"initial_prompt = \"こんにちは。元気、ですかー?ふふっ、私は……ちゃんと元気だよ!\"\n",
|
||||||
"\n",
|
"\n",
|
||||||
"!python slice.py -i {input_dir} -o {dataset_root}/{model_name}/raw\n",
|
"!python slice.py -i {input_dir} --model_name {model_name}\n",
|
||||||
"!python transcribe.py -i {dataset_root}/{model_name}/raw -o {dataset_root}/{model_name}/esd.list --speaker_name {model_name} --compute_type float16 --initial_prompt {initial_prompt}"
|
"!python transcribe.py --model_name {model_name} --compute_type float16 --initial_prompt {initial_prompt}"
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
@@ -229,7 +229,11 @@
|
|||||||
"normalize = False\n",
|
"normalize = False\n",
|
||||||
"\n",
|
"\n",
|
||||||
"# 音声ファイルの開始・終了にある無音区間を削除するかどうか\n",
|
"# 音声ファイルの開始・終了にある無音区間を削除するかどうか\n",
|
||||||
"trim = False"
|
"trim = False\n",
|
||||||
|
"\n",
|
||||||
|
"# 読みのエラーが出た場合にどうするか。\n",
|
||||||
|
"# \"raise\"ならテキスト前処理が終わったら中断、\"skip\"なら読めない行は学習に使わない、\"use\"なら無理やり使う\n",
|
||||||
|
"yomi_error = \"skip\""
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
@@ -269,6 +273,7 @@
|
|||||||
" 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",
|
||||||
|
" yomi_error=yomi_error\n",
|
||||||
")"
|
")"
|
||||||
]
|
]
|
||||||
},
|
},
|
||||||
|
|||||||
@@ -4,7 +4,7 @@ import enum
|
|||||||
# See https://huggingface.co/spaces/gradio/theme-gallery for more themes
|
# See https://huggingface.co/spaces/gradio/theme-gallery for more themes
|
||||||
GRADIO_THEME: str = "NoCrypt/miku"
|
GRADIO_THEME: str = "NoCrypt/miku"
|
||||||
|
|
||||||
LATEST_VERSION: str = "2.3"
|
LATEST_VERSION: str = "2.3.1"
|
||||||
|
|
||||||
USER_DICT_DIR = "dict_data"
|
USER_DICT_DIR = "dict_data"
|
||||||
|
|
||||||
|
|||||||
@@ -68,5 +68,5 @@
|
|||||||
"use_spectral_norm": false,
|
"use_spectral_norm": false,
|
||||||
"gin_channels": 256
|
"gin_channels": 256
|
||||||
},
|
},
|
||||||
"version": "2.3"
|
"version": "2.3.1"
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -75,5 +75,5 @@
|
|||||||
"initial_channel": 64
|
"initial_channel": 64
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
"version": "2.3-JP-Extra"
|
"version": "2.3.1-JP-Extra"
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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] [--freeze_decoder]
|
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] [--yomi_error <yomi_error>]
|
||||||
```
|
```
|
||||||
|
|
||||||
Required:
|
Required:
|
||||||
@@ -76,6 +76,7 @@ Optional:
|
|||||||
- `--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).
|
||||||
|
- `--yomi_error`: How to handle yomi errors (default: `raise`: raise an error after preprocessing all texts, `skip`: skip the texts with errors, `use`: use the texts with errors by ignoring unknown characters).
|
||||||
|
|
||||||
## 3. Train
|
## 3. Train
|
||||||
|
|
||||||
|
|||||||
@@ -74,6 +74,9 @@ if __name__ == "__main__":
|
|||||||
help="Log interval",
|
help="Log interval",
|
||||||
default=200,
|
default=200,
|
||||||
)
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--yomi_error", type=str, help="Yomi error. raise, skip, use", default="raise"
|
||||||
|
)
|
||||||
|
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
|
|
||||||
@@ -93,4 +96,5 @@ if __name__ == "__main__":
|
|||||||
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,
|
||||||
|
yomi_error=args.yomi_error,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -105,7 +105,8 @@ def merge_style(model_name_a, model_name_b, weight, output_name, style_triple_li
|
|||||||
def lerp_tensors(t, v0, v1):
|
def lerp_tensors(t, v0, v1):
|
||||||
return v0 * (1 - t) + v1 * t
|
return v0 * (1 - t) + v1 * t
|
||||||
|
|
||||||
def slerp_tensors(t, v0, v1, dot_thres = 0.998):
|
|
||||||
|
def slerp_tensors(t, v0, v1, dot_thres=0.998):
|
||||||
device = v0.device
|
device = v0.device
|
||||||
v0c = v0.cpu().numpy()
|
v0c = v0.cpu().numpy()
|
||||||
v1c = v1.cpu().numpy()
|
v1c = v1.cpu().numpy()
|
||||||
@@ -119,7 +120,10 @@ def slerp_tensors(t, v0, v1, dot_thres = 0.998):
|
|||||||
sin_th0 = np.sin(th0)
|
sin_th0 = np.sin(th0)
|
||||||
th_t = th0 * t
|
th_t = th0 * t
|
||||||
|
|
||||||
return torch.from_numpy(v0c * np.sin(th0 - th_t) / sin_th0 + v1c * np.sin(th_t) / sin_th0).to(device)
|
return torch.from_numpy(
|
||||||
|
v0c * np.sin(th0 - th_t) / sin_th0 + v1c * np.sin(th_t) / sin_th0
|
||||||
|
).to(device)
|
||||||
|
|
||||||
|
|
||||||
def merge_models(
|
def merge_models(
|
||||||
model_path_a,
|
model_path_a,
|
||||||
@@ -269,10 +273,18 @@ def load_styles_gr(model_name_a, model_name_b):
|
|||||||
config_b = json.load(f)
|
config_b = json.load(f)
|
||||||
styles_b = list(config_b["data"]["style2id"].keys())
|
styles_b = list(config_b["data"]["style2id"].keys())
|
||||||
|
|
||||||
return gr.Textbox(value=", ".join(styles_a)), gr.Textbox(value=", ".join(styles_b)), gr.TextArea(
|
return (
|
||||||
label="スタイルのマージリスト",
|
gr.Textbox(value=", ".join(styles_a)),
|
||||||
placeholder=f"{DEFAULT_STYLE}, {DEFAULT_STYLE},{DEFAULT_STYLE}\nAngry, Angry, Angry",
|
gr.Textbox(value=", ".join(styles_b)),
|
||||||
value='\n'.join(f"{sty_a}, {sty_b}, {sty_a if sty_a != sty_b else ''}{sty_b}" for sty_a in styles_a for sty_b in styles_b),
|
gr.TextArea(
|
||||||
|
label="スタイルのマージリスト",
|
||||||
|
placeholder=f"{DEFAULT_STYLE}, {DEFAULT_STYLE},{DEFAULT_STYLE}\nAngry, Angry, Angry",
|
||||||
|
value="\n".join(
|
||||||
|
f"{sty_a}, {sty_b}, {sty_a if sty_a != sty_b else ''}{sty_b}"
|
||||||
|
for sty_a in styles_a
|
||||||
|
for sty_b in styles_b
|
||||||
|
),
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -387,7 +399,7 @@ with gr.Blocks(theme=GRADIO_THEME) as app:
|
|||||||
step=0.1,
|
step=0.1,
|
||||||
)
|
)
|
||||||
use_slerp_instead_of_lerp = gr.Checkbox(
|
use_slerp_instead_of_lerp = gr.Checkbox(
|
||||||
label="lerpのかわりにslerpを使う",
|
label="線形補完のかわりに球面線形補完を使う",
|
||||||
value=False,
|
value=False,
|
||||||
)
|
)
|
||||||
with gr.Column(variant="panel"):
|
with gr.Column(variant="panel"):
|
||||||
|
|||||||
@@ -496,11 +496,11 @@ if __name__ == "__main__":
|
|||||||
value=False,
|
value=False,
|
||||||
)
|
)
|
||||||
yomi_error = gr.Radio(
|
yomi_error = gr.Radio(
|
||||||
label="読みエラーの扱い",
|
label="書き起こしが読めないファイルの扱い",
|
||||||
choices=[
|
choices=[
|
||||||
("エラー出たらテキスト前処理が終わった時点で中断", "raise"),
|
("エラー出たらテキスト前処理が終わった時点で中断", "raise"),
|
||||||
("エラーファイルは使わず続行", "skip"),
|
("読めないファイルは使わず続行", "skip"),
|
||||||
("読みを強引に埋めて使い続行", "use"),
|
("読めないファイルも無理やり読んで学習に使う", "use"),
|
||||||
],
|
],
|
||||||
value="raise",
|
value="raise",
|
||||||
)
|
)
|
||||||
@@ -647,11 +647,11 @@ if __name__ == "__main__":
|
|||||||
step=1,
|
step=1,
|
||||||
)
|
)
|
||||||
yomi_error_manual = gr.Radio(
|
yomi_error_manual = gr.Radio(
|
||||||
label="読みエラーの扱い",
|
label="書き起こしが読めないファイルの扱い",
|
||||||
choices=[
|
choices=[
|
||||||
("エラー出たらテキスト前処理が終わった時点で中断", "raise"),
|
("エラー出たらテキスト前処理が終わった時点で中断", "raise"),
|
||||||
("エラーファイルは使わず続行", "skip"),
|
("読めないファイルは使わず続行", "skip"),
|
||||||
("読みを強引に埋めて使い続行", "use"),
|
("読めないファイルも無理やり読んで学習に使う", "use"),
|
||||||
],
|
],
|
||||||
value="raise",
|
value="raise",
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user