refactor: include additional parameters
Refactor ServeClientBase and its subclasses to include additional parameters for segment handling and audio clipping.
This commit is contained in:
@@ -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()
|
||||||
|
|||||||
@@ -11,7 +11,20 @@ class ServeClientTensorRT(ServeClientBase):
|
|||||||
SINGLE_MODEL = None
|
SINGLE_MODEL = None
|
||||||
SINGLE_MODEL_LOCK = threading.Lock()
|
SINGLE_MODEL_LOCK = threading.Lock()
|
||||||
|
|
||||||
def __init__(self, websocket, task="transcribe", multilingual=False, language=None, client_uid=None, model=None, single_model=False):
|
def __init__(
|
||||||
|
self,
|
||||||
|
websocket,
|
||||||
|
task="transcribe",
|
||||||
|
multilingual=False,
|
||||||
|
language=None,
|
||||||
|
client_uid=None,
|
||||||
|
model=None,
|
||||||
|
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.
|
||||||
@@ -26,9 +39,19 @@ class ServeClientTensorRT(ServeClientBase):
|
|||||||
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.
|
||||||
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 = language if multilingual else "en"
|
self.language = language if multilingual else "en"
|
||||||
self.task = task
|
self.task = task
|
||||||
self.eos = False
|
self.eos = False
|
||||||
|
|||||||
Reference in New Issue
Block a user