Merge pull request #367 from giubots/configure-more-params
Add possibility to configure more parameters
This commit is contained in:
@@ -51,6 +51,10 @@ class TestClientCallbacks(BaseTestCase):
|
|||||||
"use_vad": True,
|
"use_vad": True,
|
||||||
"max_clients": 4,
|
"max_clients": 4,
|
||||||
"max_connection_time": 600,
|
"max_connection_time": 600,
|
||||||
|
"send_last_n_segments": 10,
|
||||||
|
"no_speech_thresh": 0.45,
|
||||||
|
"clip_audio": False,
|
||||||
|
"same_output_threshold": 10,
|
||||||
})
|
})
|
||||||
self.client.on_open(self.mock_ws_app)
|
self.client.on_open(self.mock_ws_app)
|
||||||
self.mock_ws_app.send.assert_called_with(expected_message)
|
self.mock_ws_app.send.assert_called_with(expected_message)
|
||||||
|
|||||||
@@ -10,22 +10,46 @@ class ServeClientBase(object):
|
|||||||
SERVER_READY = "SERVER_READY"
|
SERVER_READY = "SERVER_READY"
|
||||||
DISCONNECT = "DISCONNECT"
|
DISCONNECT = "DISCONNECT"
|
||||||
|
|
||||||
def __init__(self, client_uid, websocket):
|
client_uid: str
|
||||||
|
"""A unique identifier for the client."""
|
||||||
|
websocket: object
|
||||||
|
"""The WebSocket connection for the client."""
|
||||||
|
send_last_n_segments: int
|
||||||
|
"""Number of most recent segments to send to the client."""
|
||||||
|
no_speech_thresh: float
|
||||||
|
"""Segments with no speech probability above this threshold will be discarded."""
|
||||||
|
clip_audio: bool
|
||||||
|
"""Whether to clip audio with no valid segments."""
|
||||||
|
same_output_threshold: int
|
||||||
|
"""Number of repeated outputs before considering it as a valid segment."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
client_uid,
|
||||||
|
websocket,
|
||||||
|
send_last_n_segments=10,
|
||||||
|
no_speech_thresh=0.45,
|
||||||
|
clip_audio=False,
|
||||||
|
same_output_threshold=10,
|
||||||
|
):
|
||||||
self.client_uid = client_uid
|
self.client_uid = client_uid
|
||||||
self.websocket = websocket
|
self.websocket = websocket
|
||||||
|
self.send_last_n_segments = send_last_n_segments
|
||||||
|
self.no_speech_thresh = no_speech_thresh
|
||||||
|
self.clip_audio = clip_audio
|
||||||
|
self.same_output_threshold = same_output_threshold
|
||||||
|
|
||||||
self.frames = b""
|
self.frames = b""
|
||||||
self.timestamp_offset = 0.0
|
self.timestamp_offset = 0.0
|
||||||
self.frames_np = None
|
self.frames_np = None
|
||||||
self.frames_offset = 0.0
|
self.frames_offset = 0.0
|
||||||
self.text = []
|
self.text = []
|
||||||
self.current_out = ''
|
self.current_out = ""
|
||||||
self.prev_out = ''
|
self.prev_out = ""
|
||||||
self.exit = False
|
self.exit = False
|
||||||
self.same_output_count = 0
|
self.same_output_count = 0
|
||||||
self.transcript = []
|
self.transcript = []
|
||||||
self.send_last_n_segments = 10
|
self.end_time_for_same_output = None
|
||||||
self.no_speech_thresh = 0.45
|
|
||||||
self.clip_audio = False
|
|
||||||
|
|
||||||
# threading
|
# threading
|
||||||
self.lock = threading.Lock()
|
self.lock = threading.Lock()
|
||||||
|
|||||||
@@ -9,12 +9,26 @@ from whisper_live.backend.base import ServeClientBase
|
|||||||
|
|
||||||
|
|
||||||
class ServeClientFasterWhisper(ServeClientBase):
|
class ServeClientFasterWhisper(ServeClientBase):
|
||||||
|
|
||||||
SINGLE_MODEL = None
|
SINGLE_MODEL = None
|
||||||
SINGLE_MODEL_LOCK = threading.Lock()
|
SINGLE_MODEL_LOCK = threading.Lock()
|
||||||
|
|
||||||
def __init__(self, websocket, task="transcribe", device=None, language=None, client_uid=None, model="small.en",
|
def __init__(
|
||||||
initial_prompt=None, vad_parameters=None, use_vad=True, single_model=False):
|
self,
|
||||||
|
websocket,
|
||||||
|
task="transcribe",
|
||||||
|
device=None,
|
||||||
|
language=None,
|
||||||
|
client_uid=None,
|
||||||
|
model="small.en",
|
||||||
|
initial_prompt=None,
|
||||||
|
vad_parameters=None,
|
||||||
|
use_vad=True,
|
||||||
|
single_model=False,
|
||||||
|
send_last_n_segments=10,
|
||||||
|
no_speech_thresh=0.45,
|
||||||
|
clip_audio=False,
|
||||||
|
same_output_threshold=10,
|
||||||
|
):
|
||||||
"""
|
"""
|
||||||
Initialize a ServeClient instance.
|
Initialize a ServeClient instance.
|
||||||
The Whisper model is initialized based on the client's language and device availability.
|
The Whisper model is initialized based on the client's language and device availability.
|
||||||
@@ -23,15 +37,27 @@ class ServeClientFasterWhisper(ServeClientBase):
|
|||||||
|
|
||||||
Args:
|
Args:
|
||||||
websocket (WebSocket): The WebSocket connection for the client.
|
websocket (WebSocket): The WebSocket connection for the client.
|
||||||
task (str, optional): The task type, e.g., "transcribe." Defaults to "transcribe".
|
task (str, optional): The task type, e.g., "transcribe". Defaults to "transcribe".
|
||||||
device (str, optional): The device type for Whisper, "cuda" or "cpu". Defaults to None.
|
device (str, optional): The device type for Whisper, "cuda" or "cpu". Defaults to None.
|
||||||
language (str, optional): The language for transcription. Defaults to None.
|
language (str, optional): The language for transcription. Defaults to None.
|
||||||
client_uid (str, optional): A unique identifier for the client. Defaults to None.
|
client_uid (str, optional): A unique identifier for the client. Defaults to None.
|
||||||
model (str, optional): The whisper model size. Defaults to 'small.en'
|
model (str, optional): The whisper model size. Defaults to 'small.en'
|
||||||
initial_prompt (str, optional): Prompt for whisper inference. Defaults to None.
|
initial_prompt (str, optional): Prompt for whisper inference. Defaults to None.
|
||||||
single_model (bool, optional): Whether to instantiate a new model for each client connection. Defaults to False.
|
single_model (bool, optional): Whether to instantiate a new model for each client connection. Defaults to False.
|
||||||
|
send_last_n_segments (int, optional): Number of most recent segments to send to the client. Defaults to 10.
|
||||||
|
no_speech_thresh (float, optional): Segments with no speech probability above this threshold will be discarded. Defaults to 0.45.
|
||||||
|
clip_audio (bool, optional): Whether to clip audio with no valid segments. Defaults to False.
|
||||||
|
same_output_threshold (int, optional): Number of repeated outputs before considering it as a valid segment. Defaults to 10.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
super().__init__(client_uid, websocket)
|
super().__init__(
|
||||||
|
client_uid,
|
||||||
|
websocket,
|
||||||
|
send_last_n_segments,
|
||||||
|
no_speech_thresh,
|
||||||
|
clip_audio,
|
||||||
|
same_output_threshold,
|
||||||
|
)
|
||||||
self.model_sizes = [
|
self.model_sizes = [
|
||||||
"tiny", "tiny.en", "base", "base.en", "small", "small.en",
|
"tiny", "tiny.en", "base", "base.en", "small", "small.en",
|
||||||
"medium", "medium.en", "large-v2", "large-v3", "distil-small.en",
|
"medium", "medium.en", "large-v2", "large-v3", "distil-small.en",
|
||||||
@@ -45,9 +71,6 @@ class ServeClientFasterWhisper(ServeClientBase):
|
|||||||
self.initial_prompt = initial_prompt
|
self.initial_prompt = initial_prompt
|
||||||
self.vad_parameters = vad_parameters or {"onset": 0.5}
|
self.vad_parameters = vad_parameters or {"onset": 0.5}
|
||||||
|
|
||||||
self.same_output_threshold = 10
|
|
||||||
self.end_time_for_same_output = None
|
|
||||||
|
|
||||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||||
if device == "cuda":
|
if device == "cuda":
|
||||||
major, _ = torch.cuda.get_device_capability(device)
|
major, _ = torch.cuda.get_device_capability(device)
|
||||||
|
|||||||
@@ -12,8 +12,23 @@ class ServeClientOpenVINO(ServeClientBase):
|
|||||||
SINGLE_MODEL = None
|
SINGLE_MODEL = None
|
||||||
SINGLE_MODEL_LOCK = threading.Lock()
|
SINGLE_MODEL_LOCK = threading.Lock()
|
||||||
|
|
||||||
def __init__(self, websocket, task="transcribe", device=None, language=None, client_uid=None, model="small.en",
|
def __init__(
|
||||||
initial_prompt=None, vad_parameters=None, use_vad=True, single_model=False):
|
self,
|
||||||
|
websocket,
|
||||||
|
task="transcribe",
|
||||||
|
device=None,
|
||||||
|
language=None,
|
||||||
|
client_uid=None,
|
||||||
|
model="small.en",
|
||||||
|
initial_prompt=None,
|
||||||
|
vad_parameters=None,
|
||||||
|
use_vad=True,
|
||||||
|
single_model=False,
|
||||||
|
send_last_n_segments=10,
|
||||||
|
no_speech_thresh=0.45,
|
||||||
|
clip_audio=False,
|
||||||
|
same_output_threshold=10,
|
||||||
|
):
|
||||||
"""
|
"""
|
||||||
Initialize a ServeClient instance.
|
Initialize a ServeClient instance.
|
||||||
The Whisper model is initialized based on the client's language and device availability.
|
The Whisper model is initialized based on the client's language and device availability.
|
||||||
@@ -29,15 +44,25 @@ class ServeClientOpenVINO(ServeClientBase):
|
|||||||
model (str, optional): Huggingface model_id for a valid OpenVINO model.
|
model (str, optional): Huggingface model_id for a valid OpenVINO model.
|
||||||
initial_prompt (str, optional): Prompt for whisper inference. Defaults to None.
|
initial_prompt (str, optional): Prompt for whisper inference. Defaults to None.
|
||||||
single_model (bool, optional): Whether to instantiate a new model for each client connection. Defaults to False.
|
single_model (bool, optional): Whether to instantiate a new model for each client connection. Defaults to False.
|
||||||
|
send_last_n_segments (int, optional): Number of most recent segments to send to the client. Defaults to 10.
|
||||||
|
no_speech_thresh (float, optional): Segments with no speech probability above this threshold will be discarded. Defaults to 0.45.
|
||||||
|
clip_audio (bool, optional): Whether to clip audio with no valid segments. Defaults to False.
|
||||||
|
same_output_threshold (int, optional): Number of repeated outputs before considering it as a valid segment. Defaults to 10.
|
||||||
"""
|
"""
|
||||||
super().__init__(client_uid, websocket)
|
super().__init__(
|
||||||
|
client_uid,
|
||||||
|
websocket,
|
||||||
|
send_last_n_segments,
|
||||||
|
no_speech_thresh,
|
||||||
|
clip_audio,
|
||||||
|
same_output_threshold,
|
||||||
|
)
|
||||||
self.language = "en" if language is None else language
|
self.language = "en" if language is None else language
|
||||||
if not self.language.startswith("<|"):
|
if not self.language.startswith("<|"):
|
||||||
self.language = f"<|{self.language}|>"
|
self.language = f"<|{self.language}|>"
|
||||||
|
|
||||||
self.task = "transcribe" if task is None else task
|
self.task = "transcribe" if task is None else task
|
||||||
self.same_output_threshold = 10
|
|
||||||
self.end_time_for_same_output = None
|
|
||||||
self.clip_audio = True
|
self.clip_audio = True
|
||||||
|
|
||||||
core = Core()
|
core = Core()
|
||||||
|
|||||||
@@ -22,6 +22,10 @@ class ServeClientTensorRT(ServeClientBase):
|
|||||||
single_model=False,
|
single_model=False,
|
||||||
use_py_session=False,
|
use_py_session=False,
|
||||||
max_new_tokens=225,
|
max_new_tokens=225,
|
||||||
|
send_last_n_segments=10,
|
||||||
|
no_speech_thresh=0.45,
|
||||||
|
clip_audio=False,
|
||||||
|
same_output_threshold=10,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Initialize a ServeClient instance.
|
Initialize a ServeClient instance.
|
||||||
@@ -39,9 +43,20 @@ class ServeClientTensorRT(ServeClientBase):
|
|||||||
single_model (bool, optional): Whether to instantiate a new model for each client connection. Defaults to False.
|
single_model (bool, optional): Whether to instantiate a new model for each client connection. Defaults to False.
|
||||||
use_py_session (bool, optional): Use python session or cpp session. Defaults to Cpp Session.
|
use_py_session (bool, optional): Use python session or cpp session. Defaults to Cpp Session.
|
||||||
max_new_tokens (int, optional): Max number of tokens to generate.
|
max_new_tokens (int, optional): Max number of tokens to generate.
|
||||||
|
send_last_n_segments (int, optional): Number of most recent segments to send to the client. Defaults to 10.
|
||||||
|
no_speech_thresh (float, optional): Segments with no speech probability above this threshold will be discarded. Defaults to 0.45.
|
||||||
|
clip_audio (bool, optional): Whether to clip audio with no valid segments. Defaults to False.
|
||||||
|
same_output_threshold (int, optional): Number of repeated outputs before considering it as a valid segment. Defaults to 10.
|
||||||
"""
|
"""
|
||||||
super().__init__(client_uid, websocket)
|
super().__init__(
|
||||||
|
client_uid,
|
||||||
|
websocket,
|
||||||
|
send_last_n_segments,
|
||||||
|
no_speech_thresh,
|
||||||
|
clip_audio,
|
||||||
|
same_output_threshold,
|
||||||
|
)
|
||||||
|
|
||||||
self.language = language if multilingual else "en"
|
self.language = language if multilingual else "en"
|
||||||
self.task = task
|
self.task = task
|
||||||
self.eos = False
|
self.eos = False
|
||||||
|
|||||||
+38
-3
@@ -33,6 +33,10 @@ class Client:
|
|||||||
log_transcription=True,
|
log_transcription=True,
|
||||||
max_clients=4,
|
max_clients=4,
|
||||||
max_connection_time=600,
|
max_connection_time=600,
|
||||||
|
send_last_n_segments=10,
|
||||||
|
no_speech_thresh=0.45,
|
||||||
|
clip_audio=False,
|
||||||
|
same_output_threshold=10,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
Initializes a Client instance for audio recording and streaming to a server.
|
Initializes a Client instance for audio recording and streaming to a server.
|
||||||
@@ -52,6 +56,10 @@ class Client:
|
|||||||
log_transcription (bool, optional): Whether to log transcription output to the console. Default is True.
|
log_transcription (bool, optional): Whether to log transcription output to the console. Default is True.
|
||||||
max_clients (int, optional): Maximum number of client connections allowed. Default is 4.
|
max_clients (int, optional): Maximum number of client connections allowed. Default is 4.
|
||||||
max_connection_time (int, optional): Maximum allowed connection time in seconds. Default is 600.
|
max_connection_time (int, optional): Maximum allowed connection time in seconds. Default is 600.
|
||||||
|
send_last_n_segments (int, optional): Number of most recent segments to send to the client. Defaults to 10.
|
||||||
|
no_speech_thresh (float, optional): Segments with no speech probability above this threshold will be discarded. Defaults to 0.45.
|
||||||
|
clip_audio (bool, optional): Whether to clip audio with no valid segments. Defaults to False.
|
||||||
|
same_output_threshold (int, optional): Number of repeated outputs before considering it as a valid segment. Defaults to 10.
|
||||||
"""
|
"""
|
||||||
self.recording = False
|
self.recording = False
|
||||||
self.task = "transcribe"
|
self.task = "transcribe"
|
||||||
@@ -69,6 +77,10 @@ class Client:
|
|||||||
self.log_transcription = log_transcription
|
self.log_transcription = log_transcription
|
||||||
self.max_clients = max_clients
|
self.max_clients = max_clients
|
||||||
self.max_connection_time = max_connection_time
|
self.max_connection_time = max_connection_time
|
||||||
|
self.send_last_n_segments = send_last_n_segments
|
||||||
|
self.no_speech_thresh = no_speech_thresh
|
||||||
|
self.clip_audio = clip_audio
|
||||||
|
self.same_output_threshold = same_output_threshold
|
||||||
|
|
||||||
if translate:
|
if translate:
|
||||||
self.task = "translate"
|
self.task = "translate"
|
||||||
@@ -212,6 +224,10 @@ class Client:
|
|||||||
"use_vad": self.use_vad,
|
"use_vad": self.use_vad,
|
||||||
"max_clients": self.max_clients,
|
"max_clients": self.max_clients,
|
||||||
"max_connection_time": self.max_connection_time,
|
"max_connection_time": self.max_connection_time,
|
||||||
|
"send_last_n_segments": self.send_last_n_segments,
|
||||||
|
"no_speech_thresh": self.no_speech_thresh,
|
||||||
|
"clip_audio": self.clip_audio,
|
||||||
|
"same_output_threshold": self.same_output_threshold,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
@@ -682,6 +698,10 @@ class TranscriptionClient(TranscriptionTeeClient):
|
|||||||
max_clients (int, optional): Maximum number of client connections allowed. Default is 4.
|
max_clients (int, optional): Maximum number of client connections allowed. Default is 4.
|
||||||
max_connection_time (int, optional): Maximum allowed connection time in seconds. Default is 600.
|
max_connection_time (int, optional): Maximum allowed connection time in seconds. Default is 600.
|
||||||
mute_audio_playback (bool, optional): If True, mutes audio playback during file playback. Default is False.
|
mute_audio_playback (bool, optional): If True, mutes audio playback during file playback. Default is False.
|
||||||
|
send_last_n_segments (int, optional): Number of most recent segments to send to the client. Defaults to 10.
|
||||||
|
no_speech_thresh (float, optional): Segments with no speech probability above this threshold will be discarded. Defaults to 0.45.
|
||||||
|
clip_audio (bool, optional): Whether to clip audio with no valid segments. Defaults to False.
|
||||||
|
same_output_threshold (int, optional): Number of repeated outputs before considering it as a valid segment. Defaults to 10.
|
||||||
|
|
||||||
Attributes:
|
Attributes:
|
||||||
client (Client): An instance of the underlying Client class responsible for handling the WebSocket connection.
|
client (Client): An instance of the underlying Client class responsible for handling the WebSocket connection.
|
||||||
@@ -708,11 +728,26 @@ class TranscriptionClient(TranscriptionTeeClient):
|
|||||||
max_clients=4,
|
max_clients=4,
|
||||||
max_connection_time=600,
|
max_connection_time=600,
|
||||||
mute_audio_playback=False,
|
mute_audio_playback=False,
|
||||||
|
send_last_n_segments=10,
|
||||||
|
no_speech_thresh=0.45,
|
||||||
|
clip_audio=False,
|
||||||
|
same_output_threshold=10,
|
||||||
):
|
):
|
||||||
self.client = Client(
|
self.client = Client(
|
||||||
host, port, lang, translate, model, srt_file_path=output_transcription_path,
|
host,
|
||||||
use_vad=use_vad, log_transcription=log_transcription, max_clients=max_clients,
|
port,
|
||||||
max_connection_time=max_connection_time
|
lang,
|
||||||
|
translate,
|
||||||
|
model,
|
||||||
|
srt_file_path=output_transcription_path,
|
||||||
|
use_vad=use_vad,
|
||||||
|
log_transcription=log_transcription,
|
||||||
|
max_clients=max_clients,
|
||||||
|
max_connection_time=max_connection_time,
|
||||||
|
send_last_n_segments=send_last_n_segments,
|
||||||
|
no_speech_thresh=no_speech_thresh,
|
||||||
|
clip_audio=clip_audio,
|
||||||
|
same_output_threshold=same_output_threshold,
|
||||||
)
|
)
|
||||||
|
|
||||||
if save_output_recording and not output_recording_filename.endswith(".wav"):
|
if save_output_recording and not output_recording_filename.endswith(".wav"):
|
||||||
|
|||||||
@@ -169,6 +169,10 @@ class TranscriptionServer:
|
|||||||
model=whisper_tensorrt_path,
|
model=whisper_tensorrt_path,
|
||||||
single_model=self.single_model,
|
single_model=self.single_model,
|
||||||
use_py_session=trt_py_session,
|
use_py_session=trt_py_session,
|
||||||
|
send_last_n_segments=options.get("send_last_n_segments", 10),
|
||||||
|
no_speech_thresh=options.get("no_speech_thresh", 0.45),
|
||||||
|
clip_audio=options.get("clip_audio", False),
|
||||||
|
same_output_threshold=options.get("same_output_threshold", 10),
|
||||||
)
|
)
|
||||||
logging.info("Running TensorRT backend.")
|
logging.info("Running TensorRT backend.")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -192,6 +196,10 @@ class TranscriptionServer:
|
|||||||
client_uid=options["uid"],
|
client_uid=options["uid"],
|
||||||
model=options["model"],
|
model=options["model"],
|
||||||
single_model=self.single_model,
|
single_model=self.single_model,
|
||||||
|
send_last_n_segments=options.get("send_last_n_segments", 10),
|
||||||
|
no_speech_thresh=options.get("no_speech_thresh", 0.45),
|
||||||
|
clip_audio=options.get("clip_audio", False),
|
||||||
|
same_output_threshold=options.get("same_output_threshold", 10),
|
||||||
)
|
)
|
||||||
logging.info("Running OpenVINO backend.")
|
logging.info("Running OpenVINO backend.")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -221,6 +229,10 @@ class TranscriptionServer:
|
|||||||
vad_parameters=options.get("vad_parameters"),
|
vad_parameters=options.get("vad_parameters"),
|
||||||
use_vad=self.use_vad,
|
use_vad=self.use_vad,
|
||||||
single_model=self.single_model,
|
single_model=self.single_model,
|
||||||
|
send_last_n_segments=options.get("send_last_n_segments", 10),
|
||||||
|
no_speech_thresh=options.get("no_speech_thresh", 0.45),
|
||||||
|
clip_audio=options.get("clip_audio", False),
|
||||||
|
same_output_threshold=options.get("same_output_threshold", 10),
|
||||||
)
|
)
|
||||||
|
|
||||||
logging.info("Running faster_whisper backend.")
|
logging.info("Running faster_whisper backend.")
|
||||||
|
|||||||
Reference in New Issue
Block a user