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

@@ -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)

View File

@@ -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: