add model size option to client

This commit is contained in:
makaveli10
2023-12-20 18:06:02 +05:30
parent a52dc0cbf8
commit e006722da7
+19 -11
View File
@@ -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)