This commit is contained in:
litagin02
2024-03-15 12:27:46 +09:00
parent 867855ace9
commit 3d8b60c03c
7 changed files with 243 additions and 199 deletions

View File

@@ -6,12 +6,12 @@ import gradio as gr
import matplotlib.pyplot as plt
import numpy as np
import yaml
from config import config
from scipy.spatial.distance import pdist, squareform
from sklearn.cluster import DBSCAN, AgglomerativeClustering, KMeans
from sklearn.manifold import TSNE
from umap import UMAP
from config import config
from style_bert_vits2.constants import DEFAULT_STYLE, GRADIO_THEME
from style_bert_vits2.logging import logger

View File

@@ -24,32 +24,31 @@ tensorboard_executed = False
# Get path settings
with open(os.path.join("configs", "paths.yml"), "r", encoding="utf-8") as f:
path_config: dict[str, str] = yaml.safe_load(f.read())
dataset_root = path_config["dataset_root"]
# assets_root = path_config["assets_root"]
dataset_root = Path(path_config["dataset_root"])
def get_path(model_name):
def get_path(model_name: str) -> tuple[Path, Path, Path, Path, Path]:
assert model_name != "", "モデル名は空にできません"
dataset_path = os.path.join(dataset_root, model_name)
lbl_path = os.path.join(dataset_path, "esd.list")
train_path = os.path.join(dataset_path, "train.list")
val_path = os.path.join(dataset_path, "val.list")
config_path = os.path.join(dataset_path, "config.json")
dataset_path = dataset_root / model_name
lbl_path = dataset_path / "esd.list"
train_path = dataset_path / "train.list"
val_path = dataset_path / "val.list"
config_path = dataset_path / "config.json"
return dataset_path, lbl_path, train_path, val_path, config_path
def initialize(
model_name,
batch_size,
epochs,
save_every_steps,
freeze_EN_bert,
freeze_JP_bert,
freeze_ZH_bert,
freeze_style,
freeze_decoder,
use_jp_extra,
log_interval,
model_name: str,
batch_size: int,
epochs: int,
save_every_steps: int,
freeze_EN_bert: bool,
freeze_JP_bert: bool,
freeze_ZH_bert: bool,
freeze_style: bool,
freeze_decoder: bool,
use_jp_extra: bool,
log_interval: int,
):
global logger_handler
dataset_path, _, train_path, val_path, config_path = get_path(model_name)
@@ -72,8 +71,8 @@ def initialize(
with open(default_config_path, "r", encoding="utf-8") as f:
config = json.load(f)
config["model_name"] = model_name
config["data"]["training_files"] = train_path
config["data"]["validation_files"] = val_path
config["data"]["training_files"] = str(train_path)
config["data"]["validation_files"] = str(val_path)
config["train"]["batch_size"] = batch_size
config["train"]["epochs"] = epochs
config["train"]["eval_interval"] = save_every_steps
@@ -90,18 +89,18 @@ def initialize(
# 今はデフォルトであるが、以前は非JP-Extra版になくバグの原因になるので念のため
config["data"]["use_jp_extra"] = use_jp_extra
model_path = os.path.join(dataset_path, "models")
if os.path.exists(model_path):
model_path = dataset_path / "models"
if model_path.exists():
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"),
dst=dataset_path / "models_backup",
dirs_exist_ok=True,
)
shutil.rmtree(model_path)
pretrained_dir = "pretrained" if not use_jp_extra else "pretrained_jp_extra"
pretrained_dir = Path("pretrained" if not use_jp_extra else "pretrained_jp_extra")
try:
shutil.copytree(
src=pretrained_dir,
@@ -113,30 +112,29 @@ def initialize(
with open(config_path, "w", encoding="utf-8") as f:
json.dump(config, f, indent=2, ensure_ascii=False)
if not os.path.exists("config.yml"):
if not Path("config.yml").exists():
shutil.copy(src="default_config.yml", dst="config.yml")
# yml_data = safe_load(open("config.yml", "r", encoding="utf-8"))
with open("config.yml", "r", encoding="utf-8") as f:
yml_data = yaml.safe_load(f)
yml_data["model_name"] = model_name
yml_data["dataset_path"] = dataset_path
yml_data["dataset_path"] = str(dataset_path)
with open("config.yml", "w", encoding="utf-8") as f:
yaml.dump(yml_data, f, allow_unicode=True)
logger.success("Step 1: initialization finished.")
return True, "Step 1, Success: 初期設定が完了しました"
def resample(model_name, normalize, trim, num_processes):
def resample(model_name: str, normalize: bool, trim: bool, num_processes: int):
logger.info("Step 2: start resampling...")
dataset_path, _, _, _, _ = get_path(model_name)
input_dir = os.path.join(dataset_path, "raw")
output_dir = os.path.join(dataset_path, "wavs")
input_dir = dataset_path / "raw"
output_dir = dataset_path / "wavs"
cmd = [
"resample.py",
"-i",
input_dir,
str(input_dir),
"-o",
output_dir,
str(output_dir),
"--num_processes",
str(num_processes),
"--sr",
@@ -157,42 +155,30 @@ def resample(model_name, normalize, trim, num_processes):
return True, "Step 2, Success: 音声ファイルの前処理が完了しました"
def preprocess_text(model_name, use_jp_extra, val_per_lang, yomi_error):
def preprocess_text(
model_name: str, use_jp_extra: bool, val_per_lang: int, yomi_error: str
):
logger.info("Step 3: start preprocessing text...")
dataset_path, lbl_path, train_path, val_path, config_path = get_path(model_name)
try:
lines = open(lbl_path, "r", encoding="utf-8").readlines()
except FileNotFoundError:
_, lbl_path, train_path, val_path, config_path = get_path(model_name)
if not lbl_path.exists():
logger.error(f"Step 3: {lbl_path} not found.")
return False, f"Step 3, Error: 書き起こしファイル {lbl_path} が見つかりません。"
new_lines = []
for line in lines:
if len(line.strip().split("|")) != 4:
logger.error(f"Step 3: {lbl_path} has invalid format at line:\n{line}")
return (
False,
f"Step 3, Error: 書き起こしファイル次の行の形式が不正です:\n{line}",
)
path, spk, language, text = line.strip().split("|")
# pathをファイル名だけ取り出して正しいパスに変更
path = Path(dataset_path) / "wavs" / Path(path).name
new_lines.append(f"{path}|{spk}|{language}|{text}\n")
with open(lbl_path, "w", encoding="utf-8") as f:
f.writelines(new_lines)
cmd = [
"preprocess_text.py",
"--config-path",
config_path,
str(config_path),
"--transcription-path",
lbl_path,
str(lbl_path),
"--train-path",
train_path,
str(train_path),
"--val-path",
val_path,
str(val_path),
"--val-per-lang",
str(val_per_lang),
"--yomi_error",
yomi_error,
"--correct_path", # 音声ファイルのパスを正しいパスに修正する
]
if use_jp_extra:
cmd.append("--use_jp_extra")
@@ -213,17 +199,11 @@ def preprocess_text(model_name, use_jp_extra, val_per_lang, yomi_error):
return True, "Step 3, Success: 書き起こしファイルの前処理が完了しました"
def bert_gen(model_name):
def bert_gen(model_name: str):
logger.info("Step 4: start bert_gen...")
_, _, _, _, config_path = get_path(model_name)
success, message = run_script_with_log(
[
"bert_gen.py",
"--config",
config_path,
# "--num_processes", # bert_genは重いのでプロセス数いじらない
# str(num_processes),
]
["bert_gen.py", "--config", str(config_path)]
)
if not success:
logger.error("Step 4: bert_gen failed.")
@@ -238,14 +218,14 @@ def bert_gen(model_name):
return True, "Step 4, Success: BERT特徴ファイルの生成が完了しました"
def style_gen(model_name, num_processes):
def style_gen(model_name: str, num_processes: int):
logger.info("Step 5: start style_gen...")
_, _, _, _, config_path = get_path(model_name)
success, message = run_script_with_log(
[
"style_gen.py",
"--config",
config_path,
str(config_path),
"--num_processes",
str(num_processes),
]
@@ -267,22 +247,22 @@ def style_gen(model_name, num_processes):
def preprocess_all(
model_name,
batch_size,
epochs,
save_every_steps,
num_processes,
normalize,
trim,
freeze_EN_bert,
freeze_JP_bert,
freeze_ZH_bert,
freeze_style,
freeze_decoder,
use_jp_extra,
val_per_lang,
log_interval,
yomi_error,
model_name: str,
batch_size: int,
epochs: int,
save_every_steps: int,
num_processes: int,
normalize: bool,
trim: bool,
freeze_EN_bert: bool,
freeze_JP_bert: bool,
freeze_ZH_bert: bool,
freeze_style: bool,
freeze_decoder: bool,
use_jp_extra: bool,
val_per_lang: int,
log_interval: int,
yomi_error: str,
):
if model_name == "":
return False, "Error: モデル名を入力してください"
@@ -333,18 +313,23 @@ def preprocess_all(
)
def train(model_name, skip_style=False, use_jp_extra=True, speedup=False):
def train(
model_name: str,
skip_style: bool = False,
use_jp_extra: bool = True,
speedup: bool = False,
):
dataset_path, _, _, _, config_path = get_path(model_name)
# 学習再開の場合念のためconfig.ymlの名前等を更新
# 学習再開の場合を考えて念のためconfig.ymlの名前等を更新
with open("config.yml", "r", encoding="utf-8") as f:
yml_data = yaml.safe_load(f)
yml_data["model_name"] = model_name
yml_data["dataset_path"] = dataset_path
yml_data["dataset_path"] = str(dataset_path)
with open("config.yml", "w", encoding="utf-8") as f:
yaml.dump(yml_data, f, allow_unicode=True)
train_py = "train_ms.py" if not use_jp_extra else "train_ms_jp_extra.py"
cmd = [train_py, "--config", config_path, "--model", dataset_path]
cmd = [train_py, "--config", str(config_path), "--model", str(dataset_path)]
if skip_style:
cmd.append("--skip_default_style")
if speedup:
@@ -360,7 +345,7 @@ def train(model_name, skip_style=False, use_jp_extra=True, speedup=False):
return True, "Success: 学習が完了しました"
def wait_for_tensorboard(port=6006, timeout=10):
def wait_for_tensorboard(port: int = 6006, timeout: float = 10):
start_time = time.time()
while True:
try:
@@ -375,7 +360,7 @@ def wait_for_tensorboard(port=6006, timeout=10):
time.sleep(0.1)
def run_tensorboard(model_name):
def run_tensorboard(model_name: str):
global tensorboard_executed
if not tensorboard_executed:
python = sys.executable