From c1ac71ada040582aa25f6c9d819b24299c62f992 Mon Sep 17 00:00:00 2001 From: makaveli10 Date: Mon, 24 Mar 2025 16:48:33 +0530 Subject: [PATCH] Refactor :hammer: Signed-off-by: makaveli10 --- tests/test_vad.py | 2 +- whisper_live/backend/__init__.py | 0 whisper_live/backend/base.py | 188 +++++ .../backend/faster_whisper_backend.py | 381 +++++++++ whisper_live/backend/trt_backend.py | 181 +++++ whisper_live/server.py | 766 +----------------- whisper_live/transcriber/__init__.py | 0 .../{ => transcriber}/tensorrt_utils.py | 0 .../transcriber_faster_whisper.py} | 0 .../{ => transcriber}/transcriber_tensorrt.py | 7 +- 10 files changed, 788 insertions(+), 737 deletions(-) create mode 100644 whisper_live/backend/__init__.py create mode 100644 whisper_live/backend/base.py create mode 100644 whisper_live/backend/faster_whisper_backend.py create mode 100644 whisper_live/backend/trt_backend.py create mode 100644 whisper_live/transcriber/__init__.py rename whisper_live/{ => transcriber}/tensorrt_utils.py (100%) rename whisper_live/{transcriber.py => transcriber/transcriber_faster_whisper.py} (100%) rename whisper_live/{ => transcriber}/transcriber_tensorrt.py (99%) diff --git a/tests/test_vad.py b/tests/test_vad.py index cfc2d3a..0d63f33 100644 --- a/tests/test_vad.py +++ b/tests/test_vad.py @@ -1,6 +1,6 @@ import unittest import numpy as np -from whisper_live.tensorrt_utils import load_audio +from whisper_live.transcriber.tensorrt_utils import load_audio from whisper_live.vad import VoiceActivityDetector diff --git a/whisper_live/backend/__init__.py b/whisper_live/backend/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/whisper_live/backend/base.py b/whisper_live/backend/base.py new file mode 100644 index 0000000..bb09fe5 --- /dev/null +++ b/whisper_live/backend/base.py @@ -0,0 +1,188 @@ +import json +import logging +import threading +import numpy as np + + +class ServeClientBase(object): + RATE = 16000 + SERVER_READY = "SERVER_READY" + DISCONNECT = "DISCONNECT" + + def __init__(self, client_uid, websocket): + self.client_uid = client_uid + self.websocket = websocket + self.frames = b"" + self.timestamp_offset = 0.0 + self.frames_np = None + self.frames_offset = 0.0 + self.text = [] + self.current_out = '' + self.prev_out = '' + self.t_start = None + self.exit = False + self.same_output_count = 0 + self.show_prev_out_thresh = 5 # if pause(no output from whisper) show previous output for 5 seconds + self.add_pause_thresh = 3 # add a blank to segment list as a pause(no speech) for 3 seconds + self.transcript = [] + self.send_last_n_segments = 10 + + # text formatting + self.pick_previous_segments = 2 + + # threading + self.lock = threading.Lock() + + def speech_to_text(self): + raise NotImplementedError + + def transcribe_audio(self): + raise NotImplementedError + + def handle_transcription_output(self): + raise NotImplementedError + + def add_frames(self, frame_np): + """ + Add audio frames to the ongoing audio stream buffer. + + This method is responsible for maintaining the audio stream buffer, allowing the continuous addition + of audio frames as they are received. It also ensures that the buffer does not exceed a specified size + to prevent excessive memory usage. + + If the buffer size exceeds a threshold (45 seconds of audio data), it discards the oldest 30 seconds + of audio data to maintain a reasonable buffer size. If the buffer is empty, it initializes it with the provided + audio frame. The audio stream buffer is used for real-time processing of audio data for transcription. + + Args: + frame_np (numpy.ndarray): The audio frame data as a NumPy array. + + """ + self.lock.acquire() + if self.frames_np is not None and self.frames_np.shape[0] > 45*self.RATE: + self.frames_offset += 30.0 + self.frames_np = self.frames_np[int(30*self.RATE):] + # check timestamp offset(should be >= self.frame_offset) + # this basically means that there is no speech as timestamp offset hasnt updated + # and is less than frame_offset + if self.timestamp_offset < self.frames_offset: + self.timestamp_offset = self.frames_offset + if self.frames_np is None: + self.frames_np = frame_np.copy() + else: + self.frames_np = np.concatenate((self.frames_np, frame_np), axis=0) + self.lock.release() + + def clip_audio_if_no_valid_segment(self): + """ + Update the timestamp offset based on audio buffer status. + Clip audio if the current chunk exceeds 30 seconds, this basically implies that + no valid segment for the last 30 seconds from whisper + """ + with self.lock: + if self.frames_np[int((self.timestamp_offset - self.frames_offset)*self.RATE):].shape[0] > 25 * self.RATE: + duration = self.frames_np.shape[0] / self.RATE + self.timestamp_offset = self.frames_offset + duration - 5 + + def get_audio_chunk_for_processing(self): + """ + Retrieves the next chunk of audio data for processing based on the current offsets. + + Calculates which part of the audio data should be processed next, based on + the difference between the current timestamp offset and the frame's offset, scaled by + the audio sample rate (RATE). It then returns this chunk of audio data along with its + duration in seconds. + + Returns: + tuple: A tuple containing: + - input_bytes (np.ndarray): The next chunk of audio data to be processed. + - duration (float): The duration of the audio chunk in seconds. + """ + with self.lock: + samples_take = max(0, (self.timestamp_offset - self.frames_offset) * self.RATE) + input_bytes = self.frames_np[int(samples_take):].copy() + duration = input_bytes.shape[0] / self.RATE + return input_bytes, duration + + def prepare_segments(self, last_segment=None): + """ + Prepares the segments of transcribed text to be sent to the client. + + This method compiles the recent segments of transcribed text, ensuring that only the + specified number of the most recent segments are included. It also appends the most + recent segment of text if provided (which is considered incomplete because of the possibility + of the last word being truncated in the audio chunk). + + Args: + last_segment (str, optional): The most recent segment of transcribed text to be added + to the list of segments. Defaults to None. + + Returns: + list: A list of transcribed text segments to be sent to the client. + """ + segments = [] + if len(self.transcript) >= self.send_last_n_segments: + segments = self.transcript[-self.send_last_n_segments:].copy() + else: + segments = self.transcript.copy() + if last_segment is not None: + segments = segments + [last_segment] + return segments + + def get_audio_chunk_duration(self, input_bytes): + """ + Calculates the duration of the provided audio chunk. + + Args: + input_bytes (numpy.ndarray): The audio chunk for which to calculate the duration. + + Returns: + float: The duration of the audio chunk in seconds. + """ + return input_bytes.shape[0] / self.RATE + + def send_transcription_to_client(self, segments): + """ + Sends the specified transcription segments to the client over the websocket connection. + + This method formats the transcription segments into a JSON object and attempts to send + this object to the client. If an error occurs during the send operation, it logs the error. + + Returns: + segments (list): A list of transcription segments to be sent to the client. + """ + try: + self.websocket.send( + json.dumps({ + "uid": self.client_uid, + "segments": segments, + }) + ) + except Exception as e: + logging.error(f"[ERROR]: Sending data to client: {e}") + + def disconnect(self): + """ + Notify the client of disconnection and send a disconnect message. + + This method sends a disconnect message to the client via the WebSocket connection to notify them + that the transcription service is disconnecting gracefully. + + """ + self.websocket.send(json.dumps({ + "uid": self.client_uid, + "message": self.DISCONNECT + })) + + def cleanup(self): + """ + Perform cleanup tasks before exiting the transcription service. + + This method performs necessary cleanup tasks, including stopping the transcription thread, marking + the exit flag to indicate the transcription thread should exit gracefully, and destroying resources + associated with the transcription process. + + """ + logging.info("Cleaning up.") + self.exit = True + diff --git a/whisper_live/backend/faster_whisper_backend.py b/whisper_live/backend/faster_whisper_backend.py new file mode 100644 index 0000000..a00ea84 --- /dev/null +++ b/whisper_live/backend/faster_whisper_backend.py @@ -0,0 +1,381 @@ +import json +import logging +import threading +import time +import torch + +from whisper_live.transcriber.transcriber_faster_whisper import WhisperModel +from whisper_live.backend.base import ServeClientBase + + +class ServeClientFasterWhisper(ServeClientBase): + + SINGLE_MODEL = None + SINGLE_MODEL_LOCK = threading.Lock() + + def __init__(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): + """ + Initialize a ServeClient instance. + The Whisper model is initialized based on the client's language and device availability. + The transcription thread is started upon initialization. A "SERVER_READY" message is sent + to the client to indicate that the server is ready. + + Args: + websocket (WebSocket): The WebSocket connection for the client. + 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. + language (str, optional): The language for transcription. 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' + 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. + """ + super().__init__(client_uid, websocket) + self.model_sizes = [ + "tiny", "tiny.en", "base", "base.en", "small", "small.en", + "medium", "medium.en", "large-v2", "large-v3", "distil-small.en", + "distil-medium.en", "distil-large-v2", "distil-large-v3", + "large-v3-turbo", "turbo" + ] + + self.model_size_or_path = model + self.language = "en" if self.model_size_or_path.endswith("en") else language + self.task = task + self.initial_prompt = initial_prompt + self.vad_parameters = vad_parameters or {"onset": 0.5} + self.no_speech_thresh = 0.45 + self.same_output_threshold = 10 + self.end_time_for_same_output = None + + device = "cuda" if torch.cuda.is_available() else "cpu" + if device == "cuda": + major, _ = torch.cuda.get_device_capability(device) + self.compute_type = "float16" if major >= 7 else "float32" + else: + self.compute_type = "int8" + + if self.model_size_or_path is None: + return + logging.info(f"Using Device={device} with precision {self.compute_type}") + + try: + if single_model: + if ServeClientFasterWhisper.SINGLE_MODEL is None: + self.create_model(device) + ServeClientFasterWhisper.SINGLE_MODEL = self.transcriber + else: + self.transcriber = ServeClientFasterWhisper.SINGLE_MODEL + else: + self.create_model(device) + except Exception as e: + logging.error(f"Failed to load model: {e}") + self.websocket.send(json.dumps({ + "uid": self.client_uid, + "status": "ERROR", + "message": f"Failed to load model: {str(self.model_size_or_path)}" + })) + self.websocket.close() + return + + self.use_vad = use_vad + + # threading + self.trans_thread = threading.Thread(target=self.speech_to_text) + self.trans_thread.start() + self.websocket.send( + json.dumps( + { + "uid": self.client_uid, + "message": self.SERVER_READY, + "backend": "faster_whisper" + } + ) + ) + + def create_model(self, device): + """ + Instantiates a new model, sets it as the transcriber. + """ + self.transcriber = WhisperModel( + self.model_size_or_path, + device=device, + compute_type=self.compute_type, + local_files_only=False, + ) + + def check_valid_model(self, model_size): + """ + Check if it's a valid whisper model size. + + Args: + model_size (str): The name of the model size to check. + + Returns: + str: The model size if valid, None otherwise. + """ + if model_size not in self.model_sizes: + self.websocket.send( + json.dumps( + { + "uid": self.client_uid, + "status": "ERROR", + "message": f"Invalid model size {model_size}. Available choices: {self.model_sizes}" + } + ) + ) + return None + return model_size + + def set_language(self, info): + """ + Updates the language attribute based on the detected language information. + + Args: + info (object): An object containing the detected language and its probability. This object + must have at least two attributes: `language`, a string indicating the detected + language, and `language_probability`, a float representing the confidence level + of the language detection. + """ + if info.language_probability > 0.5: + self.language = info.language + logging.info(f"Detected language {self.language} with probability {info.language_probability}") + self.websocket.send(json.dumps( + {"uid": self.client_uid, "language": self.language, "language_prob": info.language_probability})) + + def transcribe_audio(self, input_sample): + """ + Transcribes the provided audio sample using the configured transcriber instance. + + If the language has not been set, it updates the session's language based on the transcription + information. + + Args: + input_sample (np.array): The audio chunk to be transcribed. This should be a NumPy + array representing the audio data. + + Returns: + The transcription result from the transcriber. The exact format of this result + depends on the implementation of the `transcriber.transcribe` method but typically + includes the transcribed text. + """ + if ServeClientFasterWhisper.SINGLE_MODEL: + ServeClientFasterWhisper.SINGLE_MODEL_LOCK.acquire() + result, info = self.transcriber.transcribe( + input_sample, + initial_prompt=self.initial_prompt, + language=self.language, + task=self.task, + vad_filter=self.use_vad, + vad_parameters=self.vad_parameters if self.use_vad else None) + if ServeClientFasterWhisper.SINGLE_MODEL: + ServeClientFasterWhisper.SINGLE_MODEL_LOCK.release() + + if self.language is None and info is not None: + self.set_language(info) + return result + + def get_previous_output(self): + """ + Retrieves previously generated transcription outputs if no new transcription is available + from the current audio chunks. + + Checks the time since the last transcription output and, if it is within a specified + threshold, returns the most recent segments of transcribed text. It also manages + adding a pause (blank segment) to indicate a significant gap in speech based on a defined + threshold. + + Returns: + segments (list): A list of transcription segments. This may include the most recent + transcribed text segments or a blank segment to indicate a pause + in speech. + """ + segments = [] + if self.t_start is None: + self.t_start = time.time() + if time.time() - self.t_start < self.show_prev_out_thresh: + segments = self.prepare_segments() + + # add a blank if there is no speech for 3 seconds + if len(self.text) and self.text[-1] != '': + if time.time() - self.t_start > self.add_pause_thresh: + self.text.append('') + return segments + + def handle_transcription_output(self, result, duration): + """ + Handle the transcription output, updating the transcript and sending data to the client. + + Args: + result (str): The result from whisper inference i.e. the list of segments. + duration (float): Duration of the transcribed audio chunk. + """ + segments = [] + if len(result): + self.t_start = None + last_segment = self.update_segments(result, duration) + segments = self.prepare_segments(last_segment) + else: + # show previous output if there is pause i.e. no output from whisper + segments = self.get_previous_output() + + if len(segments): + self.send_transcription_to_client(segments) + + def speech_to_text(self): + """ + Process an audio stream in an infinite loop, continuously transcribing the speech. + + This method continuously receives audio frames, performs real-time transcription, and sends + transcribed segments to the client via a WebSocket connection. + + If the client's language is not detected, it waits for 30 seconds of audio input to make a language prediction. + It utilizes the Whisper ASR model to transcribe the audio, continuously processing and streaming results. Segments + are sent to the client in real-time, and a history of segments is maintained to provide context.Pauses in speech + (no output from Whisper) are handled by showing the previous output for a set duration. A blank segment is added if + there is no speech for a specified duration to indicate a pause. + + Raises: + Exception: If there is an issue with audio processing or WebSocket communication. + + """ + while True: + if self.exit: + logging.info("Exiting speech to text thread") + break + + if self.frames_np is None: + continue + + self.clip_audio_if_no_valid_segment() + + input_bytes, duration = self.get_audio_chunk_for_processing() + if duration < 1.0: + time.sleep(0.1) # wait for audio chunks to arrive + continue + try: + input_sample = input_bytes.copy() + result = self.transcribe_audio(input_sample) + + if result is None or self.language is None: + self.timestamp_offset += duration + time.sleep(0.25) # wait for voice activity, result is None when no voice activity + continue + self.handle_transcription_output(result, duration) + + except Exception as e: + logging.error(f"[ERROR]: Failed to transcribe audio chunk: {e}") + time.sleep(0.01) + + def format_segment(self, start, end, text, completed=False): + """ + Formats a transcription segment with precise start and end times alongside the transcribed text. + + Args: + start (float): The start time of the transcription segment in seconds. + end (float): The end time of the transcription segment in seconds. + text (str): The transcribed text corresponding to the segment. + + Returns: + dict: A dictionary representing the formatted transcription segment, including + 'start' and 'end' times as strings with three decimal places and the 'text' + of the transcription. + """ + return { + 'start': "{:.3f}".format(start), + 'end': "{:.3f}".format(end), + 'text': text, + 'completed': completed + } + + def update_segments(self, segments, duration): + """ + Processes the segments from whisper. Appends all the segments to the list + except for the last segment assuming that it is incomplete. + + Updates the ongoing transcript with transcribed segments, including their start and end times. + Complete segments are appended to the transcript in chronological order. Incomplete segments + (assumed to be the last one) are processed to identify repeated content. If the same incomplete + segment is seen multiple times, it updates the offset and appends the segment to the transcript. + A threshold is used to detect repeated content and ensure it is only included once in the transcript. + The timestamp offset is updated based on the duration of processed segments. The method returns the + last processed segment, allowing it to be sent to the client for real-time updates. + + Args: + segments(dict) : dictionary of segments as returned by whisper + duration(float): duration of the current chunk + + Returns: + dict or None: The last processed segment with its start time, end time, and transcribed text. + Returns None if there are no valid segments to process. + """ + offset = None + self.current_out = '' + last_segment = None + + # process complete segments + if len(segments) > 1 and segments[-1].no_speech_prob <= self.no_speech_thresh: + for i, s in enumerate(segments[:-1]): + text_ = s.text + self.text.append(text_) + with self.lock: + start, end = self.timestamp_offset + s.start, self.timestamp_offset + min(duration, s.end) + + if start >= end: + continue + if s.no_speech_prob > self.no_speech_thresh: + continue + + self.transcript.append(self.format_segment(start, end, text_, completed=True)) + offset = min(duration, s.end) + + # only process the last segment if it satisfies the no_speech_thresh + if segments[-1].no_speech_prob <= self.no_speech_thresh: + self.current_out += segments[-1].text + with self.lock: + last_segment = self.format_segment( + self.timestamp_offset + segments[-1].start, + self.timestamp_offset + min(duration, segments[-1].end), + self.current_out, + completed=False + ) + + if self.current_out.strip() == self.prev_out.strip() and self.current_out != '': + self.same_output_count += 1 + + # if we remove the audio because of same output on the nth reptition we might remove the + # audio thats not yet transcribed so, capturing the time when it was repeated for the first time + if self.end_time_for_same_output is None: + self.end_time_for_same_output = segments[-1].end + time.sleep(0.1) # wait for some voice activity just in case there is an unitended pause from the speaker for better punctuations. + else: + self.same_output_count = 0 + self.end_time_for_same_output = None + + # if same incomplete segment is seen multiple times then update the offset + # and append the segment to the list + if self.same_output_count > self.same_output_threshold: + if not len(self.text) or self.text[-1].strip().lower() != self.current_out.strip().lower(): + self.text.append(self.current_out) + with self.lock: + self.transcript.append(self.format_segment( + self.timestamp_offset, + self.timestamp_offset + min(duration, self.end_time_for_same_output), + self.current_out, + completed=True + )) + self.current_out = '' + offset = min(duration, self.end_time_for_same_output) + self.same_output_count = 0 + last_segment = None + self.end_time_for_same_output = None + else: + self.prev_out = self.current_out + + # update offset + if offset is not None: + with self.lock: + self.timestamp_offset += offset + + return last_segment + diff --git a/whisper_live/backend/trt_backend.py b/whisper_live/backend/trt_backend.py new file mode 100644 index 0000000..5bfb9f7 --- /dev/null +++ b/whisper_live/backend/trt_backend.py @@ -0,0 +1,181 @@ +import json +import logging +import threading +import time + +from whisper_live.backend.base import ServeClientBase +from whisper_live.transcriber.transcriber_tensorrt import WhisperTRTLLM + + +class ServeClientTensorRT(ServeClientBase): + SINGLE_MODEL = None + SINGLE_MODEL_LOCK = threading.Lock() + + def __init__(self, websocket, task="transcribe", multilingual=False, language=None, client_uid=None, model=None, single_model=False): + """ + Initialize a ServeClient instance. + The Whisper model is initialized based on the client's language and device availability. + The transcription thread is started upon initialization. A "SERVER_READY" message is sent + to the client to indicate that the server is ready. + + Args: + websocket (WebSocket): The WebSocket connection for the client. + 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. + multilingual (bool, optional): Whether the client supports multilingual transcription. Defaults to False. + language (str, optional): The language for transcription. 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. + + """ + super().__init__(client_uid, websocket) + self.language = language if multilingual else "en" + self.task = task + self.eos = False + + if single_model: + if ServeClientTensorRT.SINGLE_MODEL is None: + self.create_model(model, multilingual) + ServeClientTensorRT.SINGLE_MODEL = self.transcriber + else: + self.transcriber = ServeClientTensorRT.SINGLE_MODEL + else: + self.create_model(model, multilingual) + + # threading + self.trans_thread = threading.Thread(target=self.speech_to_text) + self.trans_thread.start() + + self.websocket.send(json.dumps({ + "uid": self.client_uid, + "message": self.SERVER_READY, + "backend": "tensorrt" + })) + + def create_model(self, model, multilingual, warmup=True): + """ + Instantiates a new model, sets it as the transcriber and does warmup if desired. + """ + self.transcriber = WhisperTRTLLM( + model, + assets_dir="assets", + device="cuda", + is_multilingual=multilingual, + language=self.language, + task=self.task + ) + if warmup: + self.warmup() + + def warmup(self, warmup_steps=10): + """ + Warmup TensorRT since first few inferences are slow. + + Args: + warmup_steps (int): Number of steps to warm up the model for. + """ + logging.info("[INFO:] Warming up TensorRT engine..") + mel, _ = self.transcriber.log_mel_spectrogram("assets/jfk.flac") + for i in range(warmup_steps): + self.transcriber.transcribe(mel) + + def set_eos(self, eos): + """ + Sets the End of Speech (EOS) flag. + + Args: + eos (bool): The value to set for the EOS flag. + """ + self.lock.acquire() + self.eos = eos + self.lock.release() + + def handle_transcription_output(self, last_segment, duration): + """ + Handle the transcription output, updating the transcript and sending data to the client. + + Args: + last_segment (str): The last segment from the whisper output which is considered to be incomplete because + of the possibility of word being truncated. + duration (float): Duration of the transcribed audio chunk. + """ + segments = self.prepare_segments({"text": last_segment}) + self.send_transcription_to_client(segments) + if self.eos: + self.update_timestamp_offset(last_segment, duration) + + def transcribe_audio(self, input_bytes): + """ + Transcribe the audio chunk and send the results to the client. + + Args: + input_bytes (np.array): The audio chunk to transcribe. + """ + if ServeClientTensorRT.SINGLE_MODEL: + ServeClientTensorRT.SINGLE_MODEL_LOCK.acquire() + logging.info(f"[WhisperTensorRT:] Processing audio with duration: {input_bytes.shape[0] / self.RATE}") + mel, duration = self.transcriber.log_mel_spectrogram(input_bytes) + last_segment = self.transcriber.transcribe( + mel, + text_prefix=f"<|startoftranscript|><|{self.language}|><|{self.task}|><|notimestamps|>" + ) + if ServeClientTensorRT.SINGLE_MODEL: + ServeClientTensorRT.SINGLE_MODEL_LOCK.release() + if last_segment: + self.handle_transcription_output(last_segment, duration) + + def update_timestamp_offset(self, last_segment, duration): + """ + Update timestamp offset and transcript. + + Args: + last_segment (str): Last transcribed audio from the whisper model. + duration (float): Duration of the last audio chunk. + """ + if not len(self.transcript): + self.transcript.append({"text": last_segment + " "}) + elif self.transcript[-1]["text"].strip() != last_segment: + self.transcript.append({"text": last_segment + " "}) + + with self.lock: + self.timestamp_offset += duration + + def speech_to_text(self): + """ + Process an audio stream in an infinite loop, continuously transcribing the speech. + + This method continuously receives audio frames, performs real-time transcription, and sends + transcribed segments to the client via a WebSocket connection. + + If the client's language is not detected, it waits for 30 seconds of audio input to make a language prediction. + It utilizes the Whisper ASR model to transcribe the audio, continuously processing and streaming results. Segments + are sent to the client in real-time, and a history of segments is maintained to provide context.Pauses in speech + (no output from Whisper) are handled by showing the previous output for a set duration. A blank segment is added if + there is no speech for a specified duration to indicate a pause. + + Raises: + Exception: If there is an issue with audio processing or WebSocket communication. + + """ + while True: + if self.exit: + logging.info("Exiting speech to text thread") + break + + if self.frames_np is None: + time.sleep(0.02) # wait for any audio to arrive + continue + + self.clip_audio_if_no_valid_segment() + + input_bytes, duration = self.get_audio_chunk_for_processing() + if duration < 0.4: + continue + + try: + input_sample = input_bytes.copy() + logging.info(f"[WhisperTensorRT:] Processing audio with duration: {duration}") + self.transcribe_audio(input_sample) + + except Exception as e: + logging.error(f"[ERROR]: {e}") diff --git a/whisper_live/server.py b/whisper_live/server.py index f537b61..eceb353 100644 --- a/whisper_live/server.py +++ b/whisper_live/server.py @@ -7,16 +7,11 @@ import logging from enum import Enum from typing import List, Optional -import torch import numpy as np from websockets.sync.server import serve from websockets.exceptions import ConnectionClosed from whisper_live.vad import VoiceActivityDetector -from whisper_live.transcriber import WhisperModel -try: - from whisper_live.transcriber_tensorrt import WhisperTRTLLM -except Exception: - pass + logging.basicConfig(level=logging.INFO) @@ -127,6 +122,7 @@ class ClientManager: class BackendType(Enum): FASTER_WHISPER = "faster_whisper" TENSORRT = "tensorrt" + OPENVINO = "openvino" @staticmethod def valid_types() -> List[str]: @@ -141,6 +137,9 @@ class BackendType(Enum): def is_tensorrt(self) -> bool: return self == BackendType.TENSORRT + + def is_openvino(self) -> bool: + return self == BackendType.OPENVINO class TranscriptionServer: @@ -160,6 +159,7 @@ class TranscriptionServer: if self.backend.is_tensorrt(): try: + from whisper_live.backend.trt_backend import ServeClientTensorRT client = ServeClientTensorRT( websocket, multilingual=trt_multilingual, @@ -180,9 +180,33 @@ class TranscriptionServer: "Reverting to available backend: 'faster_whisper'" })) self.backend = BackendType.FASTER_WHISPER + + if self.backend.is_openvino(): + try: + from whisper_live.backend.openvino_backend import ServeClientOpenVINO + client = ServeClientOpenVINO( + websocket, + language=options["language"], + task=options["task"], + client_uid=options["uid"], + model=options["model"], + single_model=self.single_model, + ) + logging.info("Running OpenVINO backend.") + except Exception as e: + logging.error(f"OpenVINO not supported: {e}") + self.backend = BackendType.FASTER_WHISPER + self.client_uid = options["uid"] + websocket.send(json.dumps({ + "uid": self.client_uid, + "status": "WARNING", + "message": "OpenVINO not supported on Server yet. " + "Reverting to available backend: 'faster_whisper'" + })) try: if self.backend.is_faster_whisper(): + from whisper_live.backend.faster_whisper_backend import ServeClientFasterWhisper if faster_whisper_custom_model_path is not None and os.path.exists(faster_whisper_custom_model_path): logging.info(f"Using custom model {faster_whisper_custom_model_path}") options["model"] = faster_whisper_custom_model_path @@ -200,6 +224,7 @@ class TranscriptionServer: logging.info("Running faster_whisper backend.") except Exception as e: + logging.error(e) return if client is None: @@ -403,732 +428,3 @@ class TranscriptionServer: if self.client_manager.get_client(websocket): self.client_manager.remove_client(websocket) - -class ServeClientBase(object): - RATE = 16000 - SERVER_READY = "SERVER_READY" - DISCONNECT = "DISCONNECT" - - def __init__(self, client_uid, websocket): - self.client_uid = client_uid - self.websocket = websocket - self.frames = b"" - self.timestamp_offset = 0.0 - self.frames_np = None - self.frames_offset = 0.0 - self.text = [] - self.current_out = '' - self.prev_out = '' - self.t_start = None - self.exit = False - self.same_output_count = 0 - self.show_prev_out_thresh = 5 # if pause(no output from whisper) show previous output for 5 seconds - self.add_pause_thresh = 3 # add a blank to segment list as a pause(no speech) for 3 seconds - self.transcript = [] - self.send_last_n_segments = 10 - - # text formatting - self.pick_previous_segments = 2 - - # threading - self.lock = threading.Lock() - - def speech_to_text(self): - raise NotImplementedError - - def transcribe_audio(self): - raise NotImplementedError - - def handle_transcription_output(self): - raise NotImplementedError - - def add_frames(self, frame_np): - """ - Add audio frames to the ongoing audio stream buffer. - - This method is responsible for maintaining the audio stream buffer, allowing the continuous addition - of audio frames as they are received. It also ensures that the buffer does not exceed a specified size - to prevent excessive memory usage. - - If the buffer size exceeds a threshold (45 seconds of audio data), it discards the oldest 30 seconds - of audio data to maintain a reasonable buffer size. If the buffer is empty, it initializes it with the provided - audio frame. The audio stream buffer is used for real-time processing of audio data for transcription. - - Args: - frame_np (numpy.ndarray): The audio frame data as a NumPy array. - - """ - self.lock.acquire() - if self.frames_np is not None and self.frames_np.shape[0] > 45*self.RATE: - self.frames_offset += 30.0 - self.frames_np = self.frames_np[int(30*self.RATE):] - # check timestamp offset(should be >= self.frame_offset) - # this basically means that there is no speech as timestamp offset hasnt updated - # and is less than frame_offset - if self.timestamp_offset < self.frames_offset: - self.timestamp_offset = self.frames_offset - if self.frames_np is None: - self.frames_np = frame_np.copy() - else: - self.frames_np = np.concatenate((self.frames_np, frame_np), axis=0) - self.lock.release() - - def clip_audio_if_no_valid_segment(self): - """ - Update the timestamp offset based on audio buffer status. - Clip audio if the current chunk exceeds 30 seconds, this basically implies that - no valid segment for the last 30 seconds from whisper - """ - with self.lock: - if self.frames_np[int((self.timestamp_offset - self.frames_offset)*self.RATE):].shape[0] > 25 * self.RATE: - duration = self.frames_np.shape[0] / self.RATE - self.timestamp_offset = self.frames_offset + duration - 5 - - def get_audio_chunk_for_processing(self): - """ - Retrieves the next chunk of audio data for processing based on the current offsets. - - Calculates which part of the audio data should be processed next, based on - the difference between the current timestamp offset and the frame's offset, scaled by - the audio sample rate (RATE). It then returns this chunk of audio data along with its - duration in seconds. - - Returns: - tuple: A tuple containing: - - input_bytes (np.ndarray): The next chunk of audio data to be processed. - - duration (float): The duration of the audio chunk in seconds. - """ - with self.lock: - samples_take = max(0, (self.timestamp_offset - self.frames_offset) * self.RATE) - input_bytes = self.frames_np[int(samples_take):].copy() - duration = input_bytes.shape[0] / self.RATE - return input_bytes, duration - - def prepare_segments(self, last_segment=None): - """ - Prepares the segments of transcribed text to be sent to the client. - - This method compiles the recent segments of transcribed text, ensuring that only the - specified number of the most recent segments are included. It also appends the most - recent segment of text if provided (which is considered incomplete because of the possibility - of the last word being truncated in the audio chunk). - - Args: - last_segment (str, optional): The most recent segment of transcribed text to be added - to the list of segments. Defaults to None. - - Returns: - list: A list of transcribed text segments to be sent to the client. - """ - segments = [] - if len(self.transcript) >= self.send_last_n_segments: - segments = self.transcript[-self.send_last_n_segments:].copy() - else: - segments = self.transcript.copy() - if last_segment is not None: - segments = segments + [last_segment] - return segments - - def get_audio_chunk_duration(self, input_bytes): - """ - Calculates the duration of the provided audio chunk. - - Args: - input_bytes (numpy.ndarray): The audio chunk for which to calculate the duration. - - Returns: - float: The duration of the audio chunk in seconds. - """ - return input_bytes.shape[0] / self.RATE - - def send_transcription_to_client(self, segments): - """ - Sends the specified transcription segments to the client over the websocket connection. - - This method formats the transcription segments into a JSON object and attempts to send - this object to the client. If an error occurs during the send operation, it logs the error. - - Returns: - segments (list): A list of transcription segments to be sent to the client. - """ - try: - self.websocket.send( - json.dumps({ - "uid": self.client_uid, - "segments": segments, - }) - ) - except Exception as e: - logging.error(f"[ERROR]: Sending data to client: {e}") - - def disconnect(self): - """ - Notify the client of disconnection and send a disconnect message. - - This method sends a disconnect message to the client via the WebSocket connection to notify them - that the transcription service is disconnecting gracefully. - - """ - self.websocket.send(json.dumps({ - "uid": self.client_uid, - "message": self.DISCONNECT - })) - - def cleanup(self): - """ - Perform cleanup tasks before exiting the transcription service. - - This method performs necessary cleanup tasks, including stopping the transcription thread, marking - the exit flag to indicate the transcription thread should exit gracefully, and destroying resources - associated with the transcription process. - - """ - logging.info("Cleaning up.") - self.exit = True - - -class ServeClientTensorRT(ServeClientBase): - - SINGLE_MODEL = None - SINGLE_MODEL_LOCK = threading.Lock() - - def __init__(self, websocket, task="transcribe", multilingual=False, language=None, client_uid=None, model=None, single_model=False): - """ - Initialize a ServeClient instance. - The Whisper model is initialized based on the client's language and device availability. - The transcription thread is started upon initialization. A "SERVER_READY" message is sent - to the client to indicate that the server is ready. - - Args: - websocket (WebSocket): The WebSocket connection for the client. - 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. - multilingual (bool, optional): Whether the client supports multilingual transcription. Defaults to False. - language (str, optional): The language for transcription. 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. - - """ - super().__init__(client_uid, websocket) - self.language = language if multilingual else "en" - self.task = task - self.eos = False - - if single_model: - if ServeClientTensorRT.SINGLE_MODEL is None: - self.create_model(model, multilingual) - ServeClientTensorRT.SINGLE_MODEL = self.transcriber - else: - self.transcriber = ServeClientTensorRT.SINGLE_MODEL - else: - self.create_model(model, multilingual) - - # threading - self.trans_thread = threading.Thread(target=self.speech_to_text) - self.trans_thread.start() - - self.websocket.send(json.dumps({ - "uid": self.client_uid, - "message": self.SERVER_READY, - "backend": "tensorrt" - })) - - def create_model(self, model, multilingual, warmup=True): - """ - Instantiates a new model, sets it as the transcriber and does warmup if desired. - """ - self.transcriber = WhisperTRTLLM( - model, - assets_dir="assets", - device="cuda", - is_multilingual=multilingual, - language=self.language, - task=self.task - ) - if warmup: - self.warmup() - - def warmup(self, warmup_steps=10): - """ - Warmup TensorRT since first few inferences are slow. - - Args: - warmup_steps (int): Number of steps to warm up the model for. - """ - logging.info("[INFO:] Warming up TensorRT engine..") - mel, _ = self.transcriber.log_mel_spectrogram("assets/jfk.flac") - for i in range(warmup_steps): - self.transcriber.transcribe(mel) - - def set_eos(self, eos): - """ - Sets the End of Speech (EOS) flag. - - Args: - eos (bool): The value to set for the EOS flag. - """ - self.lock.acquire() - self.eos = eos - self.lock.release() - - def handle_transcription_output(self, last_segment, duration): - """ - Handle the transcription output, updating the transcript and sending data to the client. - - Args: - last_segment (str): The last segment from the whisper output which is considered to be incomplete because - of the possibility of word being truncated. - duration (float): Duration of the transcribed audio chunk. - """ - segments = self.prepare_segments({"text": last_segment}) - self.send_transcription_to_client(segments) - if self.eos: - self.update_timestamp_offset(last_segment, duration) - - def transcribe_audio(self, input_bytes): - """ - Transcribe the audio chunk and send the results to the client. - - Args: - input_bytes (np.array): The audio chunk to transcribe. - """ - if ServeClientTensorRT.SINGLE_MODEL: - ServeClientTensorRT.SINGLE_MODEL_LOCK.acquire() - logging.info(f"[WhisperTensorRT:] Processing audio with duration: {input_bytes.shape[0] / self.RATE}") - mel, duration = self.transcriber.log_mel_spectrogram(input_bytes) - last_segment = self.transcriber.transcribe( - mel, - text_prefix=f"<|startoftranscript|><|{self.language}|><|{self.task}|><|notimestamps|>" - ) - if ServeClientTensorRT.SINGLE_MODEL: - ServeClientTensorRT.SINGLE_MODEL_LOCK.release() - if last_segment: - self.handle_transcription_output(last_segment, duration) - - def update_timestamp_offset(self, last_segment, duration): - """ - Update timestamp offset and transcript. - - Args: - last_segment (str): Last transcribed audio from the whisper model. - duration (float): Duration of the last audio chunk. - """ - if not len(self.transcript): - self.transcript.append({"text": last_segment + " "}) - elif self.transcript[-1]["text"].strip() != last_segment: - self.transcript.append({"text": last_segment + " "}) - - with self.lock: - self.timestamp_offset += duration - - def speech_to_text(self): - """ - Process an audio stream in an infinite loop, continuously transcribing the speech. - - This method continuously receives audio frames, performs real-time transcription, and sends - transcribed segments to the client via a WebSocket connection. - - If the client's language is not detected, it waits for 30 seconds of audio input to make a language prediction. - It utilizes the Whisper ASR model to transcribe the audio, continuously processing and streaming results. Segments - are sent to the client in real-time, and a history of segments is maintained to provide context.Pauses in speech - (no output from Whisper) are handled by showing the previous output for a set duration. A blank segment is added if - there is no speech for a specified duration to indicate a pause. - - Raises: - Exception: If there is an issue with audio processing or WebSocket communication. - - """ - while True: - if self.exit: - logging.info("Exiting speech to text thread") - break - - if self.frames_np is None: - time.sleep(0.02) # wait for any audio to arrive - continue - - self.clip_audio_if_no_valid_segment() - - input_bytes, duration = self.get_audio_chunk_for_processing() - if duration < 0.4: - continue - - try: - input_sample = input_bytes.copy() - logging.info(f"[WhisperTensorRT:] Processing audio with duration: {duration}") - self.transcribe_audio(input_sample) - - except Exception as e: - logging.error(f"[ERROR]: {e}") - - -class ServeClientFasterWhisper(ServeClientBase): - - SINGLE_MODEL = None - SINGLE_MODEL_LOCK = threading.Lock() - - def __init__(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): - """ - Initialize a ServeClient instance. - The Whisper model is initialized based on the client's language and device availability. - The transcription thread is started upon initialization. A "SERVER_READY" message is sent - to the client to indicate that the server is ready. - - Args: - websocket (WebSocket): The WebSocket connection for the client. - 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. - language (str, optional): The language for transcription. 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' - 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. - """ - super().__init__(client_uid, websocket) - self.model_sizes = [ - "tiny", "tiny.en", "base", "base.en", "small", "small.en", - "medium", "medium.en", "large-v2", "large-v3", "distil-small.en", - "distil-medium.en", "distil-large-v2", "distil-large-v3", - "large-v3-turbo", "turbo" - ] - - self.model_size_or_path = model - self.language = "en" if self.model_size_or_path.endswith("en") else language - self.task = task - self.initial_prompt = initial_prompt - self.vad_parameters = vad_parameters or {"onset": 0.5} - self.no_speech_thresh = 0.45 - self.same_output_threshold = 10 - self.end_time_for_same_output = None - - device = "cuda" if torch.cuda.is_available() else "cpu" - if device == "cuda": - major, _ = torch.cuda.get_device_capability(device) - self.compute_type = "float16" if major >= 7 else "float32" - else: - self.compute_type = "int8" - - if self.model_size_or_path is None: - return - logging.info(f"Using Device={device} with precision {self.compute_type}") - - try: - if single_model: - if ServeClientFasterWhisper.SINGLE_MODEL is None: - self.create_model(device) - ServeClientFasterWhisper.SINGLE_MODEL = self.transcriber - else: - self.transcriber = ServeClientFasterWhisper.SINGLE_MODEL - else: - self.create_model(device) - except Exception as e: - logging.error(f"Failed to load model: {e}") - self.websocket.send(json.dumps({ - "uid": self.client_uid, - "status": "ERROR", - "message": f"Failed to load model: {str(self.model_size_or_path)}" - })) - self.websocket.close() - return - - self.use_vad = use_vad - - # threading - self.trans_thread = threading.Thread(target=self.speech_to_text) - self.trans_thread.start() - self.websocket.send( - json.dumps( - { - "uid": self.client_uid, - "message": self.SERVER_READY, - "backend": "faster_whisper" - } - ) - ) - - def create_model(self, device): - """ - Instantiates a new model, sets it as the transcriber. - """ - self.transcriber = WhisperModel( - self.model_size_or_path, - device=device, - compute_type=self.compute_type, - local_files_only=False, - ) - - def check_valid_model(self, model_size): - """ - Check if it's a valid whisper model size. - - Args: - model_size (str): The name of the model size to check. - - Returns: - str: The model size if valid, None otherwise. - """ - if model_size not in self.model_sizes: - self.websocket.send( - json.dumps( - { - "uid": self.client_uid, - "status": "ERROR", - "message": f"Invalid model size {model_size}. Available choices: {self.model_sizes}" - } - ) - ) - return None - return model_size - - def set_language(self, info): - """ - Updates the language attribute based on the detected language information. - - Args: - info (object): An object containing the detected language and its probability. This object - must have at least two attributes: `language`, a string indicating the detected - language, and `language_probability`, a float representing the confidence level - of the language detection. - """ - if info.language_probability > 0.5: - self.language = info.language - logging.info(f"Detected language {self.language} with probability {info.language_probability}") - self.websocket.send(json.dumps( - {"uid": self.client_uid, "language": self.language, "language_prob": info.language_probability})) - - def transcribe_audio(self, input_sample): - """ - Transcribes the provided audio sample using the configured transcriber instance. - - If the language has not been set, it updates the session's language based on the transcription - information. - - Args: - input_sample (np.array): The audio chunk to be transcribed. This should be a NumPy - array representing the audio data. - - Returns: - The transcription result from the transcriber. The exact format of this result - depends on the implementation of the `transcriber.transcribe` method but typically - includes the transcribed text. - """ - if ServeClientFasterWhisper.SINGLE_MODEL: - ServeClientFasterWhisper.SINGLE_MODEL_LOCK.acquire() - result, info = self.transcriber.transcribe( - input_sample, - initial_prompt=self.initial_prompt, - language=self.language, - task=self.task, - vad_filter=self.use_vad, - vad_parameters=self.vad_parameters if self.use_vad else None) - if ServeClientFasterWhisper.SINGLE_MODEL: - ServeClientFasterWhisper.SINGLE_MODEL_LOCK.release() - - if self.language is None and info is not None: - self.set_language(info) - return result - - def get_previous_output(self): - """ - Retrieves previously generated transcription outputs if no new transcription is available - from the current audio chunks. - - Checks the time since the last transcription output and, if it is within a specified - threshold, returns the most recent segments of transcribed text. It also manages - adding a pause (blank segment) to indicate a significant gap in speech based on a defined - threshold. - - Returns: - segments (list): A list of transcription segments. This may include the most recent - transcribed text segments or a blank segment to indicate a pause - in speech. - """ - segments = [] - if self.t_start is None: - self.t_start = time.time() - if time.time() - self.t_start < self.show_prev_out_thresh: - segments = self.prepare_segments() - - # add a blank if there is no speech for 3 seconds - if len(self.text) and self.text[-1] != '': - if time.time() - self.t_start > self.add_pause_thresh: - self.text.append('') - return segments - - def handle_transcription_output(self, result, duration): - """ - Handle the transcription output, updating the transcript and sending data to the client. - - Args: - result (str): The result from whisper inference i.e. the list of segments. - duration (float): Duration of the transcribed audio chunk. - """ - segments = [] - if len(result): - self.t_start = None - last_segment = self.update_segments(result, duration) - segments = self.prepare_segments(last_segment) - else: - # show previous output if there is pause i.e. no output from whisper - segments = self.get_previous_output() - - if len(segments): - self.send_transcription_to_client(segments) - - def speech_to_text(self): - """ - Process an audio stream in an infinite loop, continuously transcribing the speech. - - This method continuously receives audio frames, performs real-time transcription, and sends - transcribed segments to the client via a WebSocket connection. - - If the client's language is not detected, it waits for 30 seconds of audio input to make a language prediction. - It utilizes the Whisper ASR model to transcribe the audio, continuously processing and streaming results. Segments - are sent to the client in real-time, and a history of segments is maintained to provide context.Pauses in speech - (no output from Whisper) are handled by showing the previous output for a set duration. A blank segment is added if - there is no speech for a specified duration to indicate a pause. - - Raises: - Exception: If there is an issue with audio processing or WebSocket communication. - - """ - while True: - if self.exit: - logging.info("Exiting speech to text thread") - break - - if self.frames_np is None: - continue - - self.clip_audio_if_no_valid_segment() - - input_bytes, duration = self.get_audio_chunk_for_processing() - if duration < 1.0: - time.sleep(0.1) # wait for audio chunks to arrive - continue - try: - input_sample = input_bytes.copy() - result = self.transcribe_audio(input_sample) - - if result is None or self.language is None: - self.timestamp_offset += duration - time.sleep(0.25) # wait for voice activity, result is None when no voice activity - continue - self.handle_transcription_output(result, duration) - - except Exception as e: - logging.error(f"[ERROR]: Failed to transcribe audio chunk: {e}") - time.sleep(0.01) - - def format_segment(self, start, end, text, completed=False): - """ - Formats a transcription segment with precise start and end times alongside the transcribed text. - - Args: - start (float): The start time of the transcription segment in seconds. - end (float): The end time of the transcription segment in seconds. - text (str): The transcribed text corresponding to the segment. - - Returns: - dict: A dictionary representing the formatted transcription segment, including - 'start' and 'end' times as strings with three decimal places and the 'text' - of the transcription. - """ - return { - 'start': "{:.3f}".format(start), - 'end': "{:.3f}".format(end), - 'text': text, - 'completed': completed - } - - def update_segments(self, segments, duration): - """ - Processes the segments from whisper. Appends all the segments to the list - except for the last segment assuming that it is incomplete. - - Updates the ongoing transcript with transcribed segments, including their start and end times. - Complete segments are appended to the transcript in chronological order. Incomplete segments - (assumed to be the last one) are processed to identify repeated content. If the same incomplete - segment is seen multiple times, it updates the offset and appends the segment to the transcript. - A threshold is used to detect repeated content and ensure it is only included once in the transcript. - The timestamp offset is updated based on the duration of processed segments. The method returns the - last processed segment, allowing it to be sent to the client for real-time updates. - - Args: - segments(dict) : dictionary of segments as returned by whisper - duration(float): duration of the current chunk - - Returns: - dict or None: The last processed segment with its start time, end time, and transcribed text. - Returns None if there are no valid segments to process. - """ - offset = None - self.current_out = '' - last_segment = None - - # process complete segments - if len(segments) > 1 and segments[-1].no_speech_prob <= self.no_speech_thresh: - for i, s in enumerate(segments[:-1]): - text_ = s.text - self.text.append(text_) - with self.lock: - start, end = self.timestamp_offset + s.start, self.timestamp_offset + min(duration, s.end) - - if start >= end: - continue - if s.no_speech_prob > self.no_speech_thresh: - continue - - self.transcript.append(self.format_segment(start, end, text_, completed=True)) - offset = min(duration, s.end) - - # only process the last segment if it satisfies the no_speech_thresh - if segments[-1].no_speech_prob <= self.no_speech_thresh: - self.current_out += segments[-1].text - with self.lock: - last_segment = self.format_segment( - self.timestamp_offset + segments[-1].start, - self.timestamp_offset + min(duration, segments[-1].end), - self.current_out, - completed=False - ) - - if self.current_out.strip() == self.prev_out.strip() and self.current_out != '': - self.same_output_count += 1 - - # if we remove the audio because of same output on the nth reptition we might remove the - # audio thats not yet transcribed so, capturing the time when it was repeated for the first time - if self.end_time_for_same_output is None: - self.end_time_for_same_output = segments[-1].end - time.sleep(0.1) # wait for some voice activity just in case there is an unitended pause from the speaker for better punctuations. - else: - self.same_output_count = 0 - self.end_time_for_same_output = None - - # if same incomplete segment is seen multiple times then update the offset - # and append the segment to the list - if self.same_output_count > self.same_output_threshold: - if not len(self.text) or self.text[-1].strip().lower() != self.current_out.strip().lower(): - self.text.append(self.current_out) - with self.lock: - self.transcript.append(self.format_segment( - self.timestamp_offset, - self.timestamp_offset + min(duration, self.end_time_for_same_output), - self.current_out, - completed=True - )) - self.current_out = '' - offset = min(duration, self.end_time_for_same_output) - self.same_output_count = 0 - last_segment = None - self.end_time_for_same_output = None - else: - self.prev_out = self.current_out - - # update offset - if offset is not None: - with self.lock: - self.timestamp_offset += offset - - return last_segment diff --git a/whisper_live/transcriber/__init__.py b/whisper_live/transcriber/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/whisper_live/tensorrt_utils.py b/whisper_live/transcriber/tensorrt_utils.py similarity index 100% rename from whisper_live/tensorrt_utils.py rename to whisper_live/transcriber/tensorrt_utils.py diff --git a/whisper_live/transcriber.py b/whisper_live/transcriber/transcriber_faster_whisper.py similarity index 100% rename from whisper_live/transcriber.py rename to whisper_live/transcriber/transcriber_faster_whisper.py diff --git a/whisper_live/transcriber_tensorrt.py b/whisper_live/transcriber/transcriber_tensorrt.py similarity index 99% rename from whisper_live/transcriber_tensorrt.py rename to whisper_live/transcriber/transcriber_tensorrt.py index 8742e69..7123534 100644 --- a/whisper_live/transcriber_tensorrt.py +++ b/whisper_live/transcriber/transcriber_tensorrt.py @@ -9,7 +9,12 @@ import torch import numpy as np import torch.nn.functional as F from whisper.tokenizer import get_tokenizer -from whisper_live.tensorrt_utils import (mel_filters, load_audio_wav_format, pad_or_trim, load_audio) +from whisper_live.transcriber.tensorrt_utils import ( + mel_filters, + load_audio_wav_format, + pad_or_trim, + load_audio +) import tensorrt_llm import tensorrt_llm.logger as logger