Refactor: rename Model / ModelHolder to TTSModel / TTSModelHolder for clarification and add comments to each method

This commit is contained in:
tsukumi
2024-03-10 03:47:17 +00:00
parent b7d7c78203
commit d2fd378b56
7 changed files with 133 additions and 56 deletions

View File

@@ -23,7 +23,7 @@ from style_bert_vits2.nlp import bert_models
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 ModelHolder
from style_bert_vits2.tts_model import TTSModelHolder
# pyopenjtalk_worker を起動
@@ -151,7 +151,7 @@ def gr_util(item):
return (gr.update(visible=False), gr.update(visible=True))
def create_inference_app(model_holder: ModelHolder) -> gr.Blocks:
def create_inference_app(model_holder: TTSModelHolder) -> gr.Blocks:
def tts_fn(
model_name,
model_path,
@@ -175,7 +175,7 @@ def create_inference_app(model_holder: ModelHolder) -> gr.Blocks:
pitch_scale,
intonation_scale,
):
model_holder.load_model(model_name, model_path)
model_holder.get_model(model_name, model_path)
assert model_holder.current_model is not None
wrong_tone_message = ""
@@ -218,7 +218,7 @@ def create_inference_app(model_holder: ModelHolder) -> gr.Blocks:
reference_audio_path=reference_audio_path,
sdp_ratio=sdp_ratio,
noise=noise_scale,
noisew=noise_scale_w,
noise_w=noise_scale_w,
length=length_scale,
line_split=line_split,
split_interval=split_interval,
@@ -228,7 +228,7 @@ def create_inference_app(model_holder: ModelHolder) -> gr.Blocks:
style=style,
style_weight=style_weight,
given_tone=tone,
sid=speaker_id,
speaker_id=speaker_id,
pitch_scale=pitch_scale,
intonation_scale=intonation_scale,
)
@@ -459,7 +459,7 @@ def create_inference_app(model_holder: ModelHolder) -> gr.Blocks:
)
load_button.click(
model_holder.load_model_for_gradio,
model_holder.get_model_for_gradio,
inputs=[model_name, model_path],
outputs=[style, tts_button, speaker],
)

View File

@@ -11,7 +11,7 @@ from safetensors.torch import save_file
from style_bert_vits2.constants import DEFAULT_STYLE, GRADIO_THEME
from style_bert_vits2.logging import logger
from style_bert_vits2.tts_model import Model, ModelHolder
from style_bert_vits2.tts_model import TTSModel, TTSModelHolder
voice_keys = ["dec"]
@@ -250,11 +250,11 @@ def simple_tts(model_name, text, style=DEFAULT_STYLE, style_weight=1.0):
config_path = os.path.join(assets_root, model_name, "config.json")
style_vec_path = os.path.join(assets_root, model_name, "style_vectors.npy")
model = Model(Path(model_path), Path(config_path), Path(style_vec_path), device)
model = TTSModel(Path(model_path), Path(config_path), Path(style_vec_path), device)
return model.infer(text, style=style, style_weight=style_weight)
def update_two_model_names_dropdown(model_holder: ModelHolder):
def update_two_model_names_dropdown(model_holder: TTSModelHolder):
new_names, new_files, _ = model_holder.update_model_names_for_gradio()
return new_names, new_files, new_names, new_files
@@ -328,7 +328,7 @@ Happy, Surprise, HappySurprise
"""
def create_merge_app(model_holder: ModelHolder) -> gr.Blocks:
def create_merge_app(model_holder: TTSModelHolder) -> gr.Blocks:
model_names = model_holder.model_names
if len(model_names) == 0:
logger.error(