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