From 79f20fafbb94d54956c5bedc78bc5365ece803c6 Mon Sep 17 00:00:00 2001 From: litagin02 Date: Mon, 19 Feb 2024 19:16:41 +0900 Subject: [PATCH] Change args from output_dir to model_name --- slice.py | 27 +++++++++++++++++++-------- webui_dataset.py | 5 ++--- 2 files changed, 21 insertions(+), 11 deletions(-) diff --git a/slice.py b/slice.py index 3b6e2d7..b20fe9c 100644 --- a/slice.py +++ b/slice.py @@ -5,6 +5,7 @@ from pathlib import Path import soundfile as sf import torch +import yaml from tqdm import tqdm from common.log import logger @@ -76,6 +77,7 @@ def split_wav( os.makedirs(target_dir, exist_ok=True) total_time_ms = 0 + count = 0 # タイムスタンプに従って分割し、ファイルに保存 for i, ts in enumerate(speech_timestamps): @@ -88,8 +90,9 @@ def split_wav( sf.write(os.path.join(target_dir, f"{file_name}-{i}.wav"), segment, sr) total_time_ms += end_ms - start_ms + count += 1 - return total_time_ms / 1000 + return total_time_ms / 1000, count if __name__ == "__main__": @@ -108,11 +111,11 @@ if __name__ == "__main__": help="Directory of input wav files", ) parser.add_argument( - "--output_dir", - "-o", + "--model_name", + "-m", type=str, - default="raw", - help="Directory of output wav files", + required=True, + help="The result will be in Data/{model_name}/raw/ (if Data is dataset_root in configs/paths.yml)", ) parser.add_argument( "--min_silence_dur_ms", @@ -123,8 +126,12 @@ if __name__ == "__main__": ) args = parser.parse_args() + 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"] + input_dir = args.input_dir - output_dir = args.output_dir + output_dir = os.path.join(dataset_root, args.model_name, "raw") min_sec = args.min_sec max_sec = args.max_sec min_silence_dur_ms = args.min_silence_dur_ms @@ -137,8 +144,9 @@ if __name__ == "__main__": shutil.rmtree(output_dir) total_sec = 0 + total_count = 0 for wav_file in tqdm(wav_files, file=SAFE_STDOUT): - time_sec = split_wav( + time_sec, count = split_wav( audio_file=str(wav_file), target_dir=output_dir, min_sec=min_sec, @@ -146,5 +154,8 @@ if __name__ == "__main__": min_silence_dur_ms=min_silence_dur_ms, ) total_sec += time_sec + total_count += count - logger.info(f"Slice done! Total time: {total_sec / 60:.2f} min.") + logger.info( + f"Slice done! Total time: {total_sec / 60:.2f} min, {total_count} files." + ) diff --git a/webui_dataset.py b/webui_dataset.py index 2885e27..02fb15d 100644 --- a/webui_dataset.py +++ b/webui_dataset.py @@ -25,11 +25,10 @@ def do_slice( if model_name == "": return "Error: モデル名を入力してください。" logger.info("Start slicing...") - output_dir = os.path.join(dataset_root, model_name, "raw") cmd = [ "slice.py", - "--output_dir", - output_dir, + "--model_name", + model_name, "--min_sec", str(min_sec), "--max_sec",