diff --git a/run_server.py b/run_server.py index 728e22b..aff2cd1 100644 --- a/run_server.py +++ b/run_server.py @@ -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") diff --git a/whisper_live/server.py b/whisper_live/server.py index f368bd2..100e955 100644 --- a/whisper_live/server.py +++ b/whisper_live/server.py @@ -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() -