Add: Type Hint to null_model_params

This commit is contained in:
tsukumi
2024-10-25 22:37:09 +09:00
parent 375eb7f3a9
commit b0b12696f3
2 changed files with 41 additions and 26 deletions

View File

@@ -7,7 +7,7 @@ from typing import TYPE_CHECKING, Any, Optional, Sequence, Union
import numpy as np
import onnxruntime
from numpy.typing import NDArray
from pydantic import BaseModel
from pydantic import BaseModel, Field
from style_bert_vits2.constants import (
DEFAULT_ASSIST_TEXT_WEIGHT,
@@ -33,6 +33,20 @@ if TYPE_CHECKING:
)
class NullModelParam(BaseModel):
"""
ヌルモデルのパラメータを表す Pydantic モデル。
各パラメータは 0.0 から 1.0 の範囲で指定する。
"""
name: str # モデル名
path: Path # モデルファイルのパス
weight: float = Field(ge=0.0, le=1.0) # 声質の重み
pitch: float = Field(ge=0.0, le=1.0) # 声の高さの重み
style: float = Field(ge=0.0, le=1.0) # 話し方の重み
tempo: float = Field(ge=0.0, le=1.0) # テンポの重み
class TTSModel:
"""
Style-Bert-VITS2 の音声合成モデルを操作するクラス。
@@ -110,7 +124,7 @@ class TTSModel:
# net_g / null_model_params は PyTorch 推論時のみ遅延初期化される
self.net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra, None] = None
self.null_model_params: dict[int, dict[str, Union[float, str]]] = {}
self.null_model_params: Optional[dict[int, NullModelParam]] = None
# onnx_session は ONNX 推論時のみ遅延初期化される
self.onnx_session: Optional[onnxruntime.InferenceSession] = None
@@ -137,14 +151,14 @@ class TTSModel:
)
# ここからはヌルモデルのロード用パラメータが指定されている場合のみ
if len(self.null_model_params.keys()) == 0:
if self.null_model_params is None:
return
# 推論対象のモデルの重みとヌルモデルの重みをマージ
for null_model_info in self.null_model_params.values():
logger.info(f"Adding null model: {null_model_info['path']}...")
logger.info(f"Adding null model: {null_model_info.path}...")
null_model_add = get_net_g(
model_path=str(null_model_info["path"]),
model_path=str(null_model_info.path),
version=self.hyper_parameters.version,
device=self.device,
hps=self.hyper_parameters,
@@ -154,27 +168,27 @@ class TTSModel:
self.net_g.dec.parameters(), null_model_add.dec.parameters()
)
for v in params:
v[0].data.add_(v[1].data, alpha=float(null_model_info["weight"]))
v[0].data.add_(v[1].data, alpha=float(null_model_info.weight))
params = zip(
self.net_g.flow.parameters(), null_model_add.flow.parameters()
)
for v in params:
v[0].data.add_(v[1].data, alpha=float(null_model_info["pitch"]))
v[0].data.add_(v[1].data, alpha=float(null_model_info.pitch))
params = zip(
self.net_g.enc_p.parameters(), null_model_add.enc_p.parameters()
)
for v in params:
v[0].data.add_(v[1].data, alpha=float(null_model_info["style"]))
v[0].data.add_(v[1].data, alpha=float(null_model_info.style))
# テンポは sdp と dp 二つあるからとりあえずどっちも足す
params = zip(
self.net_g.sdp.parameters(), null_model_add.sdp.parameters()
)
for v in params:
v[0].data.add_(v[1].data, alpha=float(null_model_info["tempo"]))
v[0].data.add_(v[1].data, alpha=float(null_model_info.tempo))
params = zip(self.net_g.dp.parameters(), null_model_add.dp.parameters())
for v in params:
v[0].data.add_(v[1].data, alpha=float(null_model_info["tempo"]))
v[0].data.add_(v[1].data, alpha=float(null_model_info.tempo))
logger.info(
f"Null models merged successfully ({time.time() - start_time:.2f}s)"
@@ -334,7 +348,7 @@ class TTSModel:
given_tone: Optional[list[int]] = None,
pitch_scale: float = 1.0,
intonation_scale: float = 1.0,
null_model_params: dict[int, dict[str, Union[str, float]]] = {},
null_model_params: Optional[dict[int, NullModelParam]] = None,
force_reload_model: bool = False,
) -> tuple[int, NDArray[Any]]:
"""
@@ -360,7 +374,7 @@ class TTSModel:
given_tone (Optional[list[int]], optional): アクセントのトーンのリスト. Defaults to None.
pitch_scale (float, optional): ピッチの高さ (1.0 から変更すると若干音質が低下する). Defaults to 1.0.
intonation_scale (float, optional): 抑揚の平均からの変化幅 (1.0 から変更すると若干音質が低下する). Defaults to 1.0.
null_model_params (dict[int, dict[str, Union[str, float]]], optional): 推論時に使用するヌルモデルの名前、適用割合の dict が入った dict 。ONNX 推論では無視される。
null_model_params (Optional[dict[int, NullModelParam]], optional): 推論時に使用するヌルモデルの情報。ONNX 推論では無視される。
force_reload_model (bool, optional): モデルを強制的に再ロードするかどうか. Defaults to False.
Returns:
tuple[int, NDArray[Any]]: サンプリングレートと音声データ (16bit PCM)
@@ -392,10 +406,10 @@ class TTSModel:
from style_bert_vits2.models.infer import infer
if null_model_params is not {}:
if null_model_params is not None:
self.null_model_params = null_model_params
else:
self.null_model_params = {}
self.null_model_params = None
# force_reload_model が True のとき、メモリ上に保持されているモデルを破棄する
if force_reload_model is True: