From 28f254f44927c4f6bd989b5f71031ad1db51425f Mon Sep 17 00:00:00 2001 From: tsukumi Date: Tue, 23 Apr 2024 05:45:11 +0000 Subject: [PATCH] Improve: remove Gradio dependency from style_bert_vits2 as a library --- pyproject.toml | 1 - style_bert_vits2/tts_model.py | 65 +++++++++++++++++++++++++++++------ 2 files changed, 55 insertions(+), 11 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 2d26680..c0493a3 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -25,7 +25,6 @@ dependencies = [ 'cmudict', 'cn2an', 'g2p_en', - 'gradio', 'jieba', 'librosa==0.9.2', 'loguru', diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index dbead9e..99a9901 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -1,12 +1,9 @@ -import warnings from pathlib import Path -from typing import Any, Optional, Union +from typing import TYPE_CHECKING, Any, Optional, Union -import gradio as gr import numpy as np import pyannote.audio import torch -from gradio.processing_utils import convert_to_16_bit_wav from numpy.typing import NDArray from pydantic import BaseModel @@ -32,6 +29,13 @@ from style_bert_vits2.models.models_jp_extra import ( from style_bert_vits2.voice import adjust_voice +# Gradio の import は重いため、ここでは型チェック時のみ import する +# ライブラリとしての利用を考慮し、TTSModelHolder の _for_gradio() 系メソッド以外では Gradio に依存しないようにする +# _for_gradio() 系メソッドの戻り値の型アノテーションを文字列としているのは、Gradio なしで実行できるようにするため +if TYPE_CHECKING: + import gradio as gr + + class TTSModel: """ Style-Bert-Vits2 の音声合成モデルを操作するクラス。 @@ -137,6 +141,43 @@ class TTSModel: xvec = mean + (xvec - mean) * weight return xvec + def __convert_to_16_bit_wav(self, data: NDArray[Any]) -> NDArray[Any]: + """ + 音声データを 16-bit int 形式に変換する。 + gradio.processing_utils.convert_to_16_bit_wav() を移植したもの。 + + Args: + data (NDArray[Any]): 音声データ + + Returns: + NDArray[Any]: 16-bit int 形式の音声データ + """ + # Based on: https://docs.scipy.org/doc/scipy/reference/generated/scipy.io.wavfile.write.html + if data.dtype in [np.float64, np.float32, np.float16]: # type: ignore + data = data / np.abs(data).max() + data = data * 32767 + data = data.astype(np.int16) + elif data.dtype == np.int32: + data = data / 65536 + data = data.astype(np.int16) + elif data.dtype == np.int16: + pass + elif data.dtype == np.uint16: + data = data - 32768 + data = data.astype(np.int16) + elif data.dtype == np.uint8: + data = data * 257 - 32768 + data = data.astype(np.int16) + elif data.dtype == np.int8: + data = data * 256 + data = data.astype(np.int16) + else: + raise ValueError( + "Audio data cannot be converted automatically from " + f"{data.dtype} to 16-bit int format." + ) + return data + def infer( self, text: str, @@ -260,9 +301,7 @@ class TTSModel: pitch_scale=pitch_scale, intonation_scale=intonation_scale, ) - with warnings.catch_warnings(): - warnings.simplefilter("ignore") - audio = convert_to_16_bit_wav(audio) + audio = self.__convert_to_16_bit_wav(audio) return (self.hyper_parameters.data.sampling_rate, audio) @@ -381,7 +420,9 @@ class TTSModelHolder: def get_model_for_gradio( self, model_name: str, model_path_str: str - ) -> tuple[gr.Dropdown, gr.Button, gr.Dropdown]: + ) -> tuple["gr.Dropdown", "gr.Button", "gr.Dropdown"]: + import gradio as gr + model_path = Path(model_path_str) if model_name not in self.model_files_dict: raise ValueError(f"Model `{model_name}` is not found") @@ -413,13 +454,17 @@ class TTSModelHolder: gr.Dropdown(choices=speakers, value=speakers[0]), # type: ignore ) - def update_model_files_for_gradio(self, model_name: str) -> gr.Dropdown: + def update_model_files_for_gradio(self, model_name: str) -> "gr.Dropdown": + import gradio as gr + model_files = self.model_files_dict[model_name] return gr.Dropdown(choices=model_files, value=model_files[0]) # type: ignore def update_model_names_for_gradio( self, - ) -> tuple[gr.Dropdown, gr.Dropdown, gr.Button]: + ) -> tuple["gr.Dropdown", "gr.Dropdown", "gr.Button"]: + import gradio as gr + self.refresh() initial_model_name = self.model_names[0] initial_model_files = self.model_files_dict[initial_model_name]