diff --git a/pyproject.toml b/pyproject.toml index c0493a3..77306ad 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -26,18 +26,15 @@ dependencies = [ 'cn2an', 'g2p_en', 'jieba', - 'librosa==0.9.2', 'loguru', 'num2words', 'numba', 'numpy', - 'pyannote.audio>=3.1.0', 'pydantic>=2.0', 'pyopenjtalk-dict', 'pypinyin', 'pyworld-prebuilt', 'safetensors', - 'scipy', 'torch>=2.1', 'transformers', ] diff --git a/style_bert_vits2/models/utils/__init__.py b/style_bert_vits2/models/utils/__init__.py index edd51cc..2e750c8 100644 --- a/style_bert_vits2/models/utils/__init__.py +++ b/style_bert_vits2/models/utils/__init__.py @@ -9,7 +9,6 @@ from typing import TYPE_CHECKING, Any, Optional, Union import numpy as np import torch from numpy.typing import NDArray -from scipy.io.wavfile import read from style_bert_vits2.logging import logger from style_bert_vits2.models.utils import checkpoints # type: ignore @@ -162,6 +161,13 @@ def load_wav_to_torch(full_path: Union[str, Path]) -> tuple[torch.FloatTensor, i tuple[torch.FloatTensor, int]: 音声データのテンソルとサンプリングレート """ + # この関数は学習時以外使われないため、ライブラリとしての style_bert_vits2 が + # 重たい scipy に依存しないように遅延 import する + try: + from scipy.io.wavfile import read + except ImportError: + raise ImportError("scipy is required to load wav file") + sampling_rate, data = read(full_path) return torch.FloatTensor(data.astype(np.float32)), sampling_rate diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index 99a9901..83ebcad 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -2,7 +2,6 @@ from pathlib import Path from typing import TYPE_CHECKING, Any, Optional, Union import numpy as np -import pyannote.audio import torch from numpy.typing import NDArray from pydantic import BaseModel @@ -76,7 +75,7 @@ class TTSModel: f"Number of styles ({num_styles}) does not match the number of style2id ({len(self.style2id)})" ) - self.__style_vector_inference: Optional[pyannote.audio.Inference] = None + self.__style_vector_inference: Optional[Any] = None self.__style_vectors: NDArray[Any] = np.load(self.style_vec_path) if self.__style_vectors.shape[0] != num_styles: raise ValueError( @@ -125,8 +124,18 @@ class TTSModel: NDArray[Any]: スタイルベクトル """ - # スタイルベクトルを取得するための推論モデルを初期化 if self.__style_vector_inference is None: + + # pyannote.audio は scikit-learn などの大量の重量級ライブラリに依存しているため、 + # TTSModel.infer() に reference_audio_path を指定し音声からスタイルベクトルを推論する場合のみ遅延 import する + try: + import pyannote.audio + except ImportError: + raise ImportError( + "pyannote.audio is required to infer style vector from audio" + ) + + # スタイルベクトルを取得するための推論モデルを初期化 self.__style_vector_inference = pyannote.audio.Inference( model=pyannote.audio.Model.from_pretrained( "pyannote/wespeaker-voxceleb-resnet34-LM"