add client queue

This commit is contained in:
makaveli10
2023-08-07 22:07:51 +08:00
parent 8bdc58b0de
commit 1a4775bac2
2 changed files with 99 additions and 30 deletions
+2 -1
View File
@@ -1,6 +1,7 @@
import asyncio
from whisper_live.server import TranscriptionServer
if __name__ == "__main__":
server = TranscriptionServer()
server.run("0.0.0.0", 9090)
server.run("0.0.0.0")
+97 -29
View File
@@ -2,11 +2,12 @@ import websockets
import pickle, struct, time, pyaudio
import threading
import os, json
import base64
import wave
import textwrap
import logging
logging.basicConfig(level = logging.INFO)
# logging.basicConfig(level = logging.INFO)
from collections import deque
from dataclasses import dataclass
@@ -14,6 +15,7 @@ from websockets.sync.server import serve
import torch
import numpy as np
import time
from whisper_live.transcriber import WhisperModel
@@ -24,9 +26,19 @@ class TranscriptionServer:
Attributes:
clients (dict): A dictionary to store connected clients.
"""
RATE = 16000
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.websockets = {}
self.clients_start_time = {}
self.max_clients = 4
def recv_audio(self, websocket):
"""
@@ -35,6 +47,18 @@ class TranscriptionServer:
Args:
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 = json.loads(options)
client = ServeClient(
@@ -42,23 +66,55 @@ class TranscriptionServer:
multilingual=options["multilingual"],
language=options["language"],
task=options["task"],
client_uid=options["uid"]
)
self.clients[websocket] = client
# max 10 minutes for each client
self.clients_start_time[websocket] = time.time()
while True:
try:
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)
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:
self.clients[websocket].cleanup()
self.clients.pop(websocket)
logging.info("Connection Closed.")
self.clients_start_time.pop(websocket)
print("Connection Closed.")
print(self.clients)
del websocket
break
def run(self, host, port):
def run(self, host, port=9090):
"""
Run the transcription server.
@@ -73,8 +129,10 @@ class TranscriptionServer:
class ServeClient:
RATE = 16000
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.frames = b""
self.language = language if multilingual else "en"
@@ -86,14 +144,6 @@ class ServeClient:
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.frames_np = None
self.frames_offset = 0.0
@@ -116,7 +166,14 @@ class ServeClient:
self.websocket = websocket
self.trans_thread = threading.Thread(target=self.speech_to_text)
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):
"""
@@ -141,15 +198,6 @@ class ServeClient:
return wrapped
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:
self.frames_offset += 30.0
self.frames_np = self.frames_np[int(30*self.RATE):]
@@ -178,7 +226,8 @@ class ServeClient:
task=self.task
)
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:
if self.exit:
@@ -226,7 +275,12 @@ class ServeClient:
segments = segments + [last_segment]
try:
self.websocket.send(json.dumps(segments))
self.websocket.send(
json.dumps({
"uid": self.client_uid,
"segments": segments
})
)
except Exception as e:
logging.info(f"[ERROR]: {e}")
else:
@@ -245,7 +299,12 @@ class ServeClient:
self.text.append('')
try:
self.websocket.send(json.dumps(segments))
self.websocket.send(
json.dumps({
"uid": self.client_uid,
"segments": segments
})
)
except Exception as e:
logging.info(f"[INFO]: {e}")
except Exception as e:
@@ -320,8 +379,17 @@ class ServeClient:
return last_segment
def disconnect(self):
self.websocket.send(
json.dumps(
{
"uid": self.client_uid,
"message": self.DISCONNECT
}
)
)
def cleanup(self):
logging.info("Cleaning up.")
self.exit = True
self.transcriber.destroy()