Refactor: remove dependencies on libraries not required for style_bert_vits2 as a library
librosa: Originally not used under style_bert_vits2/. pyannote.audio: Required only when using the "infer style vector from audio" function during inference, but it is currently almost unused. scipy: Because it was used in only one utility function that was not needed outside of training.
This commit is contained in:
@@ -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',
|
||||
]
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user