From 5e3906fc7bbd4f701d60455d77ce284bfe5e7873 Mon Sep 17 00:00:00 2001 From: berkaybilik Date: Wed, 26 Jun 2024 23:59:07 +0100 Subject: [PATCH 1/3] use enum to validate backend validity in server.run --- whisper_live/server.py | 24 ++++++++++++++++++++++++ 1 file changed, 24 insertions(+) diff --git a/whisper_live/server.py b/whisper_live/server.py index 2accaac..5815d20 100644 --- a/whisper_live/server.py +++ b/whisper_live/server.py @@ -4,6 +4,9 @@ import threading import json import functools import logging +from enum import Enum +from typing import List + import torch import numpy as np from websockets.sync.server import serve @@ -121,6 +124,25 @@ class ClientManager: return False +class BackendType(Enum): + FASTER_WHISPER = "faster_whisper" + TENSORRT = "tensorrt" + + @staticmethod + def valid_types() -> List[str]: + return [backend_type.value for backend_type in BackendType] + + @staticmethod + def check_validity_of(backend: str) -> bool: + return backend in BackendType.valid_types() + + def is_faster_whisper(self) -> bool: + return self == BackendType.FASTER_WHISPER + + def is_tensorrt(self) -> bool: + return self == BackendType.TENSORRT + + class TranscriptionServer: RATE = 16000 @@ -311,6 +333,8 @@ class TranscriptionServer: # TODO: load model initially else: logging.info("Single model mode currently only works with custom models.") + if BackendType.check_validity_of(backend): + raise ValueError(f"{backend} is not a valid backend type. Choose backend from {BackendType.valid_types()}") with serve( functools.partial( self.recv_audio, From b220ccb3309b8959215e88d47e8dfd30fe1234e8 Mon Sep 17 00:00:00 2001 From: berkaybilik Date: Thu, 27 Jun 2024 00:11:06 +0100 Subject: [PATCH 2/3] fixed reference before assignment error/warning --- whisper_live/server.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/whisper_live/server.py b/whisper_live/server.py index 5815d20..749ed78 100644 --- a/whisper_live/server.py +++ b/whisper_live/server.py @@ -5,7 +5,7 @@ import json import functools import logging from enum import Enum -from typing import List +from typing import List, Optional import torch import numpy as np @@ -156,6 +156,8 @@ class TranscriptionServer: self, websocket, options, faster_whisper_custom_model_path, whisper_tensorrt_path, trt_multilingual ): + client: Optional[ServeClientBase] = None + if self.backend == "tensorrt": try: client = ServeClientTensorRT( @@ -196,6 +198,9 @@ class TranscriptionServer: ) logging.info("Running faster_whisper backend.") + if client is None: + raise ValueError(f"Backend type {self.backend} not recognised or not handled.") + self.client_manager.add_client(websocket, client) def get_audio_from_websocket(self, websocket): From 2f1c934ea29ddf11164038e98c0a5d924a7df7b7 Mon Sep 17 00:00:00 2001 From: berkaybilik Date: Thu, 27 Jun 2024 00:20:17 +0100 Subject: [PATCH 3/3] always use the BackendType enum to reference the backend inside the TranscriptionServer --- whisper_live/server.py | 22 +++++++++++----------- 1 file changed, 11 insertions(+), 11 deletions(-) diff --git a/whisper_live/server.py b/whisper_live/server.py index 749ed78..b6efc77 100644 --- a/whisper_live/server.py +++ b/whisper_live/server.py @@ -133,7 +133,7 @@ class BackendType(Enum): return [backend_type.value for backend_type in BackendType] @staticmethod - def check_validity_of(backend: str) -> bool: + def is_valid(backend: str) -> bool: return backend in BackendType.valid_types() def is_faster_whisper(self) -> bool: @@ -158,7 +158,7 @@ class TranscriptionServer: ): client: Optional[ServeClientBase] = None - if self.backend == "tensorrt": + if self.backend.is_tensorrt(): try: client = ServeClientTensorRT( websocket, @@ -179,9 +179,9 @@ class TranscriptionServer: "message": "TensorRT-LLM not supported on Server yet. " "Reverting to available backend: 'faster_whisper'" })) - self.backend = "faster_whisper" + self.backend = BackendType.FASTER_WHISPER - if self.backend == "faster_whisper": + if self.backend.is_faster_whisper(): if faster_whisper_custom_model_path is not None and os.path.exists(faster_whisper_custom_model_path): logging.info(f"Using custom model {faster_whisper_custom_model_path}") options["model"] = faster_whisper_custom_model_path @@ -199,7 +199,7 @@ class TranscriptionServer: logging.info("Running faster_whisper backend.") if client is None: - raise ValueError(f"Backend type {self.backend} not recognised or not handled.") + raise ValueError(f"Backend type {self.backend.value} not recognised or not handled.") self.client_manager.add_client(websocket, client) @@ -229,7 +229,7 @@ class TranscriptionServer: websocket.close() return False # Indicates that the connection should not continue - if self.backend == "tensorrt": + if self.backend.is_tensorrt(): self.vad_detector = VoiceActivityDetector(frame_rate=self.RATE) self.initialize_client(websocket, options, faster_whisper_custom_model_path, whisper_tensorrt_path, trt_multilingual) @@ -248,11 +248,11 @@ class TranscriptionServer: frame_np = self.get_audio_from_websocket(websocket) client = self.client_manager.get_client(websocket) if frame_np is False: - if self.backend == "tensorrt": + if self.backend.is_tensorrt(): client.set_eos(True) return False - if self.backend == "tensorrt": + if self.backend.is_tensorrt(): voice_active = self.voice_activity(websocket, frame_np) if voice_active: self.no_voice_activity_chunks = 0 @@ -265,7 +265,7 @@ class TranscriptionServer: def recv_audio(self, websocket, - backend="faster_whisper", + backend: BackendType = BackendType.FASTER_WHISPER, faster_whisper_custom_model_path=None, whisper_tensorrt_path=None, trt_multilingual=False): @@ -338,12 +338,12 @@ class TranscriptionServer: # TODO: load model initially else: logging.info("Single model mode currently only works with custom models.") - if BackendType.check_validity_of(backend): + if not BackendType.is_valid(backend): raise ValueError(f"{backend} is not a valid backend type. Choose backend from {BackendType.valid_types()}") with serve( functools.partial( self.recv_audio, - backend=backend, + backend=BackendType(backend), faster_whisper_custom_model_path=faster_whisper_custom_model_path, whisper_tensorrt_path=whisper_tensorrt_path, trt_multilingual=trt_multilingual