Merge branch 'main' of github.com:lightwastak3n/WhisperLive

This commit is contained in:
Sasa Trivic
2024-02-01 17:37:18 +01:00
6 changed files with 114 additions and 147 deletions
+12 -17
View File
@@ -155,7 +155,7 @@ class TranscriptionServer:
options["model"] = faster_whisper_custom_model_path
client = ServeClientFasterWhisper(
websocket,
multilingual=options["multilingual"],
multilingual=False,
language=options["language"],
task=options["task"],
client_uid=options["uid"],
@@ -586,13 +586,11 @@ class ServeClientFasterWhisper(ServeClientBase):
"tiny", "tiny.en", "base", "base.en", "small", "small.en",
"medium", "medium.en", "large-v2", "large-v3",
]
self.multilingual = multilingual
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:
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.initial_prompt = initial_prompt
self.vad_parameters = vad_parameters or {"threshold": 0.5}
@@ -602,7 +600,7 @@ class ServeClientFasterWhisper(ServeClientBase):
if self.model_size_or_path == None:
return
self.transcriber = WhisperModel(
self.model_size_or_path,
device=device,
@@ -623,9 +621,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:
self.websocket.send(
@@ -638,15 +642,6 @@ class ServeClientFasterWhisper(ServeClientBase):
)
)
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
def speech_to_text(self):