Refactor: Use PathConfig and pathlib instead of paths.yml loading
This commit is contained in:
@@ -1,27 +1,23 @@
|
||||
import json
|
||||
import os
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
|
||||
import gradio as gr
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
import yaml
|
||||
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 config import get_path_config
|
||||
from style_bert_vits2.constants import DEFAULT_STYLE, GRADIO_THEME
|
||||
from style_bert_vits2.logging import logger
|
||||
|
||||
|
||||
# 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(path_config["dataset_root"])
|
||||
# assets_root = path_config["assets_root"]
|
||||
path_config = get_path_config()
|
||||
dataset_root = path_config.dataset_root
|
||||
assets_root = path_config.assets_root
|
||||
|
||||
MAX_CLUSTER_NUM = 10
|
||||
MAX_AUDIO_NUM = 10
|
||||
@@ -39,11 +35,7 @@ centroids = []
|
||||
|
||||
def load(model_name: str, reduction_method: str):
|
||||
global wav_files, x, x_reduced, mean
|
||||
# wavs_dir = os.path.join(dataset_root, model_name, "wavs")
|
||||
wavs_dir = dataset_root / model_name / "wavs"
|
||||
# style_vector_files = [
|
||||
# os.path.join(wavs_dir, f) for f in os.listdir(wavs_dir) if f.endswith(".npy")
|
||||
# ]
|
||||
style_vector_files = [f for f in wavs_dir.rglob("*.npy") if f.is_file()]
|
||||
# foo.wav.npy -> foo.wav
|
||||
wav_files = [f.with_suffix("") for f in style_vector_files]
|
||||
@@ -196,21 +188,21 @@ def do_clustering_gradio(n_clusters=4, method="KMeans"):
|
||||
] * MAX_AUDIO_NUM
|
||||
|
||||
|
||||
def save_style_vectors_from_clustering(model_name, style_names_str: str):
|
||||
def save_style_vectors_from_clustering(model_name: str, style_names_str: str):
|
||||
"""centerとcentroidsを保存する"""
|
||||
result_dir = os.path.join(config.assets_root, model_name)
|
||||
os.makedirs(result_dir, exist_ok=True)
|
||||
result_dir = assets_root / model_name
|
||||
result_dir.mkdir(parents=True, exist_ok=True)
|
||||
style_vectors = np.stack([mean] + centroids)
|
||||
style_vector_path = os.path.join(result_dir, "style_vectors.npy")
|
||||
if os.path.exists(style_vector_path):
|
||||
style_vector_path = result_dir / "style_vectors.npy"
|
||||
if style_vector_path.exists():
|
||||
logger.info(f"Backup {style_vector_path} to {style_vector_path}.bak")
|
||||
shutil.copy(style_vector_path, f"{style_vector_path}.bak")
|
||||
np.save(style_vector_path, style_vectors)
|
||||
logger.success(f"Saved style vectors to {style_vector_path}")
|
||||
|
||||
# config.jsonの更新
|
||||
config_path = os.path.join(result_dir, "config.json")
|
||||
if not os.path.exists(config_path):
|
||||
config_path = result_dir / "config.json"
|
||||
if not config_path.exists():
|
||||
return f"{config_path}が存在しません。"
|
||||
style_names = [name.strip() for name in style_names_str.split(",")]
|
||||
style_name_list = [DEFAULT_STYLE] + style_names
|
||||
@@ -233,7 +225,7 @@ def save_style_vectors_from_clustering(model_name, style_names_str: str):
|
||||
|
||||
|
||||
def save_style_vectors_from_files(
|
||||
model_name, audio_files_str: str, style_names_str: str
|
||||
model_name: str, audio_files_str: str, style_names_str: str
|
||||
):
|
||||
"""音声ファイルからスタイルベクトルを作成して保存する"""
|
||||
global mean
|
||||
@@ -241,8 +233,8 @@ def save_style_vectors_from_files(
|
||||
return "Error: スタイルベクトルを読み込んでください。"
|
||||
mean = np.mean(x, axis=0)
|
||||
|
||||
result_dir = os.path.join(config.assets_root, model_name)
|
||||
os.makedirs(result_dir, exist_ok=True)
|
||||
result_dir = assets_root / model_name
|
||||
result_dir.mkdir(parents=True, exist_ok=True)
|
||||
audio_files = [name.strip() for name in audio_files_str.split(",")]
|
||||
style_names = [name.strip() for name in style_names_str.split(",")]
|
||||
if len(audio_files) != len(style_names):
|
||||
@@ -252,23 +244,23 @@ def save_style_vectors_from_files(
|
||||
return "スタイル名が重複しています。"
|
||||
style_vectors = [mean]
|
||||
|
||||
wavs_dir = os.path.join(dataset_root, model_name, "wavs")
|
||||
wavs_dir = dataset_root / model_name / "wavs"
|
||||
for audio_file in audio_files:
|
||||
path = os.path.join(wavs_dir, audio_file)
|
||||
if not os.path.exists(path):
|
||||
path = wavs_dir / audio_file
|
||||
if not path.exists():
|
||||
return f"{path}が存在しません。"
|
||||
style_vectors.append(np.load(f"{path}.npy"))
|
||||
style_vectors = np.stack(style_vectors)
|
||||
assert len(style_name_list) == len(style_vectors)
|
||||
style_vector_path = os.path.join(result_dir, "style_vectors.npy")
|
||||
if os.path.exists(style_vector_path):
|
||||
style_vector_path = result_dir / "style_vectors.npy"
|
||||
if style_vector_path.exists():
|
||||
logger.info(f"Backup {style_vector_path} to {style_vector_path}.bak")
|
||||
shutil.copy(style_vector_path, f"{style_vector_path}.bak")
|
||||
np.save(style_vector_path, style_vectors)
|
||||
|
||||
# config.jsonの更新
|
||||
config_path = os.path.join(result_dir, "config.json")
|
||||
if not os.path.exists(config_path):
|
||||
config_path = result_dir / "config.json"
|
||||
if not config_path.exists():
|
||||
return f"{config_path}が存在しません。"
|
||||
logger.info(f"Backup {config_path} to {config_path}.bak")
|
||||
shutil.copy(config_path, f"{config_path}.bak")
|
||||
|
||||
Reference in New Issue
Block a user