Refactor: rename Model / ModelHolder to TTSModel / TTSModelHolder for clarification and add comments to each method
This commit is contained in:
@@ -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],
|
||||
)
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user