doby4u/chattts
2
1import os2import random3import argparse4 5import torch6import gradio as gr7import numpy as np8 9import ChatTTS10 11print("loading ChatTTS model...")12chat = ChatTTS.Chat()13chat.load_models()14 15def generate_seed():16 new_seed = random.randint(1, 100000000)17 return {18 "__type__": "update",19 "value": new_seed20 }21 22 23def generate_audio(text, temperature, top_P, top_K, audio_seed_input, text_seed_input, refine_text_flag):24 25 torch.manual_seed(audio_seed_input)26 rand_spk = torch.randn(768)27 params_infer_code = {28 'spk_emb': rand_spk, 29 'temperature': temperature,30 'top_P': top_P,31 'top_K': top_K,32 }33 params_refine_text = {'prompt': '[oral_2][laugh_0][break_6]'}34 35 torch.manual_seed(text_seed_input)36 37 if refine_text_flag:38 text = chat.infer(text, 39 skip_refine_text=False,40 refine_text_only=True,41 params_refine_text=params_refine_text,42 params_infer_code=params_infer_code43 )44 45 wav = chat.infer(text, 46 skip_refine_text=True, 47 params_refine_text=params_refine_text, 48 params_infer_code=params_infer_code49 )50 51 audio_data = np.array(wav[0]).flatten()52 sample_rate = 2400053 text_data = text[0] if isinstance(text, list) else text54 55 return [(sample_rate, audio_data), text_data]56 57 58def main():59 60 61 with gr.Blocks() as demo:62 gr.Markdown("# ChatTTS Webui")63 gr.Markdown("ChatTTS Model: [2noise/ChatTTS](https://github.com/2noise/ChatTTS)")64 65 default_text = "四川美食确实以辣闻名,但也有不辣的选择。比如甜水面、赖汤圆、蛋烘糕、叶儿粑等,这些小吃口味温和,甜而不腻,也很受欢迎。" 66 text_input = gr.Textbox(label="Input Text", lines=4, placeholder="Please Input Text...", value=default_text)67 68 with gr.Row():69 refine_text_checkbox = gr.Checkbox(label="Refine text", value=True)70 temperature_slider = gr.Slider(minimum=0.00001, maximum=1.0, step=0.00001, value=0.3, label="Audio temperature")71 top_p_slider = gr.Slider(minimum=0.1, maximum=0.9, step=0.05, value=0.7, label="top_P")72 top_k_slider = gr.Slider(minimum=1, maximum=20, step=1, value=20, label="top_K")73 74 with gr.Row():75 audio_seed_input = gr.Number(value=42, label="Audio Seed")76 generate_audio_seed = gr.Button("\U0001F3B2")77 text_seed_input = gr.Number(value=42, label="Text Seed")78 generate_text_seed = gr.Button("\U0001F3B2")79 80 generate_button = gr.Button("Generate")81 82 text_output = gr.Textbox(label="Output Text", interactive=False)83 audio_output = gr.Audio(label="Output Audio")84 85 generate_audio_seed.click(generate_seed, 86 inputs=[], 87 outputs=audio_seed_input)88 89 generate_text_seed.click(generate_seed, 90 inputs=[], 91 outputs=text_seed_input)92 93 generate_button.click(generate_audio, 94 inputs=[text_input, temperature_slider, top_p_slider, top_k_slider, audio_seed_input, text_seed_input, refine_text_checkbox], 95 outputs=[audio_output, text_output])96 97 parser = argparse.ArgumentParser(description='ChatTTS demo Launch')98 parser.add_argument('--server_name', type=str, default='0.0.0.0', help='Server name')99 parser.add_argument('--server_port', type=int, default=8080, help='Server port')100 args = parser.parse_args()101 102 # demo.launch(server_name=args.server_name, server_port=args.server_port, inbrowser=True)103 demo.launch()104 105 106if __name__ == '__main__':107 main()