Merge remote-tracking branch 'upstream/main'
This commit is contained in:
+52
-2
@@ -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
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user