From 2f1c934ea29ddf11164038e98c0a5d924a7df7b7 Mon Sep 17 00:00:00 2001 From: berkaybilik Date: Thu, 27 Jun 2024 00:20:17 +0100 Subject: [PATCH] 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