Feat: HF whisper for transcribing (faster than faster-whisper)

This commit is contained in:
litagin02
2024-03-14 15:10:08 +09:00
parent 4f60a3d5d5
commit b1972a3d3d
4 changed files with 189 additions and 40 deletions

View File

@@ -1,6 +1,5 @@
import argparse import argparse
import shutil import shutil
import sys
from pathlib import Path from pathlib import Path
from queue import Queue from queue import Queue
from threading import Thread from threading import Thread
@@ -14,8 +13,6 @@ 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: 並列処理による高速化
def get_stamps( def get_stamps(
vad_model: Any, vad_model: Any,

0
tes.py Normal file
View File

View File

@@ -2,10 +2,10 @@ import argparse
import os import os
import sys import sys
from pathlib import Path from pathlib import Path
from typing import Optional from typing import Any, Optional
import yaml import yaml
from faster_whisper import WhisperModel from torch.utils.data import Dataset
from tqdm import tqdm from tqdm import tqdm
from style_bert_vits2.constants import Languages from style_bert_vits2.constants import Languages
@@ -13,16 +13,97 @@ 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
def transcribe( # faster-whisperは並列処理しても速度が向上しないので、単一モデルでループ処理する
wav_path: Path, initial_prompt: Optional[str] = None, language: str = "ja" def transcribe_with_faster_whisper(
model: "WhisperModel",
audio_file: Path,
initial_prompt: Optional[str] = None,
language: str = "ja",
num_beams: int = 1,
): ):
segments, _ = model.transcribe( segments, _ = model.transcribe(
str(wav_path), beam_size=5, language=language, initial_prompt=initial_prompt str(audio_file),
beam_size=num_beams,
language=language,
initial_prompt=initial_prompt,
) )
texts = [segment.text for segment in segments] texts = [segment.text for segment in segments]
return "".join(texts) return "".join(texts)
# HF pipelineで進捗表示をするために必要なDatasetクラス
class StrListDataset(Dataset[str]):
def __init__(self, original_list: list[str]) -> None:
self.original_list = original_list
def __len__(self) -> int:
return len(self.original_list)
def __getitem__(self, i: int) -> str:
return self.original_list[i]
# HFのWhisperはファイルリストを与えるとバッチ処理ができて速い
def transcribe_files_with_hf_whisper(
audio_files: list[Path],
model_id: str,
initial_prompt: Optional[str] = None,
language: str = "ja",
batch_size: int = 16,
num_beams: int = 1,
device: str = "cuda",
pbar: Optional[tqdm] = None,
) -> list[str]:
import torch
from transformers import WhisperProcessor, pipeline
processor: WhisperProcessor = WhisperProcessor.from_pretrained(model_id)
generate_kwargs: dict[str, Any] = {
"language": language,
"do_sample": False,
"num_beams": 5,
"early_stopping": True,
"num_return_sequences": 5,
}
if initial_prompt is not None:
prompt_ids: torch.Tensor = processor.get_prompt_ids(
initial_prompt, return_tensors="pt"
)
prompt_ids = prompt_ids.to(device)
generate_kwargs["prompt_ids"] = prompt_ids
pipe = pipeline(
model=model_id,
max_new_tokens=128,
chunk_length_s=30,
batch_size=batch_size,
torch_dtype=torch.float16,
device="cuda",
generate_kwargs=generate_kwargs,
)
dataset = StrListDataset([str(f) for f in audio_files])
results: list[str] = []
for whisper_result in pipe(dataset):
logger.debug(whisper_result)
for result in enumerate(whisper_result):
logger.debug(result)
logger.debug(f"Transcribed: {result['text']}")
text: str = whisper_result["text"]
# なぜかテキストの最初に" {initial_prompt}"が入るので、文字の最初からこれを削除する
# cf. https://github.com/huggingface/transformers/issues/27594
if text.startswith(f" {initial_prompt}"):
text = text[len(f" {initial_prompt}") :]
results.append(text)
if pbar is not None:
pbar.update(1)
if pbar is not None:
pbar.close()
return results
if __name__ == "__main__": if __name__ == "__main__":
parser = argparse.ArgumentParser() parser = argparse.ArgumentParser()
parser.add_argument("--model_name", type=str, required=True) parser.add_argument("--model_name", type=str, required=True)
@@ -37,6 +118,9 @@ if __name__ == "__main__":
parser.add_argument("--model", type=str, default="large-v3") parser.add_argument("--model", type=str, default="large-v3")
parser.add_argument("--device", type=str, default="cuda") parser.add_argument("--device", type=str, default="cuda")
parser.add_argument("--compute_type", type=str, default="bfloat16") parser.add_argument("--compute_type", type=str, default="bfloat16")
parser.add_argument("--use_hf_whisper", action="store_true")
parser.add_argument("--batch_size", type=int, default=16)
parser.add_argument("--num_beams", type=int, default=1)
args = parser.parse_args() args = parser.parse_args()
@@ -49,22 +133,18 @@ if __name__ == "__main__":
input_dir = dataset_root / model_name / "raw" input_dir = dataset_root / model_name / "raw"
output_file = dataset_root / model_name / "esd.list" output_file = dataset_root / model_name / "esd.list"
initial_prompt: str = args.initial_prompt initial_prompt: str = args.initial_prompt
initial_prompt = initial_prompt.strip('"')
language: str = args.language language: str = args.language
device: str = args.device device: str = args.device
compute_type: str = args.compute_type compute_type: str = args.compute_type
batch_size: int = args.batch_size
num_beams: int = args.num_beams
output_file.parent.mkdir(parents=True, exist_ok=True) output_file.parent.mkdir(parents=True, exist_ok=True)
logger.info(
f"Loading Whisper model ({args.model}) with compute_type={compute_type}"
)
try:
model = WhisperModel(args.model, device=device, compute_type=compute_type)
except ValueError as e:
logger.warning(f"Failed to load model, so use `auto` compute_type: {e}")
model = WhisperModel(args.model, device=device)
wav_files = [f for f in input_dir.rglob("*.wav") if f.is_file()] wav_files = [f for f in input_dir.rglob("*.wav") if f.is_file()]
wav_files = sorted(wav_files, key=lambda x: x.name)
if output_file.exists(): if output_file.exists():
logger.warning(f"{output_file} exists, backing up to {output_file}.bak") logger.warning(f"{output_file} exists, backing up to {output_file}.bak")
backup_path = output_file.with_name(output_file.name + ".bak") backup_path = output_file.with_name(output_file.name + ".bak")
@@ -82,10 +162,42 @@ if __name__ == "__main__":
else: else:
raise ValueError(f"{language} is not supported.") raise ValueError(f"{language} is not supported.")
wav_files = sorted(wav_files, key=lambda x: x.name) logger.info(
f"Loading Whisper model ({args.model}) with compute_type={compute_type}"
)
if not args.use_hf_whisper:
from faster_whisper import WhisperModel
try:
model = WhisperModel(args.model, device=device, compute_type=compute_type)
except ValueError as e:
logger.warning(f"Failed to load model, so use `auto` compute_type: {e}")
model = WhisperModel(args.model, device=device)
for wav_file in tqdm(wav_files, file=SAFE_STDOUT): for wav_file in tqdm(wav_files, file=SAFE_STDOUT):
text = transcribe(wav_file, initial_prompt=initial_prompt, language=language) text = transcribe_with_faster_whisper(
model=model,
audio_file=wav_file,
initial_prompt=initial_prompt,
language=language,
num_beams=num_beams,
)
with open(output_file, "a", encoding="utf-8") as f: with open(output_file, "a", encoding="utf-8") as f:
f.write(f"{wav_file.name}|{model_name}|{language_id}|{text}\n") f.write(f"{wav_file.name}|{model_name}|{language_id}|{text}\n")
else:
model_id = f"openai/whisper-{args.model}"
pbar = tqdm(total=len(wav_files), file=SAFE_STDOUT)
results = transcribe_files_with_hf_whisper(
audio_files=wav_files,
model_id=model_id,
initial_prompt=initial_prompt,
language=language,
batch_size=batch_size,
num_beams=num_beams,
device=device,
pbar=pbar,
)
with open(output_file, "w", encoding="utf-8") as f:
for wav_file, text in zip(wav_files, results):
f.write(f"{wav_file.name}|{model_name}|{language_id}|{text}\n")
sys.exit(0) sys.exit(0)

View File

@@ -41,13 +41,20 @@ def do_slice(
def do_transcribe( def do_transcribe(
model_name, whisper_model, compute_type, language, initial_prompt, device model_name,
whisper_model,
compute_type,
language,
initial_prompt,
device,
use_hf_whisper,
batch_size,
num_beams,
): ):
if model_name == "": if model_name == "":
return "Error: モデル名を入力してください。" return "Error: モデル名を入力してください。"
success, message = run_script_with_log( cmd = [
[
"transcribe.py", "transcribe.py",
"--model_name", "--model_name",
model_name, model_name,
@@ -61,8 +68,13 @@ def do_transcribe(
language, language,
"--initial_prompt", "--initial_prompt",
f'"{initial_prompt}"', f'"{initial_prompt}"',
"--num_beams",
str(num_beams),
] ]
) if use_hf_whisper:
cmd.append("--use_hf_whisper")
cmd.extend(["--batch_size", str(batch_size)])
success, message = run_script_with_log(cmd)
if not success: if not success:
return f"Error: {message}. しかし何故かエラーが起きても正常に終了している場合がほとんどなので、書き起こし結果を確認して問題なければ学習に使えます。" return f"Error: {message}. しかし何故かエラーが起きても正常に終了している場合がほとんどなので、書き起こし結果を確認して問題なければ学習に使えます。"
return "音声の文字起こしが完了しました。" return "音声の文字起こしが完了しました。"
@@ -165,6 +177,9 @@ def create_dataset_app() -> gr.Blocks:
label="Whisperモデル", label="Whisperモデル",
value="large-v3", value="large-v3",
) )
use_hf_whisper = gr.Checkbox(
label="HuggingFaceのWhisperを使う使うと速度が速いがVRAMを多く使う",
)
compute_type = gr.Dropdown( compute_type = gr.Dropdown(
[ [
"int8", "int8",
@@ -186,6 +201,23 @@ def create_dataset_app() -> gr.Blocks:
value="こんにちは。元気、ですかー?ふふっ、私は……ちゃんと元気だよ!", value="こんにちは。元気、ですかー?ふふっ、私は……ちゃんと元気だよ!",
info="このように書き起こしてほしいという例文(句読点の入れ方・笑い方・固有名詞等)", info="このように書き起こしてほしいという例文(句読点の入れ方・笑い方・固有名詞等)",
) )
num_beams = gr.Slider(
minimum=1,
maximum=10,
value=5,
step=1,
label="ビームサーチのビーム数",
info="小さいほど速度が上がり以前は5、精度は少し落ちるかもしれないがほぼ変わらない体感",
)
batch_size = gr.Slider(
minimum=1,
maximum=128,
value=32,
step=1,
label="バッチサイズ",
info="大きくすると速度が速くなるがVRAMを多く使う",
visible=False,
)
transcribe_button = gr.Button("音声の文字起こし") transcribe_button = gr.Button("音声の文字起こし")
result2 = gr.Textbox(label="結果") result2 = gr.Textbox(label="結果")
slice_button.click( slice_button.click(
@@ -210,8 +242,16 @@ def create_dataset_app() -> gr.Blocks:
language, language,
initial_prompt, initial_prompt,
device, device,
use_hf_whisper,
batch_size,
num_beams,
], ],
outputs=[result2], outputs=[result2],
) )
use_hf_whisper.change(
lambda x: gr.update(visible=x),
inputs=[use_hf_whisper],
outputs=[batch_size],
)
return app return app