[+] Inference API

This commit is contained in:
2024-07-13 02:51:16 +08:00
parent d9345d73fa
commit a3e0bc1a82
2 changed files with 102 additions and 1 deletions
+6 -1
View File
@@ -9,7 +9,6 @@ import utils
from models import SynthesizerTrn
import gradio as gr
import librosa
import webbrowser
from text import text_to_sequence, _clean_text
device = "cuda:0" if torch.cuda.is_available() else "cpu"
@@ -28,6 +27,8 @@ language_marks = {
"Mix": "",
}
lang = ['日本語', '简体中文', 'English', 'Mix']
def get_text(text, hps, is_symbol):
text_norm = text_to_sequence(text, hps.symbols, [] if is_symbol else hps.data.text_cleaners)
if hps.data.add_blank:
@@ -35,6 +36,7 @@ def get_text(text, hps, is_symbol):
text_norm = LongTensor(text_norm)
return text_norm
def create_tts_fn(model, hps, speaker_ids):
def tts_fn(text, speaker, language, speed):
if language is not None:
@@ -52,6 +54,7 @@ def create_tts_fn(model, hps, speaker_ids):
return tts_fn
def create_vc_fn(model, hps, speaker_ids):
def vc_fn(original_speaker, target_speaker, record_audio, upload_audio):
input_audio = record_audio if record_audio is not None else upload_audio
@@ -83,6 +86,8 @@ def create_vc_fn(model, hps, speaker_ids):
return "Success", (hps.data.sampling_rate, audio)
return vc_fn
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--model_dir", default="./G_latest.pth", help="directory to your fine-tuned model")