import torch
import numpy as np
import gradio as gr
from librosa import resample
import logging


def _convert_output(arr):
    return np.vstack([arr / 2, arr / 2]).T


#########
# Model #
#########

from espnet2.bin.tts_inference import Text2Speech

model_tag = "kan-bayashi/vctk_full_band_multi_spk_vits"

model = Text2Speech.from_pretrained(
    model_tag=model_tag,
    device="cpu",
    speed_control_alpha=0.8,
    noise_scale=0.667,
    noise_scale_dur=0.8,
)

def apply_model(text, speaker_id=0, speed=0.5):
    model.decode_conf["alpha"] = (1 - speed) + 0.5
    better_spk_id = (speaker_id + 1) * 10
    sids = torch.from_numpy(np.array([[speaker_id]], dtype=np.int32))
    arr = model(text, sids=sids)["wav"].numpy()
    return model.fs, arr


##########
# Server #
##########


def main(text, speaker_id, speed):
    sr, arr = apply_model(text, speaker_id=speaker_id, speed=speed)
    arr = _convert_output(arr)
    return sr, arr


iface = gr.Interface(
    fn=main, inputs=[
        gr.inputs.Textbox(placeholder="Type sentence here...", label="Text"),
        gr.inputs.Dropdown([
            "Speaker 0",
            "Speaker 1",
            "Speaker 2",
            "Speaker 3",
            "Speaker 4",
            "Speaker 5",
            "Speaker 6",
            "Speaker 7",
            "Speaker 8",
            "Speaker 9",
        ], type="index", default="Speaker 0", label="Speaker ID"),
        gr.inputs.Slider(0.0, 1.0, step=0.01, default=0.5, label="Speed"),
    ], outputs="audio",
    examples=[["Sally sells sea shells on the sea shore.", "Speaker 2", 0.7]],
    allow_screenshot=False, allow_flagging=True, server_name="0.0.0.0", server_port=7860
)
iface.launch(
    ssl=('../cert/383fff33778762e8.crt', '../cert/383fff33778762e8.key'),
)
