Feat: multiprocessing of slicing for much faster slicing
This commit is contained in:
110
slice.py
110
slice.py
@@ -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."
|
||||||
|
|||||||
@@ -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],
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user