Merge branch 'main' of github.com:lightwastak3n/WhisperLive
This commit is contained in:
+12
-17
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user