CoolFace
Apppublic

6Simple9/ChatTTS-OpenVoice

sourceHugging Facemitupdated 2y agoView on Hugging Face
9likes
app.py157 linesDownload Raw Back to root
1import spaces2import os3import random4import argparse5 6import torch7import gradio as gr8import numpy as np9 10import ChatTTS11 12from OpenVoice import se_extractor13from OpenVoice.api import ToneColorConverter14import soundfile15 16print("loading ChatTTS model...")17chat = ChatTTS.Chat()18chat.load_models()19 20 21def generate_seed():22    new_seed = random.randint(1, 100000000)23    return {24        "__type__": "update",25        "value": new_seed26        }27 28@spaces.GPU29def chat_tts(text, temperature, top_P, top_K, audio_seed_input, text_seed_input, refine_text_flag, refine_text_input, output_path=None):30 31    torch.manual_seed(audio_seed_input)32    rand_spk = torch.randn(768)33    params_infer_code = {34        'spk_emb': rand_spk, 35        'temperature': temperature,36        'top_P': top_P,37        'top_K': top_K,38        }39    params_refine_text = {'prompt': '[oral_2][laugh_0][break_6]'}40    41    torch.manual_seed(text_seed_input)42 43    if refine_text_flag:44        if refine_text_input:45           params_refine_text['prompt'] = refine_text_input46        text = chat.infer(text, 47                          skip_refine_text=False,48                          refine_text_only=True,49                          params_refine_text=params_refine_text,50                          params_infer_code=params_infer_code51                          )52        print("Text has been refined!")53    54    wav = chat.infer(text, 55                     skip_refine_text=True, 56                     params_refine_text=params_refine_text, 57                     params_infer_code=params_infer_code58                     )59    60    audio_data = np.array(wav[0]).flatten()61    sample_rate = 2400062    text_data = text[0] if isinstance(text, list) else text63 64    if output_path is None:65        return [(sample_rate, audio_data), text_data]66    else:67        soundfile.write(output_path, audio_data, sample_rate)68        return text_data69 70# OpenVoice Clone71ckpt_converter = 'OpenVoice/checkpoints/converter'72device = "cuda:0" if torch.cuda.is_available() else "cpu"73 74tone_color_converter = ToneColorConverter(f'{ckpt_converter}/config.json', device=device)75tone_color_converter.load_ckpt(f'{ckpt_converter}/checkpoint.pth')76 77def generate_audio(text, audio_ref, temperature, top_P, top_K, audio_seed_input, text_seed_input, refine_text_flag, refine_text_input):78    save_path = "output.wav"79    80    if audio_ref != "" :81      # Run the base speaker tts82      src_path = "tmp.wav"83      text_data = chat_tts(text, temperature, top_P, top_K, audio_seed_input, text_seed_input, refine_text_flag, refine_text_input, src_path)84      print("Ready for voice cloning!")85    86      source_se, audio_name = se_extractor.get_se(src_path, tone_color_converter, target_dir='processed', vad=True)87      reference_speaker = audio_ref88      target_se, audio_name = se_extractor.get_se(reference_speaker, tone_color_converter, target_dir='processed', vad=True)89 90      print("Get voices segment!")91    92      # Run the tone color converter93      # convert from file94      tone_color_converter.convert(95        audio_src_path=src_path,96        src_se=source_se,97        tgt_se=target_se,98        output_path=save_path)99    else:100      chat_tts(text, temperature, top_P, top_K, audio_seed_input, text_seed_input, refine_text_flag, refine_text_input, save_path)101 102    print("Finished!")103 104    return [save_path, text_data]105 106 107with gr.Blocks() as demo:108    gr.Markdown("# <center>๐Ÿฅณ ChatTTS x OpenVoice ๐Ÿฅณ</center>")109    gr.Markdown("## <center>๐ŸŒŸ Make it sound super natural and switch it up to any voice you want, nailing the mood and tone also!๐ŸŒŸ </center>")110 111    default_text = "Today a man knocked on my door and asked for a small donation toward the local swimming pool. I gave him a glass of water."        112    text_input = gr.Textbox(label="Input Text", lines=4, placeholder="Please Input Text...", value=default_text)113 114 115    default_refine_text = "[oral_2][laugh_0][break_6]"    116    refine_text_input = gr.Textbox(label="Refine Prompt", lines=1, placeholder="Please Refine Prompt...", value=default_refine_text)117    refine_text_checkbox = gr.Checkbox(label="Refine text", info="use oral_(0-9), laugh_(0-2), break_(0-7).'oral' means add filler words, 'laugh' means add laughter, and 'break' means add a pause.", value=True)118    with gr.Column():    119        voice_ref = gr.Audio(label="Reference Audio", type="filepath", value="Examples/speaker.mp3")120 121    with gr.Row():122        temperature_slider = gr.Slider(minimum=0.00001, maximum=1.0, step=0.00001, value=0.3, label="Audio temperature")123        top_p_slider = gr.Slider(minimum=0.1, maximum=0.9, step=0.05, value=0.7, label="top_P")124        top_k_slider = gr.Slider(minimum=1, maximum=20, step=1, value=20, label="top_K")125 126    with gr.Row():127        audio_seed_input = gr.Number(value=42, label="Speaker Seed")128        generate_audio_seed = gr.Button("\U0001F3B2")129        text_seed_input = gr.Number(value=42, label="Text Seed")130        generate_text_seed = gr.Button("\U0001F3B2")131 132    generate_button = gr.Button("Generate")133        134    text_output = gr.Textbox(label="Refined Text", interactive=False)135    audio_output = gr.Audio(label="Output Audio")136 137    generate_audio_seed.click(generate_seed, 138                              inputs=[], 139                              outputs=audio_seed_input)140        141    generate_text_seed.click(generate_seed, 142                             inputs=[], 143                             outputs=text_seed_input)144        145    generate_button.click(generate_audio, 146                          inputs=[text_input, voice_ref, temperature_slider, top_p_slider, top_k_slider, audio_seed_input, text_seed_input, refine_text_checkbox, refine_text_input], 147                          outputs=[audio_output,text_output])148 149parser = argparse.ArgumentParser(description='ChatTTS-OpenVoice Launch')150parser.add_argument('--server_name', type=str, default='0.0.0.0', help='Server name')151parser.add_argument('--server_port', type=int, default=8080, help='Server port')152args = parser.parse_args()153 154# demo.launch(server_name=args.server_name, server_port=args.server_port, inbrowser=True)155 156if __name__ == '__main__':157    demo.launch()