1 Commits

Author SHA1 Message Date
makaveli10 09670dd3c7 add eos to faster_whisper server
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-07-11 07:01:34 -04:00
4 changed files with 58 additions and 82 deletions
+1 -1
View File
@@ -1,5 +1,5 @@
faster-whisper==1.0.1
torch==2.3.0
torch
websockets
onnxruntime==1.16.0
numba
+1 -1
View File
@@ -1 +1 @@
__version__ = "0.5.1"
__version__ = "0.5.0"
+8 -24
View File
@@ -2,7 +2,6 @@ import os
import shutil
import wave
import logging
import numpy as np
import pyaudio
import threading
@@ -29,8 +28,7 @@ class Client:
translate=False,
model="small",
srt_file_path="output.srt",
use_vad=True,
log_transcription=True
use_vad=True
):
"""
Initializes a Client instance for audio recording and streaming to a server.
@@ -58,11 +56,11 @@ class Client:
self.use_vad = use_vad
self.last_segment = None
self.last_received_segment = None
self.log_transcription = log_transcription
if translate:
self.task = "translate"
self.timestamp_offset = 0.0
self.audio_bytes = None
if host is not None and port is not None:
@@ -119,11 +117,10 @@ class Client:
self.last_response_received = time.time()
self.last_received_segment = segments[-1]["text"]
if self.log_transcription:
# Truncate to last 3 entries for brevity.
text = text[-3:]
utils.clear_screen()
utils.print_transcript(text)
# Truncate to last 3 entries for brevity.
text = text[-3:]
utils.clear_screen()
utils.print_transcript(text)
def on_message(self, ws, message):
"""
@@ -434,8 +431,6 @@ class TranscriptionTeeClient:
def handle_ffmpeg_process(self, process, stream_type):
print(f"[INFO]: Connecting to {stream_type} stream...")
stderr_thread = threading.Thread(target=self.consume_stderr, args=(process,))
stderr_thread.start()
try:
# Process the stream
while True:
@@ -482,16 +477,6 @@ class TranscriptionTeeClient:
return process
def consume_stderr(self, process):
"""
Consume and log the stderr output of a process in a separate thread.
Args:
process (subprocess.Popen): The process whose stderr output will be logged.
"""
for line in iter(process.stderr.readline, b""):
logging.debug(f'[STDERR]: {line.decode()}')
def save_chunk(self, n_audio_file):
"""
Saves the current audio frames to a WAV file in a separate thread.
@@ -679,10 +664,9 @@ class TranscriptionClient(TranscriptionTeeClient):
use_vad=True,
save_output_recording=False,
output_recording_filename="./output_recording.wav",
output_transcription_path="./output.srt",
log_transcription=True,
output_transcription_path="./output.srt"
):
self.client = Client(host, port, lang, translate, model, srt_file_path=output_transcription_path, use_vad=use_vad, log_transcription=log_transcription)
self.client = Client(host, port, lang, translate, model, srt_file_path=output_transcription_path, use_vad=use_vad)
if save_output_recording and not output_recording_filename.endswith(".wav"):
raise ValueError(f"Please provide a valid `output_recording_filename`: {output_recording_filename}")
if not output_transcription_path.endswith(".srt"):
+48 -56
View File
@@ -229,8 +229,7 @@ class TranscriptionServer:
websocket.close()
return False # Indicates that the connection should not continue
if self.backend.is_tensorrt():
self.vad_detector = VoiceActivityDetector(frame_rate=self.RATE)
self.vad_detector = VoiceActivityDetector(frame_rate=self.RATE)
self.initialize_client(websocket, options, faster_whisper_custom_model_path,
whisper_tensorrt_path, trt_multilingual)
return True
@@ -248,17 +247,15 @@ class TranscriptionServer:
frame_np = self.get_audio_from_websocket(websocket)
client = self.client_manager.get_client(websocket)
if frame_np is False:
if self.backend.is_tensorrt():
client.set_eos(True)
client.set_eos(True)
return False
if self.backend.is_tensorrt():
voice_active = self.voice_activity(websocket, frame_np)
if voice_active:
self.no_voice_activity_chunks = 0
client.set_eos(False)
if self.use_vad and not voice_active:
return True
voice_active = self.voice_activity(websocket, frame_np)
if voice_active:
self.no_voice_activity_chunks = 0
client.set_eos(False)
if self.use_vad and not voice_active:
return True
client.add_frames(frame_np)
return True
@@ -331,13 +328,8 @@ class TranscriptionServer:
raise ValueError(f"Custom faster_whisper model '{faster_whisper_custom_model_path}' is not a valid path.")
if whisper_tensorrt_path is not None and not os.path.exists(whisper_tensorrt_path):
raise ValueError(f"TensorRT model '{whisper_tensorrt_path}' is not a valid path.")
if single_model:
if faster_whisper_custom_model_path or whisper_tensorrt_path:
logging.info("Custom model option was provided. Switching to single model mode.")
self.single_model = True
# TODO: load model initially
else:
logging.info("Single model mode currently only works with custom models.")
self.single_model = single_model
if not BackendType.is_valid(backend):
raise ValueError(f"{backend} is not a valid backend type. Choose backend from {BackendType.valid_types()}")
with serve(
@@ -416,6 +408,7 @@ class ServeClientBase(object):
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
self.eos = False
# text formatting
self.pick_previous_segments = 2
@@ -423,6 +416,18 @@ class ServeClientBase(object):
# threading
self.lock = threading.Lock()
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 speech_to_text(self):
raise NotImplementedError
@@ -543,7 +548,8 @@ class ServeClientBase(object):
self.websocket.send(
json.dumps({
"uid": self.client_uid,
"segments": segments,
"text": segments,
"eos": self.eos
})
)
except Exception as e:
@@ -648,17 +654,6 @@ class ServeClientTensorRT(ServeClientBase):
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.
@@ -784,24 +779,19 @@ class ServeClientFasterWhisper(ServeClientBase):
self.task = task
self.initial_prompt = initial_prompt
self.vad_parameters = vad_parameters or {"threshold": 0.5}
self.no_speech_thresh = 0.45
self.no_speech_thresh = 0.35
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}")
if single_model:
if ServeClientFasterWhisper.SINGLE_MODEL is None:
self.create_model(device)
ServeClientFasterWhisper.SINGLE_MODEL = self.transcriber
else:
print("Re-using already initialized model.")
self.transcriber = ServeClientFasterWhisper.SINGLE_MODEL
else:
self.create_model(device)
@@ -828,7 +818,7 @@ class ServeClientFasterWhisper(ServeClientBase):
self.transcriber = WhisperModel(
self.model_size_or_path,
device=device,
compute_type=self.compute_type,
compute_type="int8" if device == "cpu" else "float16",
local_files_only=False,
)
@@ -894,8 +884,9 @@ class ServeClientFasterWhisper(ServeClientBase):
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)
vad_filter=False,
vad_parameters=self.vad_parameters if self.use_vad else None,
beam_size=5)
if ServeClientFasterWhisper.SINGLE_MODEL:
ServeClientFasterWhisper.SINGLE_MODEL_LOCK.release()
@@ -938,17 +929,16 @@ class ServeClientFasterWhisper(ServeClientBase):
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)
if len(self.text):
if self.eos and last_segment is None:
self.send_transcription_to_client(' '.join([s.strip() for s in self.text]))
self.set_eos(False)
self.text = []
elif not self.eos:
self.send_transcription_to_client(' '.join([s.strip() for s in self.text]))
def speech_to_text(self):
"""
@@ -978,8 +968,12 @@ class ServeClientFasterWhisper(ServeClientBase):
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
if duration < 0.6:
if len(self.text) and self.eos:
self.send_transcription_to_client(' '.join([s.strip() for s in self.text]))
self.set_eos(False)
self.text = []
time.sleep(0.1)
continue
try:
input_sample = input_bytes.copy()
@@ -987,7 +981,7 @@ class ServeClientFasterWhisper(ServeClientBase):
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
time.sleep(0.1) # wait for voice activity, result is None when no voice activity
continue
self.handle_transcription_output(result, duration)
@@ -1036,15 +1030,13 @@ class ServeClientFasterWhisper(ServeClientBase):
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.
"""
last_segment = None
offset = None
self.current_out = ''
last_segment = None
# process complete segments
if len(segments) > 1:
for i, s in enumerate(segments[:-1]):
text_ = s.text
self.text.append(text_)
start, end = self.timestamp_offset + s.start, self.timestamp_offset + min(duration, s.end)
if start >= end:
@@ -1052,10 +1044,10 @@ class ServeClientFasterWhisper(ServeClientBase):
if s.no_speech_prob > self.no_speech_thresh:
continue
self.text.append(text_)
self.transcript.append(self.format_segment(start, end, text_))
offset = min(duration, s.end)
# only process the segments if it satisfies the no_speech_thresh
if segments[-1].no_speech_prob <= self.no_speech_thresh:
self.current_out += segments[-1].text
last_segment = self.format_segment(
@@ -1071,7 +1063,7 @@ class ServeClientFasterWhisper(ServeClientBase):
else:
self.same_output_threshold = 0
if self.same_output_threshold > 5:
if self.same_output_threshold > 2:
if not len(self.text) or self.text[-1].strip().lower() != self.current_out.strip().lower():
self.text.append(self.current_out)
self.transcript.append(self.format_segment(