From b0b12696f3b0a0b6917187120c3206d9bd68966c Mon Sep 17 00:00:00 2001 From: tsukumi Date: Fri, 25 Oct 2024 22:37:09 +0900 Subject: [PATCH] Add: Type Hint to null_model_params --- gradio_tabs/inference.py | 25 +++++++++++---------- style_bert_vits2/tts_model.py | 42 +++++++++++++++++++++++------------ 2 files changed, 41 insertions(+), 26 deletions(-) diff --git a/gradio_tabs/inference.py b/gradio_tabs/inference.py index 83f6b75..7898283 100644 --- a/gradio_tabs/inference.py +++ b/gradio_tabs/inference.py @@ -1,6 +1,7 @@ import datetime import json -from typing import Any, Optional, Union +from pathlib import Path +from typing import Optional 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.g2p_utils import g2kata_tone, kata_tone2phone_tone 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 @@ -218,16 +219,16 @@ def change_null_model_row( null_voice_pitch_weights: float, null_speech_style_weights: float, null_tempo_weights: float, - null_models: dict[int, dict[str, Any]], + null_models: dict[int, NullModelParam], ): - null_models[null_model_index] = { - "name": null_model_name, - "path": null_model_path, - "weight": null_voice_weights, - "pitch": null_voice_pitch_weights, - "style": null_speech_style_weights, - "tempo": null_tempo_weights, - } + null_models[null_model_index] = NullModelParam( + name=null_model_name, + path=Path(null_model_path), + weight=null_voice_weights, + pitch=null_voice_pitch_weights, + style=null_speech_style_weights, + tempo=null_tempo_weights, + ) if len(null_models) > null_models_frame: keys_to_keep = list(range(null_models_frame)) 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, pitch_scale, intonation_scale, - null_models: dict[int, dict[str, Union[str, float]]], + null_models: dict[int, NullModelParam], force_reload_model: bool, ): model_holder.get_model(model_name, model_path) diff --git a/style_bert_vits2/tts_model.py b/style_bert_vits2/tts_model.py index a4904e4..76acfe1 100644 --- a/style_bert_vits2/tts_model.py +++ b/style_bert_vits2/tts_model.py @@ -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: