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,
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)
@router.get("/models_info")
@router.get("/models_info", response_model=list[TTSModelInfo])
def models_info():
return model_holder.models_info

View File

@@ -1,6 +1,6 @@
import warnings
from pathlib import Path
from typing import Any, Optional, Union, TypedDict
from typing import Any, Optional, Union
import gradio as gr
import numpy as np
@@ -8,6 +8,7 @@ import pyannote.audio
import torch
from gradio.processing_utils import convert_to_16_bit_wav
from numpy.typing import NDArray
from pydantic import BaseModel
from style_bert_vits2.constants import (
DEFAULT_ASSIST_TEXT_WEIGHT,
@@ -263,7 +264,7 @@ class TTSModel:
return (self.hyper_parameters.data.sampling_rate, audio)
class TTSModelInfo(TypedDict):
class TTSModelInfo(BaseModel):
name: str
files: list[str]
styles: list[str]
@@ -341,12 +342,12 @@ class TTSModelHolder:
styles = list(style2id.keys())
spk2id: dict[str, int] = hyper_parameters.data.spk2id
speakers = list(spk2id.keys())
self.models_info.append({
"name": model_dir.name,
"files": [str(f) for f in model_files],
"styles": styles,
"speakers": speakers,
})
self.models_info.append(TTSModelInfo(
name = model_dir.name,
files = [str(f) for f in model_files],
styles = styles,
speakers = speakers,
))
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 モデルを探す
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()
sample_rate, audio_data = model.infer(
"あらゆる現実を、すべて自分のほうへねじ曲げたのだ。",