Merge pull request #102 from makaveli10/change_model_size_param_name
Server to control custom model usage.
This commit is contained in:
@@ -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
@@ -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:
|
||||||
self.multilingual = True
|
logging.info(f"Setting multilingual to false with {model_size} which is english only model.")
|
||||||
return model_size
|
self.multilingual = False
|
||||||
|
|
||||||
if not self.multilingual:
|
if not model_size.endswith("en") and not self.multilingual:
|
||||||
model_size = model_size + ".en"
|
logging.info(f"Setting multilingual to true with multilingual model {model_size}.")
|
||||||
|
self.multilingual = True
|
||||||
|
|
||||||
return model_size
|
return model_size
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user