Merge remote-tracking branch 'upstream/main'

This commit is contained in:
Sasa Trivic
2024-02-01 17:31:35 +01:00
2 changed files with 65 additions and 7 deletions
+52 -2
View File
@@ -13,6 +13,29 @@ import uuid
import time import time
def format_time(s):
"""Convert seconds (float) to SRT time format."""
hours = int(s // 3600)
minutes = int((s % 3600) // 60)
seconds = int(s % 60)
milliseconds = int((s - int(s)) * 1000)
return f"{hours:02}:{minutes:02}:{seconds:02},{milliseconds:03}"
def create_srt_file(segments, output_file):
with open(output_file, 'w', encoding='utf-8') as srt_file:
segment_number = 1
for segment in segments:
start_time = format_time(float(segment['start']))
end_time = format_time(float(segment['end']))
text = segment['text']
srt_file.write(f"{segment_number}\n")
srt_file.write(f"{start_time} --> {end_time}\n")
srt_file.write(f"{text}\n\n")
segment_number += 1
def resample(file: str, sr: int = 16000): def resample(file: str, sr: int = 16000):
""" """
# https://github.com/openai/whisper/blob/7858aa9c08d98f75575035ecd6481f462d66ca27/whisper/audio.py#L22 # https://github.com/openai/whisper/blob/7858aa9c08d98f75575035ecd6481f462d66ca27/whisper/audio.py#L22
@@ -57,6 +80,7 @@ class Client:
lang=None, lang=None,
translate=False, translate=False,
model="small", model="small",
srt_file_path="output.srt"
): ):
""" """
Initializes a Client instance for audio recording and streaming to a server. Initializes a Client instance for audio recording and streaming to a server.
@@ -89,6 +113,7 @@ class Client:
self.language = lang self.language = lang
self.model = model self.model = model
self.server_error = False self.server_error = False
self.srt_file_path = srt_file_path
if translate: if translate:
self.task = "translate" self.task = "translate"
@@ -127,6 +152,7 @@ class Client:
self.ws_thread.start() self.ws_thread.start()
self.frames = b"" self.frames = b""
self.transcript = []
print("[INFO]: * recording") print("[INFO]: * recording")
def on_message(self, ws, message): def on_message(self, ws, message):
@@ -166,6 +192,8 @@ class Client:
if "message" in message.keys() and message["message"] == "SERVER_READY": if "message" in message.keys() and message["message"] == "SERVER_READY":
self.recording = True self.recording = True
self.server_backend = message["backend"]
print(f"[INFO]: Server Running with backend {self.server_backend}")
return return
if "language" in message.keys(): if "language" in message.keys():
@@ -181,12 +209,21 @@ class Client:
message = message["segments"] message = message["segments"]
text = [] text = []
if len(message): n_segments = len(message)
for seg in message:
if n_segments:
for i, seg in enumerate(message):
if text and text[-1] == seg["text"]: if text and text[-1] == seg["text"]:
# already got it # already got it
continue continue
text.append(seg["text"]) text.append(seg["text"])
if i == n_segments-1:
self.last_segment = seg
elif self.server_backend == "faster_whisper":
if not len(self.transcript) or float(seg['start']) >= float(self.transcript[-1]['end']):
self.transcript.append(seg)
# keep only last 3 # keep only last 3
if len(text) > 3: if len(text) > 3:
text = text[-3:] text = text[-3:]
@@ -302,6 +339,9 @@ class Client:
assert self.last_response_recieved assert self.last_response_recieved
while time.time() - self.last_response_recieved < self.disconnect_if_no_response_for: while time.time() - self.last_response_recieved < self.disconnect_if_no_response_for:
continue continue
if self.server_backend == "faster_whisper":
self.write_srt_file(self.srt_file_path)
self.stream.close() self.stream.close()
self.close_websocket() self.close_websocket()
@@ -311,6 +351,8 @@ class Client:
self.stream.close() self.stream.close()
self.p.terminate() self.p.terminate()
self.close_websocket() self.close_websocket()
if self.server_backend == "faster_whisper":
self.write_srt_file(self.srt_file_path)
print("[INFO]: Keyboard interrupt.") print("[INFO]: Keyboard interrupt.")
def close_websocket(self): def close_websocket(self):
@@ -438,6 +480,8 @@ class Client:
t.start() t.start()
n_audio_file += 1 n_audio_file += 1
self.frames = b"" self.frames = b""
if self.server_backend == "faster_whisper":
self.write_srt_file(self.srt_file_path)
except KeyboardInterrupt: except KeyboardInterrupt:
if len(self.frames): if len(self.frames):
@@ -451,6 +495,8 @@ class Client:
self.close_websocket() self.close_websocket()
self.write_output_recording(n_audio_file, out_file) self.write_output_recording(n_audio_file, out_file)
if self.server_backend == "faster_whisper":
self.write_srt_file(self.srt_file_path)
def write_output_recording(self, n_audio_file, out_file): def write_output_recording(self, n_audio_file, out_file):
""" """
@@ -487,6 +533,10 @@ class Client:
os.remove(in_file) os.remove(in_file)
wavfile.close() wavfile.close()
def write_srt_file(self, output_path="output.srt"):
self.transcript.append(self.last_segment)
create_srt_file(self.transcript, output_path)
class TranscriptionClient: class TranscriptionClient:
""" """
+13 -5
View File
@@ -408,16 +408,17 @@ class ServeClientTensorRT(ServeClientBase):
json.dumps( json.dumps(
{ {
"uid": self.client_uid, "uid": self.client_uid,
"message": self.SERVER_READY "message": self.SERVER_READY,
"backend": "tensorrt"
} }
) )
) )
def warmup(self, warmup_steps=10): def warmup(self, warmup_steps=10):
logging.info("[INFO:] Warming up TensorRT engine..") logging.info("[INFO:] Warming up TensorRT engine..")
mel, duration = self.transcriber.log_mel_spectrogram("tests/jfk.flac") mel, _ = self.transcriber.log_mel_spectrogram("tests/jfk.flac")
for i in range(warmup_steps): for i in range(warmup_steps):
last_segment = self.transcriber.transcribe(mel) self.transcriber.transcribe(mel)
def set_eos(self, eos): def set_eos(self, eos):
self.lock.acquire() self.lock.acquire()
@@ -561,7 +562,7 @@ class ServeClientFasterWhisper(ServeClientBase):
multilingual=False, multilingual=False,
language=None, language=None,
client_uid=None, client_uid=None,
model="small", model="small.en",
initial_prompt=None, initial_prompt=None,
vad_parameters=None, vad_parameters=None,
): ):
@@ -595,6 +596,7 @@ class ServeClientFasterWhisper(ServeClientBase):
self.task = task self.task = task
self.initial_prompt = initial_prompt self.initial_prompt = initial_prompt
self.vad_parameters = vad_parameters or {"threshold": 0.5} self.vad_parameters = vad_parameters or {"threshold": 0.5}
self.no_speech_thresh = 0.45
device = "cuda" if torch.cuda.is_available() else "cpu" device = "cuda" if torch.cuda.is_available() else "cpu"
@@ -615,7 +617,8 @@ class ServeClientFasterWhisper(ServeClientBase):
json.dumps( json.dumps(
{ {
"uid": self.client_uid, "uid": self.client_uid,
"message": self.SERVER_READY "message": self.SERVER_READY,
"backend": "faster_whisper"
} }
) )
) )
@@ -729,6 +732,7 @@ class ServeClientFasterWhisper(ServeClientBase):
if time.time() - self.t_start > self.add_pause_thresh: if time.time() - self.t_start > self.add_pause_thresh:
self.text.append('') self.text.append('')
if not len(segments): continue
try: try:
self.websocket.send( self.websocket.send(
json.dumps({ json.dumps({
@@ -781,6 +785,10 @@ class ServeClientFasterWhisper(ServeClientBase):
text_ = s.text text_ = s.text
self.text.append(text_) self.text.append(text_)
start, end = self.timestamp_offset + s.start, self.timestamp_offset + min(duration, s.end) 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_)) self.transcript.append(self.format_segment(start, end, text_))
offset = min(duration, s.end) offset = min(duration, s.end)