Merge pull request #297 from makaveli10/support_hf_models

Support loading hf models.
This commit is contained in:
Marcus Edel
2024-11-27 13:54:03 -05:00
committed by GitHub
2 changed files with 41 additions and 30 deletions
+1 -1
View File
@@ -87,7 +87,7 @@ client = TranscriptionClient(
9090, 9090,
lang="en", lang="en",
translate=False, translate=False,
model="small", model="small", # also support hf_model => `Systran/faster-whisper-small`
use_vad=False, use_vad=False,
save_output_recording=True, # Only used for microphone input, False by Default save_output_recording=True, # Only used for microphone input, False by Default
output_recording_filename="./output_recording.wav", # Only used for microphone input output_recording_filename="./output_recording.wav", # Only used for microphone input
+38 -27
View File
@@ -181,22 +181,26 @@ class TranscriptionServer:
})) }))
self.backend = BackendType.FASTER_WHISPER self.backend = BackendType.FASTER_WHISPER
if self.backend.is_faster_whisper(): try:
if faster_whisper_custom_model_path is not None and os.path.exists(faster_whisper_custom_model_path): if self.backend.is_faster_whisper():
logging.info(f"Using custom model {faster_whisper_custom_model_path}") if faster_whisper_custom_model_path is not None and os.path.exists(faster_whisper_custom_model_path):
options["model"] = faster_whisper_custom_model_path logging.info(f"Using custom model {faster_whisper_custom_model_path}")
client = ServeClientFasterWhisper( options["model"] = faster_whisper_custom_model_path
websocket, client = ServeClientFasterWhisper(
language=options["language"], websocket,
task=options["task"], language=options["language"],
client_uid=options["uid"], task=options["task"],
model=options["model"], client_uid=options["uid"],
initial_prompt=options.get("initial_prompt"), model=options["model"],
vad_parameters=options.get("vad_parameters"), initial_prompt=options.get("initial_prompt"),
use_vad=self.use_vad, vad_parameters=options.get("vad_parameters"),
single_model=self.single_model, use_vad=self.use_vad,
) single_model=self.single_model,
logging.info("Running faster_whisper backend.") )
logging.info("Running faster_whisper backend.")
except Exception as e:
return
if client is None: if client is None:
raise ValueError(f"Backend type {self.backend.value} not recognised or not handled.") raise ValueError(f"Backend type {self.backend.value} not recognised or not handled.")
@@ -785,10 +789,7 @@ class ServeClientFasterWhisper(ServeClientBase):
"large-v3-turbo", "turbo" "large-v3-turbo", "turbo"
] ]
if not os.path.exists(model): self.model_size_or_path = model
self.model_size_or_path = self.check_valid_model(model)
else:
self.model_size_or_path = model
self.language = "en" if self.model_size_or_path.endswith("en") else language self.language = "en" if self.model_size_or_path.endswith("en") else language
self.task = task self.task = task
self.initial_prompt = initial_prompt self.initial_prompt = initial_prompt
@@ -807,14 +808,24 @@ class ServeClientFasterWhisper(ServeClientBase):
return return
logging.info(f"Using Device={device} with precision {self.compute_type}") logging.info(f"Using Device={device} with precision {self.compute_type}")
if single_model: try:
if ServeClientFasterWhisper.SINGLE_MODEL is None: if single_model:
self.create_model(device) if ServeClientFasterWhisper.SINGLE_MODEL is None:
ServeClientFasterWhisper.SINGLE_MODEL = self.transcriber self.create_model(device)
ServeClientFasterWhisper.SINGLE_MODEL = self.transcriber
else:
self.transcriber = ServeClientFasterWhisper.SINGLE_MODEL
else: else:
self.transcriber = ServeClientFasterWhisper.SINGLE_MODEL self.create_model(device)
else: except Exception as e:
self.create_model(device) logging.error(f"Failed to load model: {e}")
self.websocket.send(json.dumps({
"uid": self.client_uid,
"status": "ERROR",
"message": f"Failed to load model: {str(self.model_size_or_path)}"
}))
self.websocket.close()
return
self.use_vad = use_vad self.use_vad = use_vad