add model size option to client
This commit is contained in:
+19
-11
@@ -50,7 +50,7 @@ class Client:
|
||||
INSTANCES = {}
|
||||
|
||||
def __init__(
|
||||
self, host=None, port=None, is_multilingual=False, lang=None, translate=False
|
||||
self, host=None, port=None, is_multilingual=False, lang=None, translate=False, model_size="small"
|
||||
):
|
||||
"""
|
||||
Initializes a Client instance for audio recording and streaming to a server.
|
||||
@@ -80,7 +80,9 @@ class Client:
|
||||
self.last_response_recieved = None
|
||||
self.disconnect_if_no_response_for = 15
|
||||
self.multilingual = is_multilingual
|
||||
self.language = lang if is_multilingual else "en"
|
||||
self.language = lang
|
||||
self.model_size = model_size
|
||||
self.server_error = False
|
||||
if translate:
|
||||
self.task = "translate"
|
||||
|
||||
@@ -140,11 +142,16 @@ class Client:
|
||||
print("[ERROR]: invalid client uid")
|
||||
return
|
||||
|
||||
if "status" in message.keys() and message["status"] == "WAIT":
|
||||
self.waiting = True
|
||||
print(
|
||||
f"[INFO]:Server is full. Estimated wait time {round(message['message'])} minutes."
|
||||
)
|
||||
if "status" in message.keys():
|
||||
if message["status"] == "WAIT":
|
||||
self.waiting = True
|
||||
print(
|
||||
f"[INFO]:Server is full. Estimated wait time {round(message['message'])} minutes."
|
||||
)
|
||||
elif message["status"] == "ERROR":
|
||||
print(f"Message from Server: {message['message']}")
|
||||
self.server_error = True
|
||||
return
|
||||
|
||||
if "message" in message.keys() and message["message"] == "DISCONNECT":
|
||||
print("[INFO]: Server overtime disconnected.")
|
||||
@@ -213,6 +220,7 @@ class Client:
|
||||
"multilingual": self.multilingual,
|
||||
"language": self.language,
|
||||
"task": self.task,
|
||||
"model_size": self.model_size,
|
||||
}
|
||||
)
|
||||
)
|
||||
@@ -497,8 +505,8 @@ class TranscriptionClient:
|
||||
transcription_client()
|
||||
```
|
||||
"""
|
||||
def __init__(self, host, port, is_multilingual=False, lang=None, translate=False):
|
||||
self.client = Client(host, port, is_multilingual, lang, translate)
|
||||
def __init__(self, host, port, is_multilingual=False, lang=None, translate=False, model_size="small"):
|
||||
self.client = Client(host, port, is_multilingual, lang, translate, model_size)
|
||||
|
||||
def __call__(self, audio=None, hls_url=None):
|
||||
"""
|
||||
@@ -514,10 +522,10 @@ class TranscriptionClient:
|
||||
"""
|
||||
print("[INFO]: Waiting for server ready ...")
|
||||
while not self.client.recording:
|
||||
if self.client.waiting:
|
||||
if self.client.waiting or self.client.server_error:
|
||||
self.client.close_websocket()
|
||||
return
|
||||
pass
|
||||
|
||||
print("[INFO]: Server Ready!")
|
||||
if hls_url is not None:
|
||||
self.client.process_hls_stream(hls_url)
|
||||
|
||||
Reference in New Issue
Block a user