File size: 5,826 Bytes
2c2660b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
import argparse
import os
import soundfile as sf
import gradio as gr
import numpy as np
from groq import Groq
from models import build_model
from kokoro import generate
import torch

### NEW: We'll store the default voice name in a global variable so we can change it later.
DEFAULT_VOICE_NAME = 'bm_george'

def initialize_model():
    device = 'cuda' if torch.cuda.is_available() else 'cpu'
    model = build_model('/content/Kokoro-82M/fp16/kokoro-v0_19-half.pth', device)

    ### UPDATED: Use our DEFAULT_VOICE_NAME here, so we can swap it out later
    voice_name = DEFAULT_VOICE_NAME

    voicepack = torch.load(f'voices/{voice_name}.pt', weights_only=True).to(device)
    print(f'Loaded voice: {voice_name}')
    return model, voicepack, device

MODEL, VOICEPACK, DEVICE = initialize_model()
client = None

def initialize_groq(api_key):
    """
    Initialize the Groq client with the provided API key.
    """
    global client
    try:
        client = Groq(api_key=api_key)
        return "API key configured successfully"
    except Exception as e:
        return f"Error configuring API key: {str(e)}"

### NEW: Function to change voice pack after the program has started
def set_voice_name(voice):
    """
    Dynamically load the requested voicepack.
    """
    global VOICEPACK, DEFAULT_VOICE_NAME
    try:
        VOICEPACK = torch.load(f'voices/{voice}.pt', weights_only=True).to(DEVICE)
        DEFAULT_VOICE_NAME = voice  # store new default so we can keep track
        return f"Voice changed to {voice}"
    except FileNotFoundError:
        return f"Voice pack {voice} not found in voices/"

def answer(question):
    if not client:
        return "Please configure Groq API key first"
    try:
        chat_completion = client.chat.completions.create(
            messages=[
                {
                    "role": "system",
                    "content": "you are a helpful assistant. your answers must be short"
                },
                {
                    "role": "user",
                    "content": question
                }
            ],
            model="llama3-8b-8192",
            temperature=0.5,
            max_tokens=1024,
            top_p=1,
            stop=None,
            stream=False,
        )
        return chat_completion.choices[0].message.content
    except Exception as e:
        return f"Error: {str(e)}"

def conversation_pipeline(audio_path=None, text_input=None):
    if not client:
        return "Please configure Groq API key first", None, None
        
    if audio_path and os.path.exists(audio_path):
        with open(audio_path, "rb") as file:
            transcription = client.audio.transcriptions.create(
                file=(audio_path, file.read()),
                model="distil-whisper-large-v3-en",
                response_format="verbose_json",
            ).text
    elif text_input:
        transcription = text_input
    else:
        return None, None, None

    gemini_response = answer(transcription)
    generated_audio, _ = generate(MODEL, gemini_response, VOICEPACK, lang="a")

    return transcription, gemini_response, generated_audio

def process_audio(audio=None, text_input=None):
    if audio is not None:
        audio_path = "temp_audio.wav"
        sf.write(audio_path, audio[1], audio[0])
        transcription, response, audio_out = conversation_pipeline(audio_path=audio_path)
        os.remove(audio_path)
    else:
        transcription, response, audio_out = conversation_pipeline(text_input=text_input)

    if audio_out is not None:
        audio_out = np.array(audio_out).flatten().astype(np.float32)
        return (
            transcription if transcription else "",
            response if response else "",
            (24000, audio_out),
        )
    return "", "", None

def main():
    with gr.Blocks() as interface:
        gr.Markdown("# Voice Chat Interface")
        gr.Markdown("Configure API key and start chatting")

        # Row for API key
        with gr.Row():
            api_key = gr.Textbox(label="Groq API Key", type="password")
            api_status = gr.Textbox(label="API Status", interactive=False)
            configure_btn = gr.Button("Configure API")

        ### NEW: Row to change voice name on the fly
        with gr.Row():
            voice_name_input = gr.Textbox(label="Voice Name", value=DEFAULT_VOICE_NAME)
            voice_status = gr.Textbox(label="Voice Status", interactive=False)
            set_voice_btn = gr.Button("Set Voice")

        # Row for user input
        with gr.Row():
            audio_input = gr.Audio(sources=["microphone"], type="numpy", label="Speak")
            text_input = gr.Textbox(label="Or type your message here")

        # Row for outputs
        with gr.Row():
            transcription = gr.Textbox(label="Transcription", interactive=False)
            response = gr.Textbox(label="AI Response", interactive=False)

        with gr.Row():
            audio_output = gr.Audio(label="AI Voice Response", autoplay=True)

        # Configure API key
        configure_btn.click(
            initialize_groq,
            inputs=[api_key],
            outputs=[api_status]
        )

        ### NEW: Set voice dynamically
        set_voice_btn.click(
            set_voice_name,
            inputs=[voice_name_input],
            outputs=[voice_status]
        )

        # Process with either audio or text
        audio_input.change(
            process_audio,
            inputs=[audio_input, text_input],
            outputs=[transcription, response, audio_output]
        )

        text_input.submit(
            process_audio,
            inputs=[audio_input, text_input],
            outputs=[transcription, response, audio_output]
        )

    interface.launch(share=True)

if __name__ == "__main__":
    main()