Change args from output_dir to model_name

This commit is contained in:
litagin02
2024-02-19 19:16:41 +09:00
parent 4c8ea82959
commit 79f20fafbb
2 changed files with 21 additions and 11 deletions

View File

@@ -5,6 +5,7 @@ from pathlib import Path
import soundfile as sf import soundfile as sf
import torch import torch
import yaml
from tqdm import tqdm from tqdm import tqdm
from common.log import logger from common.log import logger
@@ -76,6 +77,7 @@ def split_wav(
os.makedirs(target_dir, exist_ok=True) os.makedirs(target_dir, exist_ok=True)
total_time_ms = 0 total_time_ms = 0
count = 0
# タイムスタンプに従って分割し、ファイルに保存 # タイムスタンプに従って分割し、ファイルに保存
for i, ts in enumerate(speech_timestamps): 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) sf.write(os.path.join(target_dir, f"{file_name}-{i}.wav"), segment, sr)
total_time_ms += end_ms - start_ms total_time_ms += end_ms - start_ms
count += 1
return total_time_ms / 1000 return total_time_ms / 1000, count
if __name__ == "__main__": if __name__ == "__main__":
@@ -108,11 +111,11 @@ if __name__ == "__main__":
help="Directory of input wav files", help="Directory of input wav files",
) )
parser.add_argument( parser.add_argument(
"--output_dir", "--model_name",
"-o", "-m",
type=str, type=str,
default="raw", required=True,
help="Directory of output wav files", help="The result will be in Data/{model_name}/raw/ (if Data is dataset_root in configs/paths.yml)",
) )
parser.add_argument( parser.add_argument(
"--min_silence_dur_ms", "--min_silence_dur_ms",
@@ -123,8 +126,12 @@ if __name__ == "__main__":
) )
args = parser.parse_args() 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 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 min_sec = args.min_sec
max_sec = args.max_sec max_sec = args.max_sec
min_silence_dur_ms = args.min_silence_dur_ms min_silence_dur_ms = args.min_silence_dur_ms
@@ -137,8 +144,9 @@ if __name__ == "__main__":
shutil.rmtree(output_dir) shutil.rmtree(output_dir)
total_sec = 0 total_sec = 0
total_count = 0
for wav_file in tqdm(wav_files, file=SAFE_STDOUT): for wav_file in tqdm(wav_files, file=SAFE_STDOUT):
time_sec = split_wav( time_sec, count = split_wav(
audio_file=str(wav_file), audio_file=str(wav_file),
target_dir=output_dir, target_dir=output_dir,
min_sec=min_sec, min_sec=min_sec,
@@ -146,5 +154,8 @@ if __name__ == "__main__":
min_silence_dur_ms=min_silence_dur_ms, min_silence_dur_ms=min_silence_dur_ms,
) )
total_sec += time_sec 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."
)

View File

@@ -25,11 +25,10 @@ def do_slice(
if model_name == "": if model_name == "":
return "Error: モデル名を入力してください。" return "Error: モデル名を入力してください。"
logger.info("Start slicing...") logger.info("Start slicing...")
output_dir = os.path.join(dataset_root, model_name, "raw")
cmd = [ cmd = [
"slice.py", "slice.py",
"--output_dir", "--model_name",
output_dir, model_name,
"--min_sec", "--min_sec",
str(min_sec), str(min_sec),
"--max_sec", "--max_sec",