Change faster whisper to work with new extension

This commit is contained in:
Sasa Trivic
2024-02-01 14:54:29 +01:00
parent 5b28ddefbd
commit 32ed089a76
+11 -16
View File
@@ -155,7 +155,7 @@ class TranscriptionServer:
options["model"] = faster_whisper_custom_model_path options["model"] = faster_whisper_custom_model_path
client = ServeClientFasterWhisper( client = ServeClientFasterWhisper(
websocket, websocket,
multilingual=options["multilingual"], multilingual=False,
language=options["language"], language=options["language"],
task=options["task"], task=options["task"],
client_uid=options["uid"], client_uid=options["uid"],
@@ -585,13 +585,11 @@ class ServeClientFasterWhisper(ServeClientBase):
"tiny", "tiny.en", "base", "base.en", "small", "small.en", "tiny", "tiny.en", "base", "base.en", "small", "small.en",
"medium", "medium.en", "large-v2", "large-v3", "medium", "medium.en", "large-v2", "large-v3",
] ]
self.multilingual = multilingual
if not os.path.exists(model): if not os.path.exists(model):
self.model_size_or_path = self.get_model_size(model) self.model_size_or_path = self.check_valid_model(model)
else: else:
self.model_size_or_path = model self.model_size_or_path = model
self.language = language if self.multilingual else "en" 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
self.vad_parameters = vad_parameters or {"threshold": 0.5} self.vad_parameters = vad_parameters or {"threshold": 0.5}
@@ -620,9 +618,15 @@ class ServeClientFasterWhisper(ServeClientBase):
) )
) )
def get_model_size(self, model_size): def check_valid_model(self, model_size):
""" """
Returns the whisper model size based on multilingual. Check if it's a valid whisper model size.
Args:
model_size (str): The name of the model size to check.
Returns:
str: The model size if valid, None otherwise.
""" """
if model_size not in self.model_sizes: if model_size not in self.model_sizes:
self.websocket.send( self.websocket.send(
@@ -635,15 +639,6 @@ class ServeClientFasterWhisper(ServeClientBase):
) )
) )
return None return None
if model_size.endswith("en") and self.multilingual:
logging.info(f"Setting multilingual to false with {model_size} which is english only model.")
self.multilingual = False
if not model_size.endswith("en") and not self.multilingual:
logging.info(f"Setting multilingual to true with multilingual model {model_size}.")
self.multilingual = True
return model_size return model_size
def speech_to_text(self): def speech_to_text(self):