update faster_whisper backend
This commit is contained in:
+27
-60
@@ -13,7 +13,6 @@ import torch
|
||||
import numpy as np
|
||||
import time
|
||||
from whisper_live.transcriber import WhisperModel
|
||||
from whisper_live.vad import VoiceActivityDetection
|
||||
|
||||
|
||||
class TranscriptionServer:
|
||||
@@ -35,8 +34,6 @@ class TranscriptionServer:
|
||||
|
||||
def __init__(self):
|
||||
# voice activity detection model
|
||||
self.vad_model = VoiceActivityDetection()
|
||||
self.vad_threshold = 0.4
|
||||
|
||||
self.clients = {}
|
||||
self.websockets = {}
|
||||
@@ -115,15 +112,6 @@ class TranscriptionServer:
|
||||
frame_data = websocket.recv()
|
||||
frame_np = np.frombuffer(frame_data, dtype=np.float32)
|
||||
|
||||
try:
|
||||
speech_prob = self.vad_model(torch.from_numpy(frame_np.copy()), self.RATE).item()
|
||||
if speech_prob < self.vad_threshold:
|
||||
continue
|
||||
|
||||
except Exception as e:
|
||||
logging.error(e)
|
||||
return
|
||||
|
||||
self.clients[websocket].add_frames(frame_np)
|
||||
|
||||
elapsed_time = time.time() - self.clients_start_time[websocket]
|
||||
@@ -322,25 +310,6 @@ class ServeClient:
|
||||
Exception: If there is an issue with audio processing or WebSocket communication.
|
||||
|
||||
"""
|
||||
# detect language
|
||||
if self.language is None:
|
||||
# wait for 30s of audio
|
||||
while self.frames_np is None or self.frames_np.shape[0] < 30*self.RATE:
|
||||
time.sleep(1)
|
||||
input_bytes = self.frames_np[-30*self.RATE:].copy()
|
||||
self.frames_np = None
|
||||
duration = input_bytes.shape[0] / self.RATE
|
||||
|
||||
self.language, lang_prob = self.transcriber.transcribe(
|
||||
input_bytes,
|
||||
initial_prompt=None,
|
||||
language=self.language,
|
||||
task=self.task
|
||||
)
|
||||
logging.info(f"Detected language {self.language} with probability {lang_prob}")
|
||||
self.websocket.send(json.dumps(
|
||||
{"uid": self.client_uid, "language": self.language, "language_prob": lang_prob}))
|
||||
|
||||
while True:
|
||||
if self.exit:
|
||||
logging.info("Exiting speech to text thread")
|
||||
@@ -358,24 +327,31 @@ class ServeClient:
|
||||
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
|
||||
if duration<1.0:
|
||||
if duration<1.0:
|
||||
continue
|
||||
try:
|
||||
input_sample = input_bytes.copy()
|
||||
# set previous complete segment as initial prompt
|
||||
if len(self.text) and self.text[-1] != '':
|
||||
initial_prompt = self.text[-1]
|
||||
else:
|
||||
initial_prompt = None
|
||||
|
||||
# whisper transcribe with prompt
|
||||
result = self.transcriber.transcribe(
|
||||
result, info = self.transcriber.transcribe(
|
||||
input_sample,
|
||||
initial_prompt=initial_prompt,
|
||||
initial_prompt=None,
|
||||
language=self.language,
|
||||
task=self.task
|
||||
task=self.task,
|
||||
vad_filter=True,
|
||||
vad_parameters={"threshold": 0.5}
|
||||
)
|
||||
|
||||
if self.language is None:
|
||||
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}))
|
||||
else:
|
||||
# detect language again
|
||||
continue
|
||||
|
||||
if len(result):
|
||||
self.t_start = None
|
||||
last_segment = self.update_segments(result, duration)
|
||||
@@ -384,17 +360,7 @@ class ServeClient:
|
||||
else:
|
||||
segments = self.transcript[-self.send_last_n_segments:]
|
||||
if last_segment is not None:
|
||||
segments = segments + [last_segment]
|
||||
|
||||
try:
|
||||
self.websocket.send(
|
||||
json.dumps({
|
||||
"uid": self.client_uid,
|
||||
"segments": segments
|
||||
})
|
||||
)
|
||||
except Exception as e:
|
||||
logging.error(f"[ERROR]: {e}")
|
||||
segments = segments + [last_segment]
|
||||
else:
|
||||
# show previous output if there is pause i.e. no output from whisper
|
||||
segments = []
|
||||
@@ -410,15 +376,16 @@ class ServeClient:
|
||||
if time.time() - self.t_start > self.add_pause_thresh:
|
||||
self.text.append('')
|
||||
|
||||
try:
|
||||
self.websocket.send(
|
||||
json.dumps({
|
||||
"uid": self.client_uid,
|
||||
"segments": segments
|
||||
})
|
||||
)
|
||||
except Exception as e:
|
||||
logging.error(f"[ERROR]: {e}")
|
||||
try:
|
||||
self.websocket.send(
|
||||
json.dumps({
|
||||
"uid": self.client_uid,
|
||||
"segments": segments
|
||||
})
|
||||
)
|
||||
except Exception as e:
|
||||
logging.error(f"[ERROR]: {e}")
|
||||
|
||||
except Exception as e:
|
||||
logging.error(f"[ERROR]: {e}")
|
||||
time.sleep(0.01)
|
||||
|
||||
Reference in New Issue
Block a user