Change faster whisper to work with new extension
This commit is contained in:
+12
-17
@@ -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}
|
||||||
@@ -600,7 +598,7 @@ class ServeClientFasterWhisper(ServeClientBase):
|
|||||||
|
|
||||||
if self.model_size_or_path == None:
|
if self.model_size_or_path == None:
|
||||||
return
|
return
|
||||||
|
|
||||||
self.transcriber = WhisperModel(
|
self.transcriber = WhisperModel(
|
||||||
self.model_size_or_path,
|
self.model_size_or_path,
|
||||||
device=device,
|
device=device,
|
||||||
@@ -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):
|
||||||
|
|||||||
Reference in New Issue
Block a user