always use the BackendType enum to reference the backend inside the TranscriptionServer

This commit is contained in:
berkaybilik
2024-06-27 00:20:17 +01:00
parent b220ccb330
commit 2f1c934ea2
+11 -11
View File
@@ -133,7 +133,7 @@ class BackendType(Enum):
return [backend_type.value for backend_type in BackendType] return [backend_type.value for backend_type in BackendType]
@staticmethod @staticmethod
def check_validity_of(backend: str) -> bool: def is_valid(backend: str) -> bool:
return backend in BackendType.valid_types() return backend in BackendType.valid_types()
def is_faster_whisper(self) -> bool: def is_faster_whisper(self) -> bool:
@@ -158,7 +158,7 @@ class TranscriptionServer:
): ):
client: Optional[ServeClientBase] = None client: Optional[ServeClientBase] = None
if self.backend == "tensorrt": if self.backend.is_tensorrt():
try: try:
client = ServeClientTensorRT( client = ServeClientTensorRT(
websocket, websocket,
@@ -179,9 +179,9 @@ class TranscriptionServer:
"message": "TensorRT-LLM not supported on Server yet. " "message": "TensorRT-LLM not supported on Server yet. "
"Reverting to available backend: 'faster_whisper'" "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): 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}") logging.info(f"Using custom model {faster_whisper_custom_model_path}")
options["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.") logging.info("Running faster_whisper backend.")
if client is None: 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) self.client_manager.add_client(websocket, client)
@@ -229,7 +229,7 @@ class TranscriptionServer:
websocket.close() websocket.close()
return False # Indicates that the connection should not continue 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.vad_detector = VoiceActivityDetector(frame_rate=self.RATE)
self.initialize_client(websocket, options, faster_whisper_custom_model_path, self.initialize_client(websocket, options, faster_whisper_custom_model_path,
whisper_tensorrt_path, trt_multilingual) whisper_tensorrt_path, trt_multilingual)
@@ -248,11 +248,11 @@ class TranscriptionServer:
frame_np = self.get_audio_from_websocket(websocket) frame_np = self.get_audio_from_websocket(websocket)
client = self.client_manager.get_client(websocket) client = self.client_manager.get_client(websocket)
if frame_np is False: if frame_np is False:
if self.backend == "tensorrt": if self.backend.is_tensorrt():
client.set_eos(True) client.set_eos(True)
return False return False
if self.backend == "tensorrt": if self.backend.is_tensorrt():
voice_active = self.voice_activity(websocket, frame_np) voice_active = self.voice_activity(websocket, frame_np)
if voice_active: if voice_active:
self.no_voice_activity_chunks = 0 self.no_voice_activity_chunks = 0
@@ -265,7 +265,7 @@ class TranscriptionServer:
def recv_audio(self, def recv_audio(self,
websocket, websocket,
backend="faster_whisper", backend: BackendType = BackendType.FASTER_WHISPER,
faster_whisper_custom_model_path=None, faster_whisper_custom_model_path=None,
whisper_tensorrt_path=None, whisper_tensorrt_path=None,
trt_multilingual=False): trt_multilingual=False):
@@ -338,12 +338,12 @@ class TranscriptionServer:
# TODO: load model initially # TODO: load model initially
else: else:
logging.info("Single model mode currently only works with custom models.") 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()}") raise ValueError(f"{backend} is not a valid backend type. Choose backend from {BackendType.valid_types()}")
with serve( with serve(
functools.partial( functools.partial(
self.recv_audio, self.recv_audio,
backend=backend, backend=BackendType(backend),
faster_whisper_custom_model_path=faster_whisper_custom_model_path, faster_whisper_custom_model_path=faster_whisper_custom_model_path,
whisper_tensorrt_path=whisper_tensorrt_path, whisper_tensorrt_path=whisper_tensorrt_path,
trt_multilingual=trt_multilingual trt_multilingual=trt_multilingual