20 Commits

Author SHA1 Message Date
makaveli bc070d6688 Bump version 0.5.1 2024-09-05 09:34:30 +05:30
Marcus Edel 8e7e329a39 Merge pull request #274 from makaveli10/fallback_to_fp32
Set compute_type based on device capability.
2024-09-03 09:17:31 -04:00
makaveli10 380f07394b Set compute_type based on device capability
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-09-03 00:47:38 -04:00
Marcus Edel 30f78a2cc6 Merge pull request #272 from makaveli10/fix_last_segment_init
Initialize last_segment to None.
2024-08-30 12:39:49 -04:00
makaveli10 01c6bc1ecd Initialize last_segment to None
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-08-30 07:47:03 -04:00
Marcus Edel bdaed45820 Merge pull request #262 from makaveli10/discard_no_speech_segments
Discard no speech segments.
2024-08-19 10:12:39 -04:00
makaveli10 4870e9fb9e Make text logging optional
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-08-08 06:09:44 -04:00
makaveli10 ccb183b4d8 Pin torch version to 2.3.0
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-08-08 06:05:44 -04:00
makaveli10 fac62aaccc Fix hallucinations with no_speech_thres
Signed-off-by: makaveli10 <vineet.suryan@collabora.com>
2024-08-08 06:05:12 -04:00
makaveli aade67736a Merge pull request #257 from sondt2709/fix-ffmpeg-subprocess-deadlock
Fix deadlock issue in FFmpeg subprocess by ensuring stderr is consumed
2024-07-19 14:49:10 +05:30
Sean Dang abfe830eee Fix deadlock issue in FFmpeg subprocess by ensuring stderr is consumed 2024-07-11 23:30:44 +07:00
makaveli cb392cbb93 Merge pull request #247 from makaveli10/pin_sliero_vad_model_version
Pin silero VAD onnx model version to v4.0
2024-07-09 12:58:27 +05:30
makaveli10 42733da59a Pin numpy version to <2
Signed-off-by: makaveli10 <suryanvineet47@gmail.com>
2024-07-02 11:49:30 +05:30
makaveli10 26c517021f Pin silero VAD onnx model version to v4.0
Signed-off-by: makaveli10 <suryanvineet47@gmail.com>
2024-07-02 11:01:40 +05:30
makaveli cf721e8b53 Merge pull request #243 from berkaybilik/making_backend_arg_safer
Making backend arg safer
2024-07-02 10:57:58 +05:30
makaveli 5985ec82b6 Merge pull request #236 from t-nil/patch-1
Backslash missing in example
2024-06-30 20:33:07 +05:30
berkaybilik 2f1c934ea2 always use the BackendType enum to reference the backend inside the TranscriptionServer 2024-06-27 00:20:17 +01:00
berkaybilik b220ccb330 fixed reference before assignment error/warning 2024-06-27 00:11:06 +01:00
berkaybilik 5e3906fc7b use enum to validate backend validity in server.run 2024-06-26 23:59:07 +01:00
Florian Meißner a8b9275013 Update README.md 2024-06-15 12:29:36 +02:00
6 changed files with 85 additions and 28 deletions
+1 -1
View File
@@ -36,7 +36,7 @@ python3 run_server.py --port 9090 \
# running with custom model # running with custom model
python3 run_server.py --port 9090 \ python3 run_server.py --port 9090 \
--backend faster_whisper --backend faster_whisper \
-fw "/path/to/custom/faster/whisper/model" -fw "/path/to/custom/faster/whisper/model"
``` ```
+2 -1
View File
@@ -1,5 +1,5 @@
faster-whisper==1.0.1 faster-whisper==1.0.1
torch torch==2.3.0
websockets websockets
onnxruntime==1.16.0 onnxruntime==1.16.0
numba numba
@@ -10,3 +10,4 @@ ffmpeg-python
scipy scipy
jiwer jiwer
evaluate evaluate
numpy<2
+1 -1
View File
@@ -1 +1 @@
__version__ = "0.5.0" __version__ = "0.5.1"
+20 -4
View File
@@ -2,6 +2,7 @@ import os
import shutil import shutil
import wave import wave
import logging
import numpy as np import numpy as np
import pyaudio import pyaudio
import threading import threading
@@ -28,7 +29,8 @@ class Client:
translate=False, translate=False,
model="small", model="small",
srt_file_path="output.srt", srt_file_path="output.srt",
use_vad=True use_vad=True,
log_transcription=True
): ):
""" """
Initializes a Client instance for audio recording and streaming to a server. Initializes a Client instance for audio recording and streaming to a server.
@@ -56,11 +58,11 @@ class Client:
self.use_vad = use_vad self.use_vad = use_vad
self.last_segment = None self.last_segment = None
self.last_received_segment = None self.last_received_segment = None
self.log_transcription = log_transcription
if translate: if translate:
self.task = "translate" self.task = "translate"
self.timestamp_offset = 0.0
self.audio_bytes = None self.audio_bytes = None
if host is not None and port is not None: if host is not None and port is not None:
@@ -117,6 +119,7 @@ class Client:
self.last_response_received = time.time() self.last_response_received = time.time()
self.last_received_segment = segments[-1]["text"] self.last_received_segment = segments[-1]["text"]
if self.log_transcription:
# Truncate to last 3 entries for brevity. # Truncate to last 3 entries for brevity.
text = text[-3:] text = text[-3:]
utils.clear_screen() utils.clear_screen()
@@ -431,6 +434,8 @@ class TranscriptionTeeClient:
def handle_ffmpeg_process(self, process, stream_type): def handle_ffmpeg_process(self, process, stream_type):
print(f"[INFO]: Connecting to {stream_type} stream...") print(f"[INFO]: Connecting to {stream_type} stream...")
stderr_thread = threading.Thread(target=self.consume_stderr, args=(process,))
stderr_thread.start()
try: try:
# Process the stream # Process the stream
while True: while True:
@@ -477,6 +482,16 @@ class TranscriptionTeeClient:
return process 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): def save_chunk(self, n_audio_file):
""" """
Saves the current audio frames to a WAV file in a separate thread. Saves the current audio frames to a WAV file in a separate thread.
@@ -664,9 +679,10 @@ class TranscriptionClient(TranscriptionTeeClient):
use_vad=True, use_vad=True,
save_output_recording=False, save_output_recording=False,
output_recording_filename="./output_recording.wav", output_recording_filename="./output_recording.wav",
output_transcription_path="./output.srt" output_transcription_path="./output.srt",
log_transcription=True,
): ):
self.client = Client(host, port, lang, translate, model, srt_file_path=output_transcription_path, use_vad=use_vad) self.client = Client(host, port, lang, translate, model, srt_file_path=output_transcription_path, use_vad=use_vad, log_transcription=log_transcription)
if save_output_recording and not output_recording_filename.endswith(".wav"): if save_output_recording and not output_recording_filename.endswith(".wav"):
raise ValueError(f"Please provide a valid `output_recording_filename`: {output_recording_filename}") raise ValueError(f"Please provide a valid `output_recording_filename`: {output_recording_filename}")
if not output_transcription_path.endswith(".srt"): if not output_transcription_path.endswith(".srt"):
+49 -9
View File
@@ -4,6 +4,9 @@ import threading
import json import json
import functools import functools
import logging import logging
from enum import Enum
from typing import List, Optional
import torch import torch
import numpy as np import numpy as np
from websockets.sync.server import serve from websockets.sync.server import serve
@@ -121,6 +124,25 @@ class ClientManager:
return False return False
class BackendType(Enum):
FASTER_WHISPER = "faster_whisper"
TENSORRT = "tensorrt"
@staticmethod
def valid_types() -> List[str]:
return [backend_type.value for backend_type in BackendType]
@staticmethod
def is_valid(backend: str) -> bool:
return backend in BackendType.valid_types()
def is_faster_whisper(self) -> bool:
return self == BackendType.FASTER_WHISPER
def is_tensorrt(self) -> bool:
return self == BackendType.TENSORRT
class TranscriptionServer: class TranscriptionServer:
RATE = 16000 RATE = 16000
@@ -134,7 +156,9 @@ class TranscriptionServer:
self, websocket, options, faster_whisper_custom_model_path, self, websocket, options, faster_whisper_custom_model_path,
whisper_tensorrt_path, trt_multilingual whisper_tensorrt_path, trt_multilingual
): ):
if self.backend == "tensorrt": client: Optional[ServeClientBase] = None
if self.backend.is_tensorrt():
try: try:
client = ServeClientTensorRT( client = ServeClientTensorRT(
websocket, websocket,
@@ -155,9 +179,9 @@ class TranscriptionServer:
"message": "TensorRT-LLM not supported on Server yet. " "message": "TensorRT-LLM not supported on Server yet. "
"Reverting to available backend: 'faster_whisper'" "Reverting to available backend: 'faster_whisper'"
})) }))
self.backend = "faster_whisper" self.backend = BackendType.FASTER_WHISPER
if self.backend == "faster_whisper": if self.backend.is_faster_whisper():
if faster_whisper_custom_model_path is not None and os.path.exists(faster_whisper_custom_model_path): 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}") logging.info(f"Using custom model {faster_whisper_custom_model_path}")
options["model"] = faster_whisper_custom_model_path options["model"] = faster_whisper_custom_model_path
@@ -174,6 +198,9 @@ class TranscriptionServer:
) )
logging.info("Running faster_whisper backend.") logging.info("Running faster_whisper backend.")
if client is None:
raise ValueError(f"Backend type {self.backend.value} not recognised or not handled.")
self.client_manager.add_client(websocket, client) self.client_manager.add_client(websocket, client)
def get_audio_from_websocket(self, websocket): def get_audio_from_websocket(self, websocket):
@@ -202,7 +229,7 @@ class TranscriptionServer:
websocket.close() websocket.close()
return False # Indicates that the connection should not continue return False # Indicates that the connection should not continue
if self.backend == "tensorrt": 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, self.initialize_client(websocket, options, faster_whisper_custom_model_path,
whisper_tensorrt_path, trt_multilingual) whisper_tensorrt_path, trt_multilingual)
@@ -221,11 +248,11 @@ class TranscriptionServer:
frame_np = self.get_audio_from_websocket(websocket) frame_np = self.get_audio_from_websocket(websocket)
client = self.client_manager.get_client(websocket) client = self.client_manager.get_client(websocket)
if frame_np is False: if frame_np is False:
if self.backend == "tensorrt": if self.backend.is_tensorrt():
client.set_eos(True) client.set_eos(True)
return False return False
if self.backend == "tensorrt": if self.backend.is_tensorrt():
voice_active = self.voice_activity(websocket, frame_np) voice_active = self.voice_activity(websocket, frame_np)
if voice_active: if voice_active:
self.no_voice_activity_chunks = 0 self.no_voice_activity_chunks = 0
@@ -238,7 +265,7 @@ class TranscriptionServer:
def recv_audio(self, def recv_audio(self,
websocket, websocket,
backend="faster_whisper", backend: BackendType = BackendType.FASTER_WHISPER,
faster_whisper_custom_model_path=None, faster_whisper_custom_model_path=None,
whisper_tensorrt_path=None, whisper_tensorrt_path=None,
trt_multilingual=False): trt_multilingual=False):
@@ -311,10 +338,12 @@ class TranscriptionServer:
# TODO: load model initially # TODO: load model initially
else: else:
logging.info("Single model mode currently only works with custom models.") logging.info("Single model mode currently only works with custom models.")
if not BackendType.is_valid(backend):
raise ValueError(f"{backend} is not a valid backend type. Choose backend from {BackendType.valid_types()}")
with serve( with serve(
functools.partial( functools.partial(
self.recv_audio, self.recv_audio,
backend=backend, backend=BackendType(backend),
faster_whisper_custom_model_path=faster_whisper_custom_model_path, faster_whisper_custom_model_path=faster_whisper_custom_model_path,
whisper_tensorrt_path=whisper_tensorrt_path, whisper_tensorrt_path=whisper_tensorrt_path,
trt_multilingual=trt_multilingual trt_multilingual=trt_multilingual
@@ -758,9 +787,15 @@ class ServeClientFasterWhisper(ServeClientBase):
self.no_speech_thresh = 0.45 self.no_speech_thresh = 0.45
device = "cuda" if torch.cuda.is_available() else "cpu" 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: if self.model_size_or_path is None:
return return
logging.info(f"Using Device={device} with precision {self.compute_type}")
if single_model: if single_model:
if ServeClientFasterWhisper.SINGLE_MODEL is None: if ServeClientFasterWhisper.SINGLE_MODEL is None:
@@ -793,7 +828,7 @@ class ServeClientFasterWhisper(ServeClientBase):
self.transcriber = WhisperModel( self.transcriber = WhisperModel(
self.model_size_or_path, self.model_size_or_path,
device=device, device=device,
compute_type="int8" if device == "cpu" else "float16", compute_type=self.compute_type,
local_files_only=False, local_files_only=False,
) )
@@ -944,6 +979,7 @@ class ServeClientFasterWhisper(ServeClientBase):
input_bytes, duration = self.get_audio_chunk_for_processing() input_bytes, duration = self.get_audio_chunk_for_processing()
if duration < 1.0: if duration < 1.0:
time.sleep(0.1) # wait for audio chunks to arrive
continue continue
try: try:
input_sample = input_bytes.copy() input_sample = input_bytes.copy()
@@ -1002,6 +1038,8 @@ class ServeClientFasterWhisper(ServeClientBase):
""" """
offset = None offset = None
self.current_out = '' self.current_out = ''
last_segment = None
# process complete segments # process complete segments
if len(segments) > 1: if len(segments) > 1:
for i, s in enumerate(segments[:-1]): for i, s in enumerate(segments[:-1]):
@@ -1017,6 +1055,8 @@ class ServeClientFasterWhisper(ServeClientBase):
self.transcript.append(self.format_segment(start, end, text_)) self.transcript.append(self.format_segment(start, end, text_))
offset = min(duration, s.end) 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 self.current_out += segments[-1].text
last_segment = self.format_segment( last_segment = self.format_segment(
self.timestamp_offset + segments[-1].start, self.timestamp_offset + segments[-1].start,
+1 -1
View File
@@ -94,7 +94,7 @@ class VoiceActivityDetection():
return stacked.cpu() return stacked.cpu()
@staticmethod @staticmethod
def download(model_url="https://github.com/snakers4/silero-vad/raw/master/files/silero_vad.onnx"): def download(model_url="https://github.com/snakers4/silero-vad/raw/v4.0/files/silero_vad.onnx"):
target_dir = os.path.expanduser("~/.cache/whisper-live/") target_dir = os.path.expanduser("~/.cache/whisper-live/")
# Ensure the target directory exists # Ensure the target directory exists