Refactor: TTSModelInfo changed from TypedDict to Pydantic model

Pydantic models are more robust and properties can be accessed by dots.
This commit is contained in:
tsukumi
2024-03-10 19:21:22 +00:00
parent 859d940916
commit 7f02b0f1d5
3 changed files with 14 additions and 13 deletions

View File

@@ -52,7 +52,7 @@ from style_bert_vits2.nlp.japanese.user_dict import (
rewrite_word, rewrite_word,
update_dict, update_dict,
) )
from style_bert_vits2.tts_model import TTSModelHolder from style_bert_vits2.tts_model import TTSModelHolder, TTSModelInfo
# ---フロントエンド部分に関する処理--- # ---フロントエンド部分に関する処理---
@@ -250,7 +250,7 @@ async def normalize(item: TextRequest):
return normalize_text(item.text) return normalize_text(item.text)
@router.get("/models_info") @router.get("/models_info", response_model=list[TTSModelInfo])
def models_info(): def models_info():
return model_holder.models_info return model_holder.models_info

View File

@@ -1,6 +1,6 @@
import warnings import warnings
from pathlib import Path from pathlib import Path
from typing import Any, Optional, Union, TypedDict from typing import Any, Optional, Union
import gradio as gr import gradio as gr
import numpy as np import numpy as np
@@ -8,6 +8,7 @@ import pyannote.audio
import torch import torch
from gradio.processing_utils import convert_to_16_bit_wav from gradio.processing_utils import convert_to_16_bit_wav
from numpy.typing import NDArray from numpy.typing import NDArray
from pydantic import BaseModel
from style_bert_vits2.constants import ( from style_bert_vits2.constants import (
DEFAULT_ASSIST_TEXT_WEIGHT, DEFAULT_ASSIST_TEXT_WEIGHT,
@@ -263,7 +264,7 @@ class TTSModel:
return (self.hyper_parameters.data.sampling_rate, audio) return (self.hyper_parameters.data.sampling_rate, audio)
class TTSModelInfo(TypedDict): class TTSModelInfo(BaseModel):
name: str name: str
files: list[str] files: list[str]
styles: list[str] styles: list[str]
@@ -341,12 +342,12 @@ class TTSModelHolder:
styles = list(style2id.keys()) styles = list(style2id.keys())
spk2id: dict[str, int] = hyper_parameters.data.spk2id spk2id: dict[str, int] = hyper_parameters.data.spk2id
speakers = list(spk2id.keys()) speakers = list(spk2id.keys())
self.models_info.append({ self.models_info.append(TTSModelInfo(
"name": model_dir.name, name = model_dir.name,
"files": [str(f) for f in model_files], files = [str(f) for f in model_files],
"styles": styles, styles = styles,
"speakers": speakers, speakers = speakers,
}) ))
def get_model(self, model_name: str, model_path_str: str) -> TTSModel: def get_model(self, model_name: str, model_path_str: str) -> TTSModel:

View File

@@ -13,12 +13,12 @@ def synthesize(device: str = 'cpu'):
# jvnv-F2-jp モデルを探す # jvnv-F2-jp モデルを探す
for model_info in model_holder.models_info: for model_info in model_holder.models_info:
if model_info['name'] == 'jvnv-F2-jp': if model_info.name == 'jvnv-F2-jp':
# すべてのスタイルに対して音声合成を実行 # すべてのスタイルに対して音声合成を実行
for style in model_info['styles']: for style in model_info.styles:
# 音声合成を実行 # 音声合成を実行
model = model_holder.get_model(model_info['name'], model_info['files'][0]) model = model_holder.get_model(model_info.name, model_info.files[0])
model.load() model.load()
sample_rate, audio_data = model.infer( sample_rate, audio_data = model.infer(
"あらゆる現実を、すべて自分のほうへねじ曲げたのだ。", "あらゆる現実を、すべて自分のほうへねじ曲げたのだ。",