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 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")
+97 -29
View File
@@ -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()