Add: Type Hint to null_model_params
This commit is contained in:
@@ -1,6 +1,7 @@
|
|||||||
import datetime
|
import datetime
|
||||||
import json
|
import json
|
||||||
from typing import Any, Optional, Union
|
from pathlib import Path
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
import gradio as gr
|
import gradio as gr
|
||||||
|
|
||||||
@@ -22,7 +23,7 @@ from style_bert_vits2.nlp import InvalidToneError
|
|||||||
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk
|
from style_bert_vits2.nlp.japanese import pyopenjtalk_worker as pyopenjtalk
|
||||||
from style_bert_vits2.nlp.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone
|
from style_bert_vits2.nlp.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone
|
||||||
from style_bert_vits2.nlp.japanese.normalizer import normalize_text
|
from style_bert_vits2.nlp.japanese.normalizer import normalize_text
|
||||||
from style_bert_vits2.tts_model import TTSModelHolder
|
from style_bert_vits2.tts_model import NullModelParam, TTSModelHolder
|
||||||
from style_bert_vits2.utils import torch_device_to_onnx_providers
|
from style_bert_vits2.utils import torch_device_to_onnx_providers
|
||||||
|
|
||||||
|
|
||||||
@@ -218,16 +219,16 @@ def change_null_model_row(
|
|||||||
null_voice_pitch_weights: float,
|
null_voice_pitch_weights: float,
|
||||||
null_speech_style_weights: float,
|
null_speech_style_weights: float,
|
||||||
null_tempo_weights: float,
|
null_tempo_weights: float,
|
||||||
null_models: dict[int, dict[str, Any]],
|
null_models: dict[int, NullModelParam],
|
||||||
):
|
):
|
||||||
null_models[null_model_index] = {
|
null_models[null_model_index] = NullModelParam(
|
||||||
"name": null_model_name,
|
name=null_model_name,
|
||||||
"path": null_model_path,
|
path=Path(null_model_path),
|
||||||
"weight": null_voice_weights,
|
weight=null_voice_weights,
|
||||||
"pitch": null_voice_pitch_weights,
|
pitch=null_voice_pitch_weights,
|
||||||
"style": null_speech_style_weights,
|
style=null_speech_style_weights,
|
||||||
"tempo": null_tempo_weights,
|
tempo=null_tempo_weights,
|
||||||
}
|
)
|
||||||
if len(null_models) > null_models_frame:
|
if len(null_models) > null_models_frame:
|
||||||
keys_to_keep = list(range(null_models_frame))
|
keys_to_keep = list(range(null_models_frame))
|
||||||
result = {k: null_models[k] for k in keys_to_keep}
|
result = {k: null_models[k] for k in keys_to_keep}
|
||||||
@@ -259,7 +260,7 @@ def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks:
|
|||||||
speaker,
|
speaker,
|
||||||
pitch_scale,
|
pitch_scale,
|
||||||
intonation_scale,
|
intonation_scale,
|
||||||
null_models: dict[int, dict[str, Union[str, float]]],
|
null_models: dict[int, NullModelParam],
|
||||||
force_reload_model: bool,
|
force_reload_model: bool,
|
||||||
):
|
):
|
||||||
model_holder.get_model(model_name, model_path)
|
model_holder.get_model(model_name, model_path)
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ from typing import TYPE_CHECKING, Any, Optional, Sequence, Union
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
import onnxruntime
|
import onnxruntime
|
||||||
from numpy.typing import NDArray
|
from numpy.typing import NDArray
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
from style_bert_vits2.constants import (
|
from style_bert_vits2.constants import (
|
||||||
DEFAULT_ASSIST_TEXT_WEIGHT,
|
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:
|
class TTSModel:
|
||||||
"""
|
"""
|
||||||
Style-Bert-VITS2 の音声合成モデルを操作するクラス。
|
Style-Bert-VITS2 の音声合成モデルを操作するクラス。
|
||||||
@@ -110,7 +124,7 @@ class TTSModel:
|
|||||||
|
|
||||||
# net_g / null_model_params は PyTorch 推論時のみ遅延初期化される
|
# net_g / null_model_params は PyTorch 推論時のみ遅延初期化される
|
||||||
self.net_g: Union[SynthesizerTrn, SynthesizerTrnJPExtra, None] = None
|
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 推論時のみ遅延初期化される
|
# onnx_session は ONNX 推論時のみ遅延初期化される
|
||||||
self.onnx_session: Optional[onnxruntime.InferenceSession] = None
|
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
|
return
|
||||||
|
|
||||||
# 推論対象のモデルの重みとヌルモデルの重みをマージ
|
# 推論対象のモデルの重みとヌルモデルの重みをマージ
|
||||||
for null_model_info in self.null_model_params.values():
|
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(
|
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,
|
version=self.hyper_parameters.version,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
hps=self.hyper_parameters,
|
hps=self.hyper_parameters,
|
||||||
@@ -154,27 +168,27 @@ class TTSModel:
|
|||||||
self.net_g.dec.parameters(), null_model_add.dec.parameters()
|
self.net_g.dec.parameters(), null_model_add.dec.parameters()
|
||||||
)
|
)
|
||||||
for v in params:
|
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(
|
params = zip(
|
||||||
self.net_g.flow.parameters(), null_model_add.flow.parameters()
|
self.net_g.flow.parameters(), null_model_add.flow.parameters()
|
||||||
)
|
)
|
||||||
for v in params:
|
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(
|
params = zip(
|
||||||
self.net_g.enc_p.parameters(), null_model_add.enc_p.parameters()
|
self.net_g.enc_p.parameters(), null_model_add.enc_p.parameters()
|
||||||
)
|
)
|
||||||
for v in params:
|
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 二つあるからとりあえずどっちも足す
|
# テンポは sdp と dp 二つあるからとりあえずどっちも足す
|
||||||
params = zip(
|
params = zip(
|
||||||
self.net_g.sdp.parameters(), null_model_add.sdp.parameters()
|
self.net_g.sdp.parameters(), null_model_add.sdp.parameters()
|
||||||
)
|
)
|
||||||
for v in params:
|
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())
|
params = zip(self.net_g.dp.parameters(), null_model_add.dp.parameters())
|
||||||
for v in params:
|
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(
|
logger.info(
|
||||||
f"Null models merged successfully ({time.time() - start_time:.2f}s)"
|
f"Null models merged successfully ({time.time() - start_time:.2f}s)"
|
||||||
@@ -334,7 +348,7 @@ class TTSModel:
|
|||||||
given_tone: Optional[list[int]] = None,
|
given_tone: Optional[list[int]] = None,
|
||||||
pitch_scale: float = 1.0,
|
pitch_scale: float = 1.0,
|
||||||
intonation_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,
|
force_reload_model: bool = False,
|
||||||
) -> tuple[int, NDArray[Any]]:
|
) -> tuple[int, NDArray[Any]]:
|
||||||
"""
|
"""
|
||||||
@@ -360,7 +374,7 @@ class TTSModel:
|
|||||||
given_tone (Optional[list[int]], optional): アクセントのトーンのリスト. Defaults to None.
|
given_tone (Optional[list[int]], optional): アクセントのトーンのリスト. Defaults to None.
|
||||||
pitch_scale (float, optional): ピッチの高さ (1.0 から変更すると若干音質が低下する). Defaults to 1.0.
|
pitch_scale (float, optional): ピッチの高さ (1.0 から変更すると若干音質が低下する). Defaults to 1.0.
|
||||||
intonation_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.
|
force_reload_model (bool, optional): モデルを強制的に再ロードするかどうか. Defaults to False.
|
||||||
Returns:
|
Returns:
|
||||||
tuple[int, NDArray[Any]]: サンプリングレートと音声データ (16bit PCM)
|
tuple[int, NDArray[Any]]: サンプリングレートと音声データ (16bit PCM)
|
||||||
@@ -392,10 +406,10 @@ class TTSModel:
|
|||||||
|
|
||||||
from style_bert_vits2.models.infer import infer
|
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
|
self.null_model_params = null_model_params
|
||||||
else:
|
else:
|
||||||
self.null_model_params = {}
|
self.null_model_params = None
|
||||||
|
|
||||||
# force_reload_model が True のとき、メモリ上に保持されているモデルを破棄する
|
# force_reload_model が True のとき、メモリ上に保持されているモデルを破棄する
|
||||||
if force_reload_model is True:
|
if force_reload_model is True:
|
||||||
|
|||||||
Reference in New Issue
Block a user