Merge pull request #102 from makaveli10/change_model_size_param_name

Server to control custom model usage.
This commit is contained in:
Marcus Edel
2024-01-18 11:07:18 -05:00
committed by GitHub
2 changed files with 24 additions and 27 deletions
+5 -9
View File
@@ -56,8 +56,7 @@ class Client:
is_multilingual=False, is_multilingual=False,
lang=None, lang=None,
translate=False, translate=False,
model_size="small", model="small",
use_custom_model=False
): ):
""" """
Initializes a Client instance for audio recording and streaming to a server. Initializes a Client instance for audio recording and streaming to a server.
@@ -88,9 +87,8 @@ class Client:
self.disconnect_if_no_response_for = 15 self.disconnect_if_no_response_for = 15
self.multilingual = is_multilingual self.multilingual = is_multilingual
self.language = lang self.language = lang
self.model_size = model_size self.model = model
self.server_error = False self.server_error = False
self.use_custom_model = use_custom_model
if translate: if translate:
self.task = "translate" self.task = "translate"
@@ -229,8 +227,7 @@ class Client:
"multilingual": self.multilingual, "multilingual": self.multilingual,
"language": self.language, "language": self.language,
"task": self.task, "task": self.task,
"model_size": self.model_size, "model": self.model,
"use_custom_model": self.use_custom_model # if runnning your own server with a custom model
} }
) )
) )
@@ -521,10 +518,9 @@ class TranscriptionClient:
is_multilingual=False, is_multilingual=False,
lang=None, lang=None,
translate=False, translate=False,
model_size="small", model="small",
use_custom_model=False
): ):
self.client = Client(host, port, is_multilingual, lang, translate, model_size, use_custom_model) self.client = Client(host, port, is_multilingual, lang, translate, model)
def __call__(self, audio=None, hls_url=None): def __call__(self, audio=None, hls_url=None):
""" """
+19 -18
View File
@@ -6,7 +6,7 @@ import json
import textwrap import textwrap
import logging import logging
# logging.basicConfig(level = logging.INFO) logging.basicConfig(level = logging.INFO)
from websockets.sync.server import serve from websockets.sync.server import serve
@@ -100,9 +100,9 @@ class TranscriptionServer:
return return
# validate custom model # validate custom model
if options["use_custom_model"]: if custom_model_path is not None and os.path.exists(custom_model_path):
if custom_model_path is None or not os.path.exists(custom_model_path): logging.info(f"Using custom model {custom_model_path}")
options["use_custom_model"] = False options["model"] = custom_model_path
client = ServeClient( client = ServeClient(
websocket, websocket,
@@ -110,10 +110,9 @@ class TranscriptionServer:
language=options["language"], language=options["language"],
task=options["task"], task=options["task"],
client_uid=options["uid"], client_uid=options["uid"],
model_size_or_path=custom_model_path if options["use_custom_model"] else options["model_size"], model=options["model"],
initial_prompt=options.get("initial_prompt"), initial_prompt=options.get("initial_prompt"),
vad_parameters=options.get("vad_parameters"), vad_parameters=options.get("vad_parameters"),
use_custom_model=options["use_custom_model"]
) )
self.clients[websocket] = client self.clients[websocket] = client
@@ -206,10 +205,9 @@ class ServeClient:
multilingual=False, multilingual=False,
language=None, language=None,
client_uid=None, client_uid=None,
model_size_or_path="small", model="small",
initial_prompt=None, initial_prompt=None,
vad_parameters=None, vad_parameters=None,
use_custom_model=False
): ):
""" """
Initialize a ServeClient instance. Initialize a ServeClient instance.
@@ -230,13 +228,15 @@ class ServeClient:
self.data = b"" self.data = b""
self.frames = b"" self.frames = b""
self.model_sizes = [ self.model_sizes = [
"tiny", "base", "small", "medium", "large-v2", "large-v3" "tiny", "tiny.en", "base", "base.en", "small", "small.en",
"medium", "medium.en", "large-v2", "large-v3",
] ]
self.multilingual = multilingual self.multilingual = multilingual
if not use_custom_model: if not os.path.exists(model):
self.model_size_or_path = self.get_model_size(model_size_or_path) self.model_size_or_path = self.get_model_size(model)
else: else:
self.model_size_or_path = model_size_or_path self.model_size_or_path = model
self.language = language if self.multilingual else "en" self.language = language if self.multilingual else "en"
self.task = task self.task = task
@@ -246,7 +246,7 @@ class ServeClient:
device = "cuda" if torch.cuda.is_available() else "cpu" device = "cuda" if torch.cuda.is_available() else "cpu"
if self.model_size_or_path == None: if self.model_size_or_path is None:
return return
self.transcriber = WhisperModel( self.transcriber = WhisperModel(
@@ -302,12 +302,13 @@ class ServeClient:
) )
return None return None
if model_size in ["large-v2", "large-v3"]: 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 self.multilingual = True
return model_size
if not self.multilingual:
model_size = model_size + ".en"
return model_size return model_size