This commit is contained in:
litagin02
2024-01-03 10:36:00 +09:00
parent 704491d796
commit 30067335e1
5 changed files with 40 additions and 18 deletions

View File

@@ -260,20 +260,9 @@ with open(os.path.join("configs", "paths.yml"), "r", encoding="utf-8") as f:
# - dataset_root: the root directory of the dataset, default to "Data"
# - assets_root: the root directory of the assets, default to "model_assets"
parser = argparse.ArgumentParser()
# 为避免与以前的config.json起冲突将其更名如下
parser.add_argument(
"-y",
"--yml_config",
type=str,
default="config.yml",
help="Path to setting yaml file.",
)
args = parser.parse_args()
try:
config = Config(args.yml_config, path_config)
config = Config("config.yml", path_config)
except TypeError:
logger.warning("Old config.yml found. Replace it with default_config.yml.")
shutil.copy(src="default_config.yml", dst="config.yml")

View File

@@ -3,6 +3,8 @@ import csv
import warnings
from pathlib import Path
import matplotlib.pyplot as plt
import pandas as pd
import torch
from tqdm import tqdm
@@ -71,9 +73,6 @@ for model_file in tqdm(safetensors_files):
score = predictor(torch.from_numpy(audio).unsqueeze(0), sr).item()
scores.append(score)
logger.info(f"score: {score}")
mean = sum(scores) / len(scores)
logger.success(f"mean: {mean}")
scores.append(mean)
# `test_e10_s1000.safetensors`` -> 1000を取り出す
step = int(model_file.stem.split("_")[2][1:])
results.append((model_file.name, step, scores))
@@ -90,3 +89,36 @@ with open(args.output, "w", encoding="utf-8", newline="") as f:
writer.writerow(["model_path"] + ["step"] + test_texts + ["mean"])
for model_file, step, scores in results:
writer.writerow([model_file] + [step] + scores)
# step countと各MOSの値を格納するリストを初期化
steps = []
mos_values = []
# resultsからデータを抽出
for _, step, scores in results:
steps.append(step)
mos_values.append(scores) # scores は MOS1, MOS2, MOS3,... のリスト
# DataFrame形式に変換
df = pd.DataFrame(mos_values, index=steps)
# ステップ数でソート
df = df.sort_index()
# 各MOSと平均値についての折れ線グラフを描画
for col in df.columns:
plt.plot(df.index, df[col], label=f"MOS{col + 1}")
# 平均値の計算と描画
df["mean"] = df.mean(axis=1)
plt.plot(df.index, df["mean"], label="Mean", color="black", linewidth=2)
# グラフのタイトルと軸ラベルを設定
plt.title("TTS Model Naturalness MOS")
plt.xlabel("Step Count")
plt.ylabel("MOS")
# 凡例を表示
plt.legend()
# グラフを表示
plt.show()

View File

@@ -70,6 +70,7 @@ if __name__ == "__main__":
language = Languages.ZH
else:
raise ValueError(f"{language} is not supported.")
logger.debug(f"Initial prompt: {initial_prompt}")
with open(output_file, "w", encoding="utf-8") as f:
for wav_file in tqdm(wav_files, file=SAFE_STDOUT):
file_name = os.path.basename(wav_file)

View File

@@ -56,7 +56,7 @@ def do_transcribe(model_name, whisper_model, compute_type, language, initial_pro
"--language",
language,
"--initial_prompt",
'"initial_prompt"',
f'"{initial_prompt}"',
]
)
return "音声の文字起こしが完了しました。"

View File

@@ -302,9 +302,9 @@ if __name__ == "__main__":
label="エポック数",
info="100もあれば十分そうだけどもっと回すと質が上がるかもしれない",
value=100,
minimum=1,
minimum=10,
maximum=1000,
step=1,
step=10,
)
save_every_steps = gr.Slider(
label="何ステップごとに結果を保存するか",