From b1972a3d3d90bc7d2bf8045dba331bacbb0c916d Mon Sep 17 00:00:00 2001 From: litagin02 Date: Thu, 14 Mar 2024 15:10:08 +0900 Subject: [PATCH] Feat: HF whisper for transcribing (faster than faster-whisper) --- slice.py | 3 - tes.py | 0 transcribe.py | 150 +++++++++++++++++++++++++++++++++++++++++------ webui/dataset.py | 76 ++++++++++++++++++------ 4 files changed, 189 insertions(+), 40 deletions(-) create mode 100644 tes.py diff --git a/slice.py b/slice.py index 367c62b..c51e352 100644 --- a/slice.py +++ b/slice.py @@ -1,6 +1,5 @@ import argparse import shutil -import sys from pathlib import Path from queue import Queue from threading import Thread @@ -14,8 +13,6 @@ from tqdm import tqdm from style_bert_vits2.logging import logger from style_bert_vits2.utils.stdout_wrapper import SAFE_STDOUT -# TODO: 並列処理による高速化 - def get_stamps( vad_model: Any, diff --git a/tes.py b/tes.py new file mode 100644 index 0000000..e69de29 diff --git a/transcribe.py b/transcribe.py index f44ecc3..cf61dde 100644 --- a/transcribe.py +++ b/transcribe.py @@ -2,10 +2,10 @@ import argparse import os import sys from pathlib import Path -from typing import Optional +from typing import Any, Optional import yaml -from faster_whisper import WhisperModel +from torch.utils.data import Dataset from tqdm import tqdm 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 -def transcribe( - wav_path: Path, initial_prompt: Optional[str] = None, language: str = "ja" +# faster-whisperは並列処理しても速度が向上しないので、単一モデルでループ処理する +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( - 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] 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__": parser = argparse.ArgumentParser() 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("--device", type=str, default="cuda") 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() @@ -49,22 +133,18 @@ if __name__ == "__main__": input_dir = dataset_root / model_name / "raw" output_file = dataset_root / model_name / "esd.list" initial_prompt: str = args.initial_prompt + initial_prompt = initial_prompt.strip('"') language: str = args.language device: str = args.device 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) - 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 = sorted(wav_files, key=lambda x: x.name) + if output_file.exists(): logger.warning(f"{output_file} exists, backing up to {output_file}.bak") backup_path = output_file.with_name(output_file.name + ".bak") @@ -82,10 +162,42 @@ if __name__ == "__main__": else: 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): + 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: + 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") - for wav_file in tqdm(wav_files, file=SAFE_STDOUT): - text = transcribe(wav_file, initial_prompt=initial_prompt, language=language) - with open(output_file, "a", encoding="utf-8") as f: - f.write(f"{wav_file.name}|{model_name}|{language_id}|{text}\n") sys.exit(0) diff --git a/webui/dataset.py b/webui/dataset.py index 1305e21..ffd1b95 100644 --- a/webui/dataset.py +++ b/webui/dataset.py @@ -41,28 +41,40 @@ def do_slice( 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 == "": return "Error: モデル名を入力してください。" - success, message = run_script_with_log( - [ - "transcribe.py", - "--model_name", - model_name, - "--model", - whisper_model, - "--compute_type", - compute_type, - "--device", - device, - "--language", - language, - "--initial_prompt", - f'"{initial_prompt}"', - ] - ) + cmd = [ + "transcribe.py", + "--model_name", + model_name, + "--model", + whisper_model, + "--compute_type", + compute_type, + "--device", + device, + "--language", + language, + "--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: return f"Error: {message}. しかし何故かエラーが起きても正常に終了している場合がほとんどなので、書き起こし結果を確認して問題なければ学習に使えます。" return "音声の文字起こしが完了しました。" @@ -165,6 +177,9 @@ def create_dataset_app() -> gr.Blocks: label="Whisperモデル", value="large-v3", ) + use_hf_whisper = gr.Checkbox( + label="HuggingFaceのWhisperを使う(使うと速度が速いがVRAMを多く使う)", + ) compute_type = gr.Dropdown( [ "int8", @@ -186,6 +201,23 @@ def create_dataset_app() -> gr.Blocks: value="こんにちは。元気、ですかー?ふふっ、私は……ちゃんと元気だよ!", 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("音声の文字起こし") result2 = gr.Textbox(label="結果") slice_button.click( @@ -210,8 +242,16 @@ def create_dataset_app() -> gr.Blocks: language, initial_prompt, device, + use_hf_whisper, + batch_size, + num_beams, ], outputs=[result2], ) + use_hf_whisper.change( + lambda x: gr.update(visible=x), + inputs=[use_hf_whisper], + outputs=[batch_size], + ) return app