Support loading hf models
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
This commit is contained in:
@@ -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
|
||||||
|
|||||||
+40
-29
@@ -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.")
|
||||||
@@ -224,7 +228,7 @@ class TranscriptionServer:
|
|||||||
logging.info("New client connected")
|
logging.info("New client connected")
|
||||||
options = websocket.recv()
|
options = websocket.recv()
|
||||||
options = json.loads(options)
|
options = json.loads(options)
|
||||||
|
|
||||||
if self.client_manager is None:
|
if self.client_manager is None:
|
||||||
max_clients = options.get('max_clients', 4)
|
max_clients = options.get('max_clients', 4)
|
||||||
max_connection_time = options.get('max_connection_time', 600)
|
max_connection_time = options.get('max_connection_time', 600)
|
||||||
@@ -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
|
||||||
@@ -806,15 +807,25 @@ class ServeClientFasterWhisper(ServeClientBase):
|
|||||||
if self.model_size_or_path is None:
|
if self.model_size_or_path is None:
|
||||||
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
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user