add option to use custom model
This commit is contained in:
+9
-1
@@ -1,5 +1,13 @@
|
|||||||
|
import argparse
|
||||||
from whisper_live.server import TranscriptionServer
|
from whisper_live.server import TranscriptionServer
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
server = TranscriptionServer()
|
server = TranscriptionServer()
|
||||||
server.run("0.0.0.0")
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument('--model_path', type=str, default=None, help="Custom Faster Whisper Model")
|
||||||
|
args = parser.parse_args()
|
||||||
|
server.run(
|
||||||
|
"0.0.0.0",
|
||||||
|
9090,
|
||||||
|
custom_model_path=args.model_path
|
||||||
|
)
|
||||||
|
|||||||
+21
-3
@@ -50,7 +50,14 @@ class Client:
|
|||||||
INSTANCES = {}
|
INSTANCES = {}
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self, host=None, port=None, is_multilingual=False, lang=None, translate=False, model_size="small"
|
self,
|
||||||
|
host=None,
|
||||||
|
port=None,
|
||||||
|
is_multilingual=False,
|
||||||
|
lang=None,
|
||||||
|
translate=False,
|
||||||
|
model_size="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.
|
||||||
@@ -83,6 +90,8 @@ class Client:
|
|||||||
self.language = lang
|
self.language = lang
|
||||||
self.model_size = model_size
|
self.model_size = model_size
|
||||||
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"
|
||||||
|
|
||||||
@@ -221,6 +230,7 @@ class Client:
|
|||||||
"language": self.language,
|
"language": self.language,
|
||||||
"task": self.task,
|
"task": self.task,
|
||||||
"model_size": self.model_size,
|
"model_size": self.model_size,
|
||||||
|
"use_custom_model": self.use_custom_model # if runnning your own server with a custom model
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
@@ -505,8 +515,16 @@ class TranscriptionClient:
|
|||||||
transcription_client()
|
transcription_client()
|
||||||
```
|
```
|
||||||
"""
|
"""
|
||||||
def __init__(self, host, port, is_multilingual=False, lang=None, translate=False, model_size="small"):
|
def __init__(self,
|
||||||
self.client = Client(host, port, is_multilingual, lang, translate, model_size)
|
host,
|
||||||
|
port,
|
||||||
|
is_multilingual=False,
|
||||||
|
lang=None,
|
||||||
|
translate=False,
|
||||||
|
model_size="small",
|
||||||
|
use_custom_model=False
|
||||||
|
):
|
||||||
|
self.client = Client(host, port, is_multilingual, lang, translate, model_size, use_custom_model)
|
||||||
|
|
||||||
def __call__(self, audio=None, hls_url=None):
|
def __call__(self, audio=None, hls_url=None):
|
||||||
"""
|
"""
|
||||||
|
|||||||
+31
-10
@@ -1,3 +1,4 @@
|
|||||||
|
import os
|
||||||
import websockets
|
import websockets
|
||||||
import time
|
import time
|
||||||
import threading
|
import threading
|
||||||
@@ -12,6 +13,8 @@ from websockets.sync.server import serve
|
|||||||
import torch
|
import torch
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import time
|
import time
|
||||||
|
import functools
|
||||||
|
|
||||||
from whisper_live.transcriber import WhisperModel
|
from whisper_live.transcriber import WhisperModel
|
||||||
|
|
||||||
|
|
||||||
@@ -58,7 +61,7 @@ class TranscriptionServer:
|
|||||||
|
|
||||||
return wait_time / 60
|
return wait_time / 60
|
||||||
|
|
||||||
def recv_audio(self, websocket):
|
def recv_audio(self, websocket, custom_model_path=None):
|
||||||
"""
|
"""
|
||||||
Receive audio chunks from a client in an infinite loop.
|
Receive audio chunks from a client in an infinite loop.
|
||||||
|
|
||||||
@@ -95,6 +98,11 @@ class TranscriptionServer:
|
|||||||
websocket.close()
|
websocket.close()
|
||||||
del websocket
|
del websocket
|
||||||
return
|
return
|
||||||
|
|
||||||
|
# validate custom model
|
||||||
|
if options["use_custom_model"]:
|
||||||
|
if custom_model_path is None or not os.path.exists(custom_model_path):
|
||||||
|
options["use_custom_model"] = False
|
||||||
|
|
||||||
client = ServeClient(
|
client = ServeClient(
|
||||||
websocket,
|
websocket,
|
||||||
@@ -102,9 +110,10 @@ class TranscriptionServer:
|
|||||||
language=options["language"],
|
language=options["language"],
|
||||||
task=options["task"],
|
task=options["task"],
|
||||||
client_uid=options["uid"],
|
client_uid=options["uid"],
|
||||||
model_size=options["model_size"],
|
model_size_or_path=custom_model_path if options["use_custom_model"] else options["model_size"],
|
||||||
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
|
||||||
@@ -137,7 +146,7 @@ class TranscriptionServer:
|
|||||||
del websocket
|
del websocket
|
||||||
break
|
break
|
||||||
|
|
||||||
def run(self, host, port=9090):
|
def run(self, host, port=9090, custom_model_path=None):
|
||||||
"""
|
"""
|
||||||
Run the transcription server.
|
Run the transcription server.
|
||||||
|
|
||||||
@@ -145,7 +154,14 @@ class TranscriptionServer:
|
|||||||
host (str): The host address to bind the server.
|
host (str): The host address to bind the server.
|
||||||
port (int): The port number to bind the server.
|
port (int): The port number to bind the server.
|
||||||
"""
|
"""
|
||||||
with serve(self.recv_audio, host, port) as server:
|
with serve(
|
||||||
|
functools.partial(
|
||||||
|
self.recv_audio,
|
||||||
|
custom_model_path=custom_model_path
|
||||||
|
),
|
||||||
|
host,
|
||||||
|
port
|
||||||
|
) as server:
|
||||||
server.serve_forever()
|
server.serve_forever()
|
||||||
|
|
||||||
|
|
||||||
@@ -190,9 +206,10 @@ class ServeClient:
|
|||||||
multilingual=False,
|
multilingual=False,
|
||||||
language=None,
|
language=None,
|
||||||
client_uid=None,
|
client_uid=None,
|
||||||
model_size="small",
|
model_size_or_path="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.
|
||||||
@@ -216,7 +233,11 @@ class ServeClient:
|
|||||||
"tiny", "base", "small", "medium", "large-v2", "large-v3"
|
"tiny", "base", "small", "medium", "large-v2", "large-v3"
|
||||||
]
|
]
|
||||||
self.multilingual = multilingual
|
self.multilingual = multilingual
|
||||||
self.model_size = self.get_model_size(model_size)
|
if not use_custom_model:
|
||||||
|
self.model_size_or_path = self.get_model_size(model_size_or_path)
|
||||||
|
else:
|
||||||
|
self.model_size_or_path = model_size_or_path
|
||||||
|
|
||||||
self.language = language if self.multilingual else "en"
|
self.language = language if self.multilingual else "en"
|
||||||
self.task = task
|
self.task = task
|
||||||
self.websocket = websocket
|
self.websocket = websocket
|
||||||
@@ -225,11 +246,11 @@ class ServeClient:
|
|||||||
|
|
||||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||||
|
|
||||||
if self.model_size == None:
|
if self.model_size_or_path == None:
|
||||||
return
|
return
|
||||||
|
|
||||||
self.transcriber = WhisperModel(
|
self.transcriber = WhisperModel(
|
||||||
self.model_size,
|
self.model_size_or_path,
|
||||||
device=device,
|
device=device,
|
||||||
compute_type="int8" if device=="cpu" else "float16",
|
compute_type="int8" if device=="cpu" else "float16",
|
||||||
local_files_only=False,
|
local_files_only=False,
|
||||||
|
|||||||
Reference in New Issue
Block a user