Feat: multiprocessing of slicing for much faster slicing

This commit is contained in:
litagin02
2024-03-13 18:56:51 +09:00
parent dd407d882d
commit 4f60a3d5d5
2 changed files with 111 additions and 21 deletions

110
slice.py
View File

@@ -1,6 +1,10 @@
import argparse import argparse
import shutil import shutil
import sys
from pathlib import Path from pathlib import Path
from queue import Queue
from threading import Thread
from typing import Any, Optional
import soundfile as sf import soundfile as sf
import torch import torch
@@ -10,20 +14,12 @@ from tqdm import tqdm
from style_bert_vits2.logging import logger from style_bert_vits2.logging import logger
from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT
# TODO: 並列処理による高速化 # TODO: 並列処理による高速化
vad_model, utils = torch.hub.load(
repo_or_dir="litagin02/silero-vad",
model="silero_vad",
onnx=True,
trust_repo=True,
)
(get_speech_timestamps, _, read_audio, *_) = utils
def get_stamps( def get_stamps(
vad_model: Any,
utils: Any,
audio_file: Path, audio_file: Path,
min_silence_dur_ms: int = 700, min_silence_dur_ms: int = 700,
min_sec: float = 2, min_sec: float = 2,
@@ -42,6 +38,7 @@ def get_stamps(
この秒数より大きい発話は無視する。 この秒数より大きい発話は無視する。
""" """
(get_speech_timestamps, _, read_audio, *_) = utils
sampling_rate = 16000 # 16kHzか8kHzのみ対応 sampling_rate = 16000 # 16kHzか8kHzのみ対応
min_ms = int(min_sec * 1000) min_ms = int(min_sec * 1000)
@@ -60,6 +57,8 @@ def get_stamps(
def split_wav( def split_wav(
vad_model: Any,
utils: Any,
audio_file: Path, audio_file: Path,
target_dir: Path, target_dir: Path,
min_sec: float = 2, min_sec: float = 2,
@@ -69,7 +68,9 @@ def split_wav(
) -> tuple[float, int]: ) -> tuple[float, int]:
margin: int = 200 # ミリ秒単位で、音声の前後に余裕を持たせる margin: int = 200 # ミリ秒単位で、音声の前後に余裕を持たせる
speech_timestamps = get_stamps( speech_timestamps = get_stamps(
audio_file, vad_model=vad_model,
utils=utils,
audio_file=audio_file,
min_silence_dur_ms=min_silence_dur_ms, min_silence_dur_ms=min_silence_dur_ms,
min_sec=min_sec, min_sec=min_sec,
max_sec=max_sec, max_sec=max_sec,
@@ -139,6 +140,12 @@ if __name__ == "__main__":
action="store_true", action="store_true",
help="Make the filename end with -start_ms-end_ms when saving wav.", help="Make the filename end with -start_ms-end_ms when saving wav.",
) )
parser.add_argument(
"--num_processes",
type=int,
default=3,
help="Number of processes to use. Default 3 seems to be the best.",
)
args = parser.parse_args() args = parser.parse_args()
with open(Path("configs/paths.yml"), "r", encoding="utf-8") as f: with open(Path("configs/paths.yml"), "r", encoding="utf-8") as f:
@@ -152,6 +159,7 @@ if __name__ == "__main__":
max_sec: float = args.max_sec max_sec: float = args.max_sec
min_silence_dur_ms: int = args.min_silence_dur_ms min_silence_dur_ms: int = args.min_silence_dur_ms
time_suffix: bool = args.time_suffix time_suffix: bool = args.time_suffix
num_processes: int = args.num_processes
wav_files = Path(input_dir).glob("**/*.wav") wav_files = Path(input_dir).glob("**/*.wav")
wav_files = list(wav_files) wav_files = list(wav_files)
@@ -160,19 +168,89 @@ if __name__ == "__main__":
logger.warning(f"Output directory {output_dir} already exists, deleting...") logger.warning(f"Output directory {output_dir} already exists, deleting...")
shutil.rmtree(output_dir) shutil.rmtree(output_dir)
total_sec = 0 # Silero VADのモデルは、同じインスタンスで並列処理するとおかしくなるらしい
total_count = 0 # ワーカーごとにモデルをロードするようにするため、Queueを使って処理する
for wav_file in tqdm(wav_files, file=SAFE_STDOUT): def process_queue(
q: Queue[Optional[Path]],
result_queue: Queue[tuple[float, int]],
error_queue: Queue[tuple[Path, Exception]],
):
# logger.debug("Worker started.")
vad_model, utils = torch.hub.load(
repo_or_dir="litagin02/silero-vad",
model="silero_vad",
onnx=True,
trust_repo=True,
)
while True:
file = q.get()
if file is None: # 終了シグナルを確認
q.task_done()
break
try:
time_sec, count = split_wav( time_sec, count = split_wav(
audio_file=wav_file, vad_model=vad_model,
utils=utils,
audio_file=file,
target_dir=output_dir, target_dir=output_dir,
min_sec=min_sec, min_sec=min_sec,
max_sec=max_sec, max_sec=max_sec,
min_silence_dur_ms=min_silence_dur_ms, min_silence_dur_ms=min_silence_dur_ms,
time_suffix=time_suffix, time_suffix=time_suffix,
) )
total_sec += time_sec result_queue.put((time_sec, count))
except Exception as e:
logger.error(f"Error processing {file}: {e}")
error_queue.put((file, e))
result_queue.put((0, 0))
finally:
q.task_done()
q: Queue[Optional[Path]] = Queue()
result_queue: Queue[tuple[float, int]] = Queue()
error_queue: Queue[tuple[Path, Exception]] = Queue()
# ファイル数が少ない場合は、ワーカー数をファイル数に合わせる
num_processes = min(num_processes, len(wav_files))
threads = [
Thread(target=process_queue, args=(q, result_queue, error_queue))
for _ in range(num_processes)
]
for t in threads:
t.start()
pbar = tqdm(total=len(wav_files), file=SAFE_STDOUT)
for file in wav_files:
q.put(file)
# result_queueを監視し、要素が追加されるごとに結果を加算しプログレスバーを更新
total_sec = 0
total_count = 0
for _ in range(len(wav_files)):
time, count = result_queue.get()
total_sec += time
total_count += count total_count += count
pbar.update(1)
# 全ての処理が終わるまで待つ
q.join()
# 終了シグナル None を送る
for _ in range(num_processes):
q.put(None)
for t in threads:
t.join()
pbar.close()
if not error_queue.empty():
error_str = "Error slicing some files:"
while not error_queue.empty():
file, e = error_queue.get()
error_str += f"\n{file}: {e}"
raise RuntimeError(error_str)
logger.info( logger.info(
f"Slice done! Total time: {total_sec / 60:.2f} min, {total_count} files." f"Slice done! Total time: {total_sec / 60:.2f} min, {total_count} files."

View File

@@ -11,6 +11,7 @@ def do_slice(
min_silence_dur_ms: int, min_silence_dur_ms: int,
time_suffix: bool, time_suffix: bool,
input_dir: str, input_dir: str,
num_processes: int = 3,
): ):
if model_name == "": if model_name == "":
return "Error: モデル名を入力してください。" return "Error: モデル名を入力してください。"
@@ -25,6 +26,8 @@ def do_slice(
str(max_sec), str(max_sec),
"--min_silence_dur_ms", "--min_silence_dur_ms",
str(min_silence_dur_ms), str(min_silence_dur_ms),
"--num_processes",
str(num_processes),
] ]
if time_suffix: if time_suffix:
cmd.append("--time_suffix") cmd.append("--time_suffix")
@@ -137,6 +140,14 @@ def create_dataset_app() -> gr.Blocks:
value=False, value=False,
label="WAVファイル名の末尾に元ファイルの時間範囲を付与する", label="WAVファイル名の末尾に元ファイルの時間範囲を付与する",
) )
num_processes = gr.Slider(
minimum=1,
maximum=10,
value=3,
step=1,
label="並列処理数(速度向上のため)",
info="3で十分高速、多くしてもCPU負荷が増すだけでそこまで速度は変わらない",
)
slice_button = gr.Button("スライスを実行") slice_button = gr.Button("スライスを実行")
result1 = gr.Textbox(label="結果") result1 = gr.Textbox(label="結果")
with gr.Row(): with gr.Row():
@@ -186,6 +197,7 @@ def create_dataset_app() -> gr.Blocks:
min_silence_dur_ms, min_silence_dur_ms,
time_suffix, time_suffix,
input_dir, input_dir,
num_processes,
], ],
outputs=[result1], outputs=[result1],
) )