Refactor: separate adjust_voice() function from tts_model.py

This commit is contained in:
tsukumi
2024-03-09 17:45:56 +00:00
parent 61e2a1deae
commit 96d22102f3
4 changed files with 127 additions and 102 deletions

View File

@@ -25,54 +25,23 @@ from style_bert_vits2.models.infer import get_net_g, infer
from style_bert_vits2.models.models import SynthesizerTrn from style_bert_vits2.models.models import SynthesizerTrn
from style_bert_vits2.models.models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra from style_bert_vits2.models.models_jp_extra import SynthesizerTrn as SynthesizerTrnJPExtra
from style_bert_vits2.logging import logger from style_bert_vits2.logging import logger
from style_bert_vits2.voice import adjust_voice
def adjust_voice(
fs: int,
wave: NDArray[Any],
pitch_scale: float,
intonation_scale: float,
) -> tuple[int, NDArray[Any]]:
if pitch_scale == 1.0 and intonation_scale == 1.0:
# 初期値の場合は、音質劣化を避けるためにそのまま返す
return fs, wave
try:
import pyworld
except ImportError:
raise ImportError(
"pyworld is not installed. Please install it by `pip install pyworld`"
)
# pyworld で f0 を加工して合成
# pyworld よりもよいのがあるかもしれないが……
## pyworld は Cython で書かれているが、スタブファイルがないため型補完が全く効かない…
wave = wave.astype(np.double)
# 質が高そうだしとりあえずharvestにしておく
f0, t = pyworld.harvest(wave, fs) # type: ignore
sp = pyworld.cheaptrick(wave, f0, t, fs) # type: ignore
ap = pyworld.d4c(wave, f0, t, fs) # type: ignore
non_zero_f0 = [f for f in f0 if f != 0]
f0_mean = sum(non_zero_f0) / len(non_zero_f0)
for i, f in enumerate(f0):
if f == 0:
continue
f0[i] = pitch_scale * f0_mean + intonation_scale * (f - f0_mean)
wave = pyworld.synthesize(f0, sp, ap, fs) # type: ignore
return fs, wave
class Model: class Model:
"""
Style-Bert-Vits2 の音声合成モデルを操作するためのクラス
モデル/ハイパーパラメータ/スタイルベクトルのパスとデバイスを指定して初期化し、model.infer() メソッドを呼び出すと音声合成を行える
"""
def __init__( def __init__(
self, model_path: Path, config_path: Path, style_vec_path: Path, device: str self,
): model_path: Path,
config_path: Path,
style_vec_path: Path,
device: str,
) -> None:
self.model_path: Path = model_path self.model_path: Path = model_path
self.config_path: Path = config_path self.config_path: Path = config_path
self.style_vec_path: Path = style_vec_path self.style_vec_path: Path = style_vec_path
@@ -99,7 +68,8 @@ class Model:
self.net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra, None] = None self.net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra, None] = None
def load_net_g(self):
def load_net_g(self) -> None:
self.net_g = get_net_g( self.net_g = get_net_g(
model_path=str(self.model_path), model_path=str(self.model_path),
version=self.hps.version, version=self.hps.version,
@@ -107,15 +77,15 @@ class Model:
hps=self.hps, hps=self.hps,
) )
def get_style_vector(self, style_id: int, weight: float = 1.0) -> NDArray[Any]: def get_style_vector(self, style_id: int, weight: float = 1.0) -> NDArray[Any]:
mean = self.style_vectors[0] mean = self.style_vectors[0]
style_vec = self.style_vectors[style_id] style_vec = self.style_vectors[style_id]
style_vec = mean + (style_vec - mean) * weight style_vec = mean + (style_vec - mean) * weight
return style_vec return style_vec
def get_style_vector_from_audio(
self, audio_path: str, weight: float = 1.0 def get_style_vector_from_audio(self, audio_path: str, weight: float = 1.0) -> NDArray[Any]:
) -> NDArray[Any]:
from style_gen import get_style_vector from style_gen import get_style_vector
xvec = get_style_vector(audio_path) xvec = get_style_vector(audio_path)
@@ -123,6 +93,7 @@ class Model:
xvec = mean + (xvec - mean) * weight xvec = mean + (xvec - mean) * weight
return xvec return xvec
def infer( def infer(
self, self,
text: str, text: str,
@@ -223,8 +194,13 @@ class Model:
class ModelHolder: class ModelHolder:
def __init__(self, root_dir: Path, device: str): """
self.root_dir: Path = root_dir Style-Bert-Vits2 の音声合成モデルを管理するためのクラス
"""
def __init__(self, model_root_dir: Path, device: str) -> None:
self.root_dir: Path = model_root_dir
self.device: str = device self.device: str = device
self.model_files_dict: dict[str, list[Path]] = {} self.model_files_dict: dict[str, list[Path]] = {}
self.current_model: Optional[Model] = None self.current_model: Optional[Model] = None
@@ -233,7 +209,8 @@ class ModelHolder:
self.models_info: list[dict[str, Union[str, list[str]]]] = [] self.models_info: list[dict[str, Union[str, list[str]]]] = []
self.refresh() self.refresh()
def refresh(self):
def refresh(self) -> None:
self.model_files_dict = {} self.model_files_dict = {}
self.model_names = [] self.model_names = []
self.current_model = None self.current_model = None
@@ -269,7 +246,8 @@ class ModelHolder:
"speakers": speakers, "speakers": speakers,
}) })
def load_model(self, model_name: str, model_path_str: str):
def load_model(self, model_name: str, model_path_str: str) -> Model:
model_path = Path(model_path_str) model_path = Path(model_path_str)
if model_name not in self.model_files_dict: if model_name not in self.model_files_dict:
raise ValueError(f"Model `{model_name}` is not found") raise ValueError(f"Model `{model_name}` is not found")
@@ -284,9 +262,8 @@ class ModelHolder:
) )
return self.current_model return self.current_model
def load_model_gr(
self, model_name: str, model_path_str: str def load_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]:
model_path = Path(model_path_str) model_path = Path(model_path_str)
if model_name not in self.model_files_dict: if model_name not in self.model_files_dict:
raise ValueError(f"Model `{model_name}` is not found") raise ValueError(f"Model `{model_name}` is not found")
@@ -318,11 +295,13 @@ class ModelHolder:
gr.Dropdown(choices=speakers, value=speakers[0]), # type: ignore gr.Dropdown(choices=speakers, value=speakers[0]), # type: ignore
) )
def update_model_files_gr(self, model_name: str) -> gr.Dropdown:
def update_model_files_for_gradio(self, model_name: str) -> gr.Dropdown:
model_files = self.model_files_dict[model_name] model_files = self.model_files_dict[model_name]
return gr.Dropdown(choices=model_files, value=model_files[0]) # type: ignore return gr.Dropdown(choices=model_files, value=model_files[0]) # type: ignore
def update_model_names_gr(self) -> tuple[gr.Dropdown, gr.Dropdown, gr.Button]:
def update_model_names_for_gradio(self) -> tuple[gr.Dropdown, gr.Dropdown, gr.Button]:
self.refresh() self.refresh()
initial_model_name = self.model_names[0] initial_model_name = self.model_names[0]
initial_model_files = self.model_files_dict[initial_model_name] initial_model_files = self.model_files_dict[initial_model_name]

