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:
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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(
|
||||||
"あらゆる現実を、すべて自分のほうへねじ曲げたのだ。",
|
"あらゆる現実を、すべて自分のほうへねじ曲げたのだ。",
|
||||||
|
|||||||
Reference in New Issue
Block a user