add client queue
This commit is contained in:
+2
-1
@@ -1,6 +1,7 @@
|
|||||||
|
|
||||||
|
import asyncio
|
||||||
from whisper_live.server import TranscriptionServer
|
from whisper_live.server import TranscriptionServer
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
server = TranscriptionServer()
|
server = TranscriptionServer()
|
||||||
server.run("0.0.0.0", 9090)
|
server.run("0.0.0.0")
|
||||||
|
|||||||
+96
-28
@@ -2,11 +2,12 @@ import websockets
|
|||||||
import pickle, struct, time, pyaudio
|
import pickle, struct, time, pyaudio
|
||||||
import threading
|
import threading
|
||||||
import os, json
|
import os, json
|
||||||
|
import base64
|
||||||
import wave
|
import wave
|
||||||
import textwrap
|
import textwrap
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
logging.basicConfig(level = logging.INFO)
|
# logging.basicConfig(level = logging.INFO)
|
||||||
|
|
||||||
from collections import deque
|
from collections import deque
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
@@ -14,6 +15,7 @@ from websockets.sync.server import serve
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
import time
|
||||||
from whisper_live.transcriber import WhisperModel
|
from whisper_live.transcriber import WhisperModel
|
||||||
|
|
||||||
|
|
||||||
@@ -24,9 +26,19 @@ class TranscriptionServer:
|
|||||||
Attributes:
|
Attributes:
|
||||||
clients (dict): A dictionary to store connected clients.
|
clients (dict): A dictionary to store connected clients.
|
||||||
"""
|
"""
|
||||||
|
RATE = 16000
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
|
# voice activity detection model
|
||||||
|
self.vad_model, _ = torch.hub.load(repo_or_dir='snakers4/silero-vad',
|
||||||
|
model='silero_vad',
|
||||||
|
force_reload=True,
|
||||||
|
onnx=True
|
||||||
|
)
|
||||||
|
self.vad_threshold = 0.4
|
||||||
self.clients = {}
|
self.clients = {}
|
||||||
|
self.websockets = {}
|
||||||
|
self.clients_start_time = {}
|
||||||
|
self.max_clients = 4
|
||||||
|
|
||||||
def recv_audio(self, websocket):
|
def recv_audio(self, websocket):
|
||||||
"""
|
"""
|
||||||
@@ -35,6 +47,18 @@ class TranscriptionServer:
|
|||||||
Args:
|
Args:
|
||||||
websocket (WebSocket): The WebSocket connection for the client.
|
websocket (WebSocket): The WebSocket connection for the client.
|
||||||
"""
|
"""
|
||||||
|
# Check if the maximum number of clients is reached
|
||||||
|
print("New client connected")
|
||||||
|
if len(self.clients) >= self.max_clients:
|
||||||
|
# Send response to the new client to come back later
|
||||||
|
response = {
|
||||||
|
"status": "error",
|
||||||
|
"message": "Server is currently full. Please try again later.",
|
||||||
|
}
|
||||||
|
websocket.send(json.dumps(response))
|
||||||
|
websocket.close()
|
||||||
|
return
|
||||||
|
|
||||||
options = websocket.recv()
|
options = websocket.recv()
|
||||||
options = json.loads(options)
|
options = json.loads(options)
|
||||||
client = ServeClient(
|
client = ServeClient(
|
||||||
@@ -42,23 +66,55 @@ class TranscriptionServer:
|
|||||||
multilingual=options["multilingual"],
|
multilingual=options["multilingual"],
|
||||||
language=options["language"],
|
language=options["language"],
|
||||||
task=options["task"],
|
task=options["task"],
|
||||||
|
client_uid=options["uid"]
|
||||||
)
|
)
|
||||||
|
|
||||||
self.clients[websocket] = client
|
self.clients[websocket] = client
|
||||||
|
# max 10 minutes for each client
|
||||||
|
self.clients_start_time[websocket] = time.time()
|
||||||
|
|
||||||
while True:
|
while True:
|
||||||
try:
|
try:
|
||||||
frame_data = websocket.recv()
|
frame_data = websocket.recv()
|
||||||
frame_np = np.frombuffer(frame_data, np.float32)
|
data = json.loads(frame_data)
|
||||||
|
base64_audio = data["audio"]
|
||||||
|
binary_audio = base64.b64decode(base64_audio)
|
||||||
|
frame_np = np.frombuffer(binary_audio, dtype=np.float32)
|
||||||
|
|
||||||
|
try:
|
||||||
|
speech_prob = self.vad_model(torch.from_numpy(frame_np.copy()), self.RATE).item()
|
||||||
|
if speech_prob < self.vad_threshold:
|
||||||
|
continue
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logging.error(e)
|
||||||
|
return
|
||||||
self.clients[websocket].add_frames(frame_np)
|
self.clients[websocket].add_frames(frame_np)
|
||||||
|
elapsed_time = time.time() - self.clients_start_time[websocket]
|
||||||
|
if elapsed_time >= 45: # 10 minutes in seconds
|
||||||
|
# send a disconnection message
|
||||||
|
self.clients[websocket].disconnect()
|
||||||
|
print(f"{self.clients[websocket]} Client disconnected due to overtime.")
|
||||||
|
print()
|
||||||
|
self.clients[websocket].cleanup()
|
||||||
|
self.clients.pop(websocket)
|
||||||
|
self.clients_start_time.pop(websocket)
|
||||||
|
websocket.close()
|
||||||
|
del websocket
|
||||||
|
break
|
||||||
|
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
self.clients[websocket].cleanup()
|
self.clients[websocket].cleanup()
|
||||||
self.clients.pop(websocket)
|
self.clients.pop(websocket)
|
||||||
logging.info("Connection Closed.")
|
self.clients_start_time.pop(websocket)
|
||||||
|
print("Connection Closed.")
|
||||||
|
print(self.clients)
|
||||||
|
|
||||||
|
del websocket
|
||||||
break
|
break
|
||||||
|
|
||||||
def run(self, host, port):
|
def run(self, host, port=9090):
|
||||||
"""
|
"""
|
||||||
Run the transcription server.
|
Run the transcription server.
|
||||||
|
|
||||||
@@ -73,8 +129,10 @@ class TranscriptionServer:
|
|||||||
class ServeClient:
|
class ServeClient:
|
||||||
RATE = 16000
|
RATE = 16000
|
||||||
SERVER_READY = "SERVER_READY"
|
SERVER_READY = "SERVER_READY"
|
||||||
|
DISCONNECT = "DISCONNECT"
|
||||||
|
|
||||||
def __init__(self, websocket, task="transcribe", device=None, multilingual=False, language=None):
|
def __init__(self, websocket, task="transcribe", device=None, multilingual=False, language=None, client_uid=None):
|
||||||
|
self.client_uid = client_uid
|
||||||
self.data = b""
|
self.data = b""
|
||||||
self.frames = b""
|
self.frames = b""
|
||||||
self.language = language if multilingual else "en"
|
self.language = language if multilingual else "en"
|
||||||
@@ -86,14 +144,6 @@ class ServeClient:
|
|||||||
local_files_only=False,
|
local_files_only=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
# voice activity detection model
|
|
||||||
self.vad_model, _ = torch.hub.load(repo_or_dir='snakers4/silero-vad',
|
|
||||||
model='silero_vad',
|
|
||||||
force_reload=True,
|
|
||||||
onnx=True
|
|
||||||
)
|
|
||||||
self.vad_threshold = 0.4
|
|
||||||
|
|
||||||
self.timestamp_offset = 0.0
|
self.timestamp_offset = 0.0
|
||||||
self.frames_np = None
|
self.frames_np = None
|
||||||
self.frames_offset = 0.0
|
self.frames_offset = 0.0
|
||||||
@@ -116,7 +166,14 @@ class ServeClient:
|
|||||||
self.websocket = websocket
|
self.websocket = websocket
|
||||||
self.trans_thread = threading.Thread(target=self.speech_to_text)
|
self.trans_thread = threading.Thread(target=self.speech_to_text)
|
||||||
self.trans_thread.start()
|
self.trans_thread.start()
|
||||||
self.websocket.send(json.dumps(self.SERVER_READY))
|
self.websocket.send(
|
||||||
|
json.dumps(
|
||||||
|
{
|
||||||
|
"uid": self.client_uid,
|
||||||
|
"message": self.SERVER_READY
|
||||||
|
}
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
def fill_output(self, output):
|
def fill_output(self, output):
|
||||||
"""
|
"""
|
||||||
@@ -141,15 +198,6 @@ class ServeClient:
|
|||||||
return wrapped
|
return wrapped
|
||||||
|
|
||||||
def add_frames(self, frame_np):
|
def add_frames(self, frame_np):
|
||||||
try:
|
|
||||||
speech_prob = self.vad_model(torch.from_numpy(frame_np.copy()), self.RATE).item()
|
|
||||||
if speech_prob < self.vad_threshold:
|
|
||||||
return
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
logging.error(e)
|
|
||||||
return
|
|
||||||
|
|
||||||
if self.frames_np is not None and self.frames_np.shape[0] > 45*self.RATE:
|
if self.frames_np is not None and self.frames_np.shape[0] > 45*self.RATE:
|
||||||
self.frames_offset += 30.0
|
self.frames_offset += 30.0
|
||||||
self.frames_np = self.frames_np[int(30*self.RATE):]
|
self.frames_np = self.frames_np[int(30*self.RATE):]
|
||||||
@@ -178,7 +226,8 @@ class ServeClient:
|
|||||||
task=self.task
|
task=self.task
|
||||||
)
|
)
|
||||||
logging.info(f"Detected language {self.language} with probability {lang_prob}")
|
logging.info(f"Detected language {self.language} with probability {lang_prob}")
|
||||||
self.websocket.send(json.dumps({"language": self.language, "language_prob": lang_prob}))
|
self.websocket.send(json.dumps(
|
||||||
|
{"uid": self.client_uid, "language": self.language, "language_prob": lang_prob}))
|
||||||
|
|
||||||
while True:
|
while True:
|
||||||
if self.exit:
|
if self.exit:
|
||||||
@@ -226,7 +275,12 @@ class ServeClient:
|
|||||||
segments = segments + [last_segment]
|
segments = segments + [last_segment]
|
||||||
|
|
||||||
try:
|
try:
|
||||||
self.websocket.send(json.dumps(segments))
|
self.websocket.send(
|
||||||
|
json.dumps({
|
||||||
|
"uid": self.client_uid,
|
||||||
|
"segments": segments
|
||||||
|
})
|
||||||
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logging.info(f"[ERROR]: {e}")
|
logging.info(f"[ERROR]: {e}")
|
||||||
else:
|
else:
|
||||||
@@ -245,7 +299,12 @@ class ServeClient:
|
|||||||
self.text.append('')
|
self.text.append('')
|
||||||
|
|
||||||
try:
|
try:
|
||||||
self.websocket.send(json.dumps(segments))
|
self.websocket.send(
|
||||||
|
json.dumps({
|
||||||
|
"uid": self.client_uid,
|
||||||
|
"segments": segments
|
||||||
|
})
|
||||||
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logging.info(f"[INFO]: {e}")
|
logging.info(f"[INFO]: {e}")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -320,8 +379,17 @@ class ServeClient:
|
|||||||
|
|
||||||
return last_segment
|
return last_segment
|
||||||
|
|
||||||
|
def disconnect(self):
|
||||||
|
self.websocket.send(
|
||||||
|
json.dumps(
|
||||||
|
{
|
||||||
|
"uid": self.client_uid,
|
||||||
|
"message": self.DISCONNECT
|
||||||
|
}
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
def cleanup(self):
|
def cleanup(self):
|
||||||
logging.info("Cleaning up.")
|
logging.info("Cleaning up.")
|
||||||
self.exit = True
|
self.exit = True
|
||||||
self.transcriber.destroy()
|
self.transcriber.destroy()
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user