46
style_bert_vits2/voice.py Normal file
View File

@@ -0,0 +1,46 @@
from typing import Any
import numpy as np
from numpy.typing import NDArray
def adjust_voice(
fs: int,
wave: NDArray[Any],
pitch_scale: float = 1.0,
intonation_scale: float = 1.0,
) -> tuple[int, NDArray[Any]]:
if pitch_scale == 1.0 and intonation_scale == 1.0:
# 初期値の場合は、音質劣化を避けるためにそのまま返す
return fs, wave
try:
import pyworld
except ImportError:
raise ImportError(
"pyworld is not installed. Please install it by `pip install pyworld`"
)
# pyworld で f0 を加工して合成
# pyworld よりもよいのがあるかもしれないが……
## pyworld は Cython で書かれているが、スタブファイルがないため型補完が全く効かない…
wave = wave.astype(np.double)
# 質が高そうだしとりあえずharvestにしておく
f0, t = pyworld.harvest(wave, fs) # type: ignore
sp = pyworld.cheaptrick(wave, f0, t, fs) # type: ignore
ap = pyworld.d4c(wave, f0, t, fs) # type: ignore
non_zero_f0 = [f for f in f0 if f != 0]
f0_mean = sum(non_zero_f0) / len(non_zero_f0)
for i, f in enumerate(f0):
if f == 0:
continue
f0[i] = pitch_scale * f0_mean + intonation_scale * (f - f0_mean)
wave = pyworld.synthesize(f0, sp, ap, fs) # type: ignore
return fs, wave

View File

@@ -446,7 +446,7 @@ def create_inference_app(model_holder: ModelHolder) -> gr.Blocks:
) )
model_name.change( model_name.change(
model_holder.update_model_files_gr, model_holder.update_model_files_for_gradio,
inputs=[model_name], inputs=[model_name],
outputs=[model_path], outputs=[model_path],
) )
@@ -454,12 +454,12 @@ def create_inference_app(model_holder: ModelHolder) -> gr.Blocks:
model_path.change(make_non_interactive, outputs=[tts_button]) model_path.change(make_non_interactive, outputs=[tts_button])
refresh_button.click( refresh_button.click(
model_holder.update_model_names_gr, model_holder.update_model_names_for_gradio,
outputs=[model_name, model_path, tts_button], outputs=[model_name, model_path, tts_button],
) )
load_button.click( load_button.click(
model_holder.load_model_gr, model_holder.load_model_for_gradio,
inputs=[model_name, model_path], inputs=[model_name, model_path],
outputs=[style, tts_button, speaker], outputs=[style, tts_button, speaker],
) )

View File

@@ -255,7 +255,7 @@ def simple_tts(model_name, text, style=DEFAULT_STYLE, style_weight=1.0):
def update_two_model_names_dropdown(model_holder: ModelHolder): def update_two_model_names_dropdown(model_holder: ModelHolder):
new_names, new_files, _ = model_holder.update_model_names_gr() new_names, new_files, _ = model_holder.update_model_names_for_gradio()
return new_names, new_files, new_names, new_files return new_names, new_files, new_names, new_files
@@ -444,12 +444,12 @@ def create_merge_app(model_holder: ModelHolder) -> gr.Blocks:
audio_output = gr.Audio(label="結果") audio_output = gr.Audio(label="結果")
model_name_a.change( model_name_a.change(
model_holder.update_model_files_gr, model_holder.update_model_files_for_gradio,
inputs=[model_name_a], inputs=[model_name_a],
outputs=[model_path_a], outputs=[model_path_a],
) )
model_name_b.change( model_name_b.change(
model_holder.update_model_files_gr, model_holder.update_model_files_for_gradio,
inputs=[model_name_b], inputs=[model_name_b],
outputs=[model_path_b], outputs=[model_path_b],
) )