haouarin commited on
Commit
2c2660b
·
verified ·
1 Parent(s): 94ed16a

Create app.py

Browse files
Files changed (1) hide show
  1. app.py +177 -0
app.py ADDED
@@ -0,0 +1,177 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import argparse
2
+ import os
3
+ import soundfile as sf
4
+ import gradio as gr
5
+ import numpy as np
6
+ from groq import Groq
7
+ from models import build_model
8
+ from kokoro import generate
9
+ import torch
10
+
11
+ ### NEW: We'll store the default voice name in a global variable so we can change it later.
12
+ DEFAULT_VOICE_NAME = 'bm_george'
13
+
14
+ def initialize_model():
15
+ device = 'cuda' if torch.cuda.is_available() else 'cpu'
16
+ model = build_model('/content/Kokoro-82M/fp16/kokoro-v0_19-half.pth', device)
17
+
18
+ ### UPDATED: Use our DEFAULT_VOICE_NAME here, so we can swap it out later
19
+ voice_name = DEFAULT_VOICE_NAME
20
+
21
+ voicepack = torch.load(f'voices/{voice_name}.pt', weights_only=True).to(device)
22
+ print(f'Loaded voice: {voice_name}')
23
+ return model, voicepack, device
24
+
25
+ MODEL, VOICEPACK, DEVICE = initialize_model()
26
+ client = None
27
+
28
+ def initialize_groq(api_key):
29
+ """
30
+ Initialize the Groq client with the provided API key.
31
+ """
32
+ global client
33
+ try:
34
+ client = Groq(api_key=api_key)
35
+ return "API key configured successfully"
36
+ except Exception as e:
37
+ return f"Error configuring API key: {str(e)}"
38
+
39
+ ### NEW: Function to change voice pack after the program has started
40
+ def set_voice_name(voice):
41
+ """
42
+ Dynamically load the requested voicepack.
43
+ """
44
+ global VOICEPACK, DEFAULT_VOICE_NAME
45
+ try:
46
+ VOICEPACK = torch.load(f'voices/{voice}.pt', weights_only=True).to(DEVICE)
47
+ DEFAULT_VOICE_NAME = voice # store new default so we can keep track
48
+ return f"Voice changed to {voice}"
49
+ except FileNotFoundError:
50
+ return f"Voice pack {voice} not found in voices/"
51
+
52
+ def answer(question):
53
+ if not client:
54
+ return "Please configure Groq API key first"
55
+ try:
56
+ chat_completion = client.chat.completions.create(
57
+ messages=[
58
+ {
59
+ "role": "system",
60
+ "content": "you are a helpful assistant. your answers must be short"
61
+ },
62
+ {
63
+ "role": "user",
64
+ "content": question
65
+ }
66
+ ],
67
+ model="llama3-8b-8192",
68
+ temperature=0.5,
69
+ max_tokens=1024,
70
+ top_p=1,
71
+ stop=None,
72
+ stream=False,
73
+ )
74
+ return chat_completion.choices[0].message.content
75
+ except Exception as e:
76
+ return f"Error: {str(e)}"
77
+
78
+ def conversation_pipeline(audio_path=None, text_input=None):
79
+ if not client:
80
+ return "Please configure Groq API key first", None, None
81
+
82
+ if audio_path and os.path.exists(audio_path):
83
+ with open(audio_path, "rb") as file:
84
+ transcription = client.audio.transcriptions.create(
85
+ file=(audio_path, file.read()),
86
+ model="distil-whisper-large-v3-en",
87
+ response_format="verbose_json",
88
+ ).text
89
+ elif text_input:
90
+ transcription = text_input
91
+ else:
92
+ return None, None, None
93
+
94
+ gemini_response = answer(transcription)
95
+ generated_audio, _ = generate(MODEL, gemini_response, VOICEPACK, lang="a")
96
+
97
+ return transcription, gemini_response, generated_audio
98
+
99
+ def process_audio(audio=None, text_input=None):
100
+ if audio is not None:
101
+ audio_path = "temp_audio.wav"
102
+ sf.write(audio_path, audio[1], audio[0])
103
+ transcription, response, audio_out = conversation_pipeline(audio_path=audio_path)
104
+ os.remove(audio_path)
105
+ else:
106
+ transcription, response, audio_out = conversation_pipeline(text_input=text_input)
107
+
108
+ if audio_out is not None:
109
+ audio_out = np.array(audio_out).flatten().astype(np.float32)
110
+ return (
111
+ transcription if transcription else "",
112
+ response if response else "",
113
+ (24000, audio_out),
114
+ )
115
+ return "", "", None
116
+
117
+ def main():
118
+ with gr.Blocks() as interface:
119
+ gr.Markdown("# Voice Chat Interface")
120
+ gr.Markdown("Configure API key and start chatting")
121
+
122
+ # Row for API key
123
+ with gr.Row():
124
+ api_key = gr.Textbox(label="Groq API Key", type="password")
125
+ api_status = gr.Textbox(label="API Status", interactive=False)
126
+ configure_btn = gr.Button("Configure API")
127
+
128
+ ### NEW: Row to change voice name on the fly
129
+ with gr.Row():
130
+ voice_name_input = gr.Textbox(label="Voice Name", value=DEFAULT_VOICE_NAME)
131
+ voice_status = gr.Textbox(label="Voice Status", interactive=False)
132
+ set_voice_btn = gr.Button("Set Voice")
133
+
134
+ # Row for user input
135
+ with gr.Row():
136
+ audio_input = gr.Audio(sources=["microphone"], type="numpy", label="Speak")
137
+ text_input = gr.Textbox(label="Or type your message here")
138
+
139
+ # Row for outputs
140
+ with gr.Row():
141
+ transcription = gr.Textbox(label="Transcription", interactive=False)
142
+ response = gr.Textbox(label="AI Response", interactive=False)
143
+
144
+ with gr.Row():
145
+ audio_output = gr.Audio(label="AI Voice Response", autoplay=True)
146
+
147
+ # Configure API key
148
+ configure_btn.click(
149
+ initialize_groq,
150
+ inputs=[api_key],
151
+ outputs=[api_status]
152
+ )
153
+
154
+ ### NEW: Set voice dynamically
155
+ set_voice_btn.click(
156
+ set_voice_name,
157
+ inputs=[voice_name_input],
158
+ outputs=[voice_status]
159
+ )
160
+
161
+ # Process with either audio or text
162
+ audio_input.change(
163
+ process_audio,
164
+ inputs=[audio_input, text_input],
165
+ outputs=[transcription, response, audio_output]
166
+ )
167
+
168
+ text_input.submit(
169
+ process_audio,
170
+ inputs=[audio_input, text_input],
171
+ outputs=[transcription, response, audio_output]
172
+ )
173
+
174
+ interface.launch(share=True)
175
+
176
+ if __name__ == "__main__":
177
+ main()