Merge pull request #146 from makaveli10/code_formatting
Code formatting
This commit is contained in:
@@ -11,12 +11,11 @@ on:
|
||||
types: [opened, synchronize, reopened]
|
||||
|
||||
jobs:
|
||||
test:
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 60
|
||||
run-tests:
|
||||
runs-on: ubuntu-22.04
|
||||
strategy:
|
||||
matrix:
|
||||
python-version: [3.8, 3.9, '3.10', '3.11']
|
||||
python-version: [3.8, 3.9, '3.10', 3.11]
|
||||
steps:
|
||||
- uses: actions/checkout@v2
|
||||
|
||||
@@ -48,10 +47,36 @@ jobs:
|
||||
run: |
|
||||
echo "Running tests with Python ${{ matrix.python-version }}"
|
||||
python -m unittest discover -s tests
|
||||
|
||||
check-code-format:
|
||||
runs-on: ubuntu-22.04
|
||||
strategy:
|
||||
matrix:
|
||||
python-version: [3.8, 3.9, '3.10', 3.11]
|
||||
|
||||
build-and-push:
|
||||
needs: test
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v2
|
||||
|
||||
- name: Set up Python ${{ matrix.python-version }}
|
||||
uses: actions/setup-python@v2
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
python -m pip install flake8
|
||||
|
||||
- name: Lint with flake8
|
||||
run: |
|
||||
# stop the build if there are Python syntax errors or undefined names
|
||||
flake8 . --count --select=E9,F63,F7,F82 --show-source --statistics
|
||||
# exit-zero treats all errors as warnings. The GitHub editor is 127 chars wide
|
||||
flake8 . --count --exit-zero --max-complexity=10 --max-line-length=127 --statistics
|
||||
|
||||
publish-to-pypi:
|
||||
needs: [run-tests, check-code-format]
|
||||
runs-on: ubuntu-22.04
|
||||
if: github.event_name == 'push' && startsWith(github.ref, 'refs/tags')
|
||||
steps:
|
||||
- uses: actions/checkout@v2
|
||||
|
||||
+5
-5
@@ -4,15 +4,15 @@ from whisper_live.server import TranscriptionServer
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument('--port', '-p',
|
||||
type=int,
|
||||
type=int,
|
||||
default=9090,
|
||||
help="Websocket port to run the server on.")
|
||||
parser.add_argument('--backend', '-b',
|
||||
type=str,
|
||||
default='faster_whisper',
|
||||
type=str,
|
||||
default='faster_whisper',
|
||||
help='Backends from ["tensorrt", "faster_whisper"]')
|
||||
parser.add_argument('--faster_whisper_custom_model_path', '-fw',
|
||||
type=str, default=None,
|
||||
type=str, default=None,
|
||||
help="Custom Faster Whisper Model")
|
||||
parser.add_argument('--trt_model_path', '-trt',
|
||||
type=str,
|
||||
@@ -30,7 +30,7 @@ if __name__ == "__main__":
|
||||
server = TranscriptionServer()
|
||||
server.run(
|
||||
"0.0.0.0",
|
||||
port=args.port,
|
||||
port=args.port,
|
||||
backend=args.backend,
|
||||
faster_whisper_custom_model_path=args.faster_whisper_custom_model_path,
|
||||
whisper_tensorrt_path=args.trt_model_path,
|
||||
|
||||
@@ -10,36 +10,38 @@ HERE = pathlib.Path(__file__).parent
|
||||
README = (HERE / "README.md").read_text()
|
||||
|
||||
# This call to setup() does all the work
|
||||
setup(name="whisper-live",
|
||||
version=__version__,
|
||||
description="A nearly-live implementation of OpenAI's Whisper.",
|
||||
long_description=README,
|
||||
long_description_content_type="text/markdown",
|
||||
include_package_data=True,
|
||||
url="https://github.com/collabora/WhisperLive",
|
||||
author="Collabora Ltd",
|
||||
author_email="vineet.suryan@collabora.com",
|
||||
license="MIT",
|
||||
classifiers=[
|
||||
"Development Status :: 4 - Beta",
|
||||
"Intended Audience :: Developers",
|
||||
"Intended Audience :: Science/Research",
|
||||
"License :: OSI Approved :: MIT License",
|
||||
"Programming Language :: Python :: 3",
|
||||
"Programming Language :: Python :: 3 :: Only",
|
||||
"Programming Language :: Python :: 3.8",
|
||||
"Programming Language :: Python :: 3.9",
|
||||
"Topic :: Scientific/Engineering :: Artificial Intelligence",
|
||||
],
|
||||
packages=find_packages(
|
||||
exclude=("examples",
|
||||
"Audio-Transcription-Chrome",
|
||||
"Audio-Transcription-Firefox",
|
||||
"requirements",
|
||||
"whisper-finetuning"
|
||||
)
|
||||
),
|
||||
install_requires=[
|
||||
setup(
|
||||
name="whisper-live",
|
||||
version=__version__,
|
||||
description="A nearly-live implementation of OpenAI's Whisper.",
|
||||
long_description=README,
|
||||
long_description_content_type="text/markdown",
|
||||
include_package_data=True,
|
||||
url="https://github.com/collabora/WhisperLive",
|
||||
author="Collabora Ltd",
|
||||
author_email="vineet.suryan@collabora.com",
|
||||
license="MIT",
|
||||
classifiers=[
|
||||
"Development Status :: 4 - Beta",
|
||||
"Intended Audience :: Developers",
|
||||
"Intended Audience :: Science/Research",
|
||||
"License :: OSI Approved :: MIT License",
|
||||
"Programming Language :: Python :: 3",
|
||||
"Programming Language :: Python :: 3 :: Only",
|
||||
"Programming Language :: Python :: 3.8",
|
||||
"Programming Language :: Python :: 3.9",
|
||||
"Topic :: Scientific/Engineering :: Artificial Intelligence",
|
||||
],
|
||||
packages=find_packages(
|
||||
exclude=(
|
||||
"examples",
|
||||
"Audio-Transcription-Chrome",
|
||||
"Audio-Transcription-Firefox",
|
||||
"requirements",
|
||||
"whisper-finetuning"
|
||||
)
|
||||
),
|
||||
install_requires=[
|
||||
"PyAudio",
|
||||
"faster-whisper==0.10.0",
|
||||
"torch",
|
||||
@@ -53,6 +55,6 @@ setup(name="whisper-live",
|
||||
"openai-whisper",
|
||||
"kaldialign",
|
||||
"soundfile",
|
||||
],
|
||||
python_requires=">=3.8"
|
||||
],
|
||||
python_requires=">=3.8"
|
||||
)
|
||||
|
||||
Binary file not shown.
@@ -4,7 +4,8 @@ import scipy
|
||||
import websocket
|
||||
import unittest
|
||||
from unittest.mock import patch, MagicMock
|
||||
from whisper_live.client import TranscriptionClient, resample
|
||||
from whisper_live.client import TranscriptionClient
|
||||
from whisper_live.utils import resample
|
||||
|
||||
|
||||
class BaseTestCase(unittest.TestCase):
|
||||
@@ -68,7 +69,7 @@ class TestClientCallbacks(BaseTestCase):
|
||||
]
|
||||
})
|
||||
self.client.on_message(self.mock_ws_app, message)
|
||||
|
||||
|
||||
# Assert that the transcript was updated correctly
|
||||
self.assertEqual(len(self.client.transcript), 2)
|
||||
self.assertEqual(self.client.transcript[1]['text'], "Test transcript 2")
|
||||
@@ -79,14 +80,14 @@ class TestClientCallbacks(BaseTestCase):
|
||||
self.client.on_close(self.mock_ws_app, close_status_code, close_msg)
|
||||
|
||||
self.assertFalse(self.client.recording)
|
||||
self.assertFalse(self.client.server_error)
|
||||
self.assertFalse(self.client.server_error)
|
||||
self.assertFalse(self.client.waiting)
|
||||
|
||||
|
||||
def test_on_error(self):
|
||||
error_message = "Test Error"
|
||||
self.client.on_error(self.mock_ws_app, error_message)
|
||||
|
||||
self.assertTrue(self.client.server_error)
|
||||
self.assertTrue(self.client.server_error)
|
||||
self.assertEqual(self.client.error_message, error_message)
|
||||
|
||||
|
||||
@@ -95,10 +96,10 @@ class TestAudioResampling(unittest.TestCase):
|
||||
original_audio = "assets/jfk.flac"
|
||||
expected_sr = 16000
|
||||
resampled_audio = resample(original_audio, expected_sr)
|
||||
|
||||
|
||||
sr, _ = scipy.io.wavfile.read(resampled_audio)
|
||||
self.assertEqual(sr, expected_sr)
|
||||
|
||||
|
||||
os.remove(resampled_audio)
|
||||
|
||||
|
||||
|
||||
+54
-22
@@ -6,6 +6,8 @@ from unittest import mock
|
||||
|
||||
import numpy as np
|
||||
import evaluate
|
||||
|
||||
from websockets.exceptions import ConnectionClosed
|
||||
from whisper_live.server import TranscriptionServer
|
||||
from whisper_live.client import TranscriptionClient
|
||||
from whisper.normalizers import EnglishTextNormalizer
|
||||
@@ -14,32 +16,31 @@ from whisper.normalizers import EnglishTextNormalizer
|
||||
class TestTranscriptionServerInitialization(unittest.TestCase):
|
||||
def test_initialization(self):
|
||||
server = TranscriptionServer()
|
||||
self.assertEqual(server.max_clients, 4)
|
||||
self.assertEqual(server.max_connection_time, 600)
|
||||
self.assertDictEqual(server.clients, {})
|
||||
self.assertDictEqual(server.websockets, {})
|
||||
self.assertDictEqual(server.clients_start_time, {})
|
||||
self.assertEqual(server.client_manager.max_clients, 4)
|
||||
self.assertEqual(server.client_manager.max_connection_time, 600)
|
||||
self.assertDictEqual(server.client_manager.clients, {})
|
||||
self.assertDictEqual(server.client_manager.start_times, {})
|
||||
|
||||
|
||||
class TestGetWaitTime(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.server = TranscriptionServer()
|
||||
self.server.clients_start_time = {
|
||||
self.server.client_manager.start_times = {
|
||||
'client1': time.time() - 120,
|
||||
'client2': time.time() - 300
|
||||
}
|
||||
self.server.max_connection_time = 600
|
||||
self.server.client_manager.max_connection_time = 600
|
||||
|
||||
def test_get_wait_time(self):
|
||||
expected_wait_time = (600 - (time.time() - self.server.clients_start_time['client2'])) / 60
|
||||
print(self.server.get_wait_time(), expected_wait_time)
|
||||
self.assertAlmostEqual(self.server.get_wait_time(), expected_wait_time, places=2)
|
||||
expected_wait_time = (600 - (time.time() - self.server.client_manager.start_times['client2'])) / 60
|
||||
print(self.server.client_manager.get_wait_time(), expected_wait_time)
|
||||
self.assertAlmostEqual(self.server.client_manager.get_wait_time(), expected_wait_time, places=2)
|
||||
|
||||
|
||||
|
||||
class TestServerConnection(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.server = TranscriptionServer()
|
||||
|
||||
|
||||
@mock.patch('websockets.WebSocketCommonProtocol')
|
||||
def test_connection(self, mock_websocket):
|
||||
mock_websocket.recv.return_value = json.dumps({
|
||||
@@ -50,7 +51,6 @@ class TestServerConnection(unittest.TestCase):
|
||||
})
|
||||
self.server.recv_audio(mock_websocket, "faster_whisper")
|
||||
|
||||
|
||||
@mock.patch('websockets.WebSocketCommonProtocol')
|
||||
def test_recv_audio_exception_handling(self, mock_websocket):
|
||||
mock_websocket.recv.side_effect = [json.dumps({
|
||||
@@ -58,12 +58,12 @@ class TestServerConnection(unittest.TestCase):
|
||||
'language': 'en',
|
||||
'task': 'transcribe',
|
||||
'model': 'tiny.en'
|
||||
}), np.array([1, 2, 3]).tobytes()]
|
||||
|
||||
}), np.array([1, 2, 3]).tobytes()]
|
||||
|
||||
with self.assertLogs(level="ERROR"):
|
||||
self.server.recv_audio(mock_websocket, "faster_whisper")
|
||||
|
||||
self.assertNotIn(mock_websocket, self.server.clients)
|
||||
|
||||
self.assertNotIn(mock_websocket, self.server.client_manager.clients)
|
||||
|
||||
|
||||
class TestServerInferenceAccuracy(unittest.TestCase):
|
||||
@@ -71,12 +71,12 @@ class TestServerInferenceAccuracy(unittest.TestCase):
|
||||
def setUpClass(cls):
|
||||
cls.server_process = subprocess.Popen(["python", "run_server.py"]) # Adjust the command as needed
|
||||
time.sleep(2)
|
||||
|
||||
|
||||
@classmethod
|
||||
def tearDownClass(cls):
|
||||
cls.server_process.terminate()
|
||||
cls.server_process.wait()
|
||||
|
||||
|
||||
@mock.patch('pyaudio.PyAudio')
|
||||
def setUp(self, mock_pyaudio):
|
||||
self.mock_pyaudio = mock_pyaudio.return_value
|
||||
@@ -84,16 +84,16 @@ class TestServerInferenceAccuracy(unittest.TestCase):
|
||||
self.mock_pyaudio.open.return_value = self.mock_stream
|
||||
self.metric = evaluate.load("wer")
|
||||
self.normalizer = EnglishTextNormalizer()
|
||||
self.client = TranscriptionClient(
|
||||
self.client = TranscriptionClient(
|
||||
"localhost", "9090", model="base.en", lang="en",
|
||||
)
|
||||
|
||||
|
||||
def test_inference(self):
|
||||
gt = "And so my fellow Americans, ask not, what your country can do for you. Ask what you can do for your country!"
|
||||
self.client("assets/jfk.flac")
|
||||
with open("output.srt", "r") as f:
|
||||
lines = f.readlines()
|
||||
prediction = " ".join([l.strip() for l in lines[2::4]])
|
||||
prediction = " ".join([line.strip() for line in lines[2::4]])
|
||||
prediction_normalized = self.normalizer(prediction)
|
||||
gt_normalized = self.normalizer(gt)
|
||||
|
||||
@@ -103,3 +103,35 @@ class TestServerInferenceAccuracy(unittest.TestCase):
|
||||
references=[gt_normalized]
|
||||
)
|
||||
self.assertLess(wer, 0.05)
|
||||
|
||||
|
||||
class TestExceptionHandling(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.server = TranscriptionServer()
|
||||
|
||||
@mock.patch('websockets.WebSocketCommonProtocol')
|
||||
def test_connection_closed_exception(self, mock_websocket):
|
||||
mock_websocket.recv.side_effect = ConnectionClosed(1001, "testing connection closed")
|
||||
|
||||
with self.assertLogs(level="INFO") as log:
|
||||
self.server.recv_audio(mock_websocket, "faster_whisper")
|
||||
self.assertTrue(any("Connection closed by client" in message for message in log.output))
|
||||
|
||||
@mock.patch('websockets.WebSocketCommonProtocol')
|
||||
def test_json_decode_exception(self, mock_websocket):
|
||||
mock_websocket.recv.return_value = "invalid json"
|
||||
|
||||
with self.assertLogs(level="ERROR") as log:
|
||||
self.server.recv_audio(mock_websocket, "faster_whisper")
|
||||
self.assertTrue(any("Failed to decode JSON from client" in message for message in log.output))
|
||||
|
||||
@mock.patch('websockets.WebSocketCommonProtocol')
|
||||
def test_unexpected_exception_handling(self, mock_websocket):
|
||||
mock_websocket.recv.side_effect = RuntimeError("Unexpected error")
|
||||
|
||||
with self.assertLogs(level="ERROR") as log:
|
||||
self.server.recv_audio(mock_websocket, "faster_whisper")
|
||||
for message in log.output:
|
||||
print(message)
|
||||
print()
|
||||
self.assertTrue(any("Unexpected error: Unexpected error" in message for message in log.output))
|
||||
|
||||
+7
-9
@@ -1,14 +1,12 @@
|
||||
import unittest
|
||||
import numpy as np
|
||||
import torch
|
||||
import scipy.io as sio
|
||||
from whisper_live.tensorrt_utils import load_audio
|
||||
from whisper_live.vad import VoiceActivityDetection
|
||||
from whisper_live.vad import VoiceActivityDetector
|
||||
|
||||
|
||||
class TestVoiceActivityDetection(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.vad = VoiceActivityDetection()
|
||||
self.vad = VoiceActivityDetector()
|
||||
self.sample_rate = 16000
|
||||
|
||||
def generate_silence(self, duration_seconds):
|
||||
@@ -19,10 +17,10 @@ class TestVoiceActivityDetection(unittest.TestCase):
|
||||
|
||||
def test_vad_silence_detection(self):
|
||||
silence = self.generate_silence(3)
|
||||
speech_prob = self.vad(torch.from_numpy(silence.copy()), self.sample_rate).item()
|
||||
self.assertLess(speech_prob, 0.5, "VAD incorrectly identified silence as speech.")
|
||||
is_speech_present = self.vad(silence.copy())
|
||||
self.assertFalse(is_speech_present, "VAD incorrectly identified silence as speech.")
|
||||
|
||||
def test_vad_speech_detection(self):
|
||||
audio_tensor = torch.from_numpy(load_audio("assets/jfk.flac"))
|
||||
speech_prob = self.vad(audio_tensor, self.sample_rate).item()
|
||||
self.assertGreater(speech_prob, 0.5, "VAD failed to identify speech segment.")
|
||||
audio_tensor = load_audio("assets/jfk.flac")
|
||||
is_speech_present = self.vad(audio_tensor)
|
||||
self.assertTrue(is_speech_present, "VAD failed to identify speech segment.")
|
||||
|
||||
@@ -1 +1 @@
|
||||
__version__="0.1.0"
|
||||
__version__ = "0.1.0"
|
||||
|
||||
+58
-125
@@ -2,68 +2,14 @@ import os
|
||||
import wave
|
||||
|
||||
import numpy as np
|
||||
import scipy
|
||||
import ffmpeg
|
||||
import pyaudio
|
||||
import threading
|
||||
import textwrap
|
||||
import json
|
||||
import websocket
|
||||
import uuid
|
||||
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):
|
||||
"""
|
||||
# https://github.com/openai/whisper/blob/7858aa9c08d98f75575035ecd6481f462d66ca27/whisper/audio.py#L22
|
||||
Open an audio file and read as mono waveform, resampling as necessary,
|
||||
save the resampled audio
|
||||
|
||||
Args:
|
||||
file (str): The audio file to open
|
||||
sr (int): The sample rate to resample the audio if necessary
|
||||
|
||||
Returns:
|
||||
resampled_file (str): The resampled audio file
|
||||
"""
|
||||
try:
|
||||
# This launches a subprocess to decode audio while down-mixing and resampling as necessary.
|
||||
# Requires the ffmpeg CLI and `ffmpeg-python` package to be installed.
|
||||
out, _ = (
|
||||
ffmpeg.input(file, threads=0)
|
||||
.output("-", format="s16le", acodec="pcm_s16le", ac=1, ar=sr)
|
||||
.run(cmd=["ffmpeg", "-nostdin"], capture_stdout=True, capture_stderr=True)
|
||||
)
|
||||
except ffmpeg.Error as e:
|
||||
raise RuntimeError(f"Failed to load audio: {e.stderr.decode()}") from e
|
||||
np_buffer = np.frombuffer(out, dtype=np.int16)
|
||||
|
||||
resampled_file = f"{file.split('.')[0]}_resampled.wav"
|
||||
scipy.io.wavfile.write(resampled_file, sr, np_buffer.astype(np.int16))
|
||||
return resampled_file
|
||||
import ffmpeg
|
||||
import whisper_live.utils as utils
|
||||
|
||||
|
||||
class Client:
|
||||
@@ -150,10 +96,40 @@ class Client:
|
||||
self.transcript = []
|
||||
print("[INFO]: * recording")
|
||||
|
||||
def handle_status_messages(self, message_data):
|
||||
"""Handles server status messages."""
|
||||
status = message_data["status"]
|
||||
if status == "WAIT":
|
||||
self.waiting = True
|
||||
print(f"[INFO]: Server is full. Estimated wait time {round(message_data['message'])} minutes.")
|
||||
elif status == "ERROR":
|
||||
print(f"Message from Server: {message_data['message']}")
|
||||
self.server_error = True
|
||||
elif status == "WARNING":
|
||||
print(f"Message from Server: {message_data['message']}")
|
||||
|
||||
def process_segments(self, segments):
|
||||
"""Processes transcript segments."""
|
||||
text = []
|
||||
for i, seg in enumerate(segments):
|
||||
if not text or text[-1] != seg["text"]:
|
||||
text.append(seg["text"])
|
||||
if i == len(segments) - 1:
|
||||
self.last_segment = seg
|
||||
elif (self.server_backend == "faster_whisper" and
|
||||
(not self.transcript or
|
||||
float(seg['start']) >= float(self.transcript[-1]['end']))):
|
||||
self.transcript.append(seg)
|
||||
|
||||
# Truncate to last 3 entries for brevity.
|
||||
text = text[-3:]
|
||||
utils.clear_screen()
|
||||
utils.print_transcript(text)
|
||||
|
||||
def on_message(self, ws, message):
|
||||
"""
|
||||
Callback function called when a message is received from the server.
|
||||
|
||||
|
||||
It updates various attributes of the client based on the received message, including
|
||||
recording status, language detection, and server messages. If a disconnect message
|
||||
is received, it sets the recording status to False.
|
||||
@@ -171,18 +147,11 @@ class Client:
|
||||
return
|
||||
|
||||
if "status" in message.keys():
|
||||
if message["status"] == "WAIT":
|
||||
self.waiting = True
|
||||
print(
|
||||
f"[INFO]:Server is full. Estimated wait time {round(message['message'])} minutes."
|
||||
)
|
||||
elif message["status"] == "ERROR":
|
||||
print(f"Message from Server: {message['message']}")
|
||||
self.server_error = True
|
||||
self.handle_status_messages(message)
|
||||
return
|
||||
|
||||
if "message" in message.keys() and message["message"] == "DISCONNECT":
|
||||
print("[INFO]: Server overtime disconnected.")
|
||||
print("[INFO]: Server disconnected due to overtime.")
|
||||
self.recording = False
|
||||
|
||||
if "message" in message.keys() and message["message"] == "SERVER_READY":
|
||||
@@ -199,38 +168,8 @@ class Client:
|
||||
)
|
||||
return
|
||||
|
||||
if "segments" not in message.keys():
|
||||
return
|
||||
|
||||
message = message["segments"]
|
||||
text = []
|
||||
n_segments = len(message)
|
||||
|
||||
if n_segments:
|
||||
for i, seg in enumerate(message):
|
||||
if text and text[-1] == seg["text"]:
|
||||
# already got it
|
||||
continue
|
||||
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
|
||||
if len(text) > 3:
|
||||
text = text[-3:]
|
||||
wrapper = textwrap.TextWrapper(width=60)
|
||||
word_list = wrapper.wrap(text="".join(text))
|
||||
# Print each line.
|
||||
if os.name == "nt":
|
||||
os.system("cls")
|
||||
else:
|
||||
os.system("clear")
|
||||
for element in word_list:
|
||||
print(element)
|
||||
if "segments" in message.keys():
|
||||
self.process_segments(message["segments"])
|
||||
|
||||
def on_error(self, ws, error):
|
||||
print(f"[ERROR] WebSocket Error: {error}")
|
||||
@@ -246,7 +185,7 @@ class Client:
|
||||
def on_open(self, ws):
|
||||
"""
|
||||
Callback function called when the WebSocket connection is successfully opened.
|
||||
|
||||
|
||||
Sends an initial configuration message to the server, including client UID,
|
||||
language selection, and task type.
|
||||
|
||||
@@ -270,8 +209,8 @@ class Client:
|
||||
def bytes_to_float_array(audio_bytes):
|
||||
"""
|
||||
Convert audio data from bytes to a NumPy float array.
|
||||
|
||||
It assumes that the audio data is in 16-bit PCM format. The audio data is normalized to
|
||||
|
||||
It assumes that the audio data is in 16-bit PCM format. The audio data is normalized to
|
||||
have values between -1 and 1.
|
||||
|
||||
Args:
|
||||
@@ -299,10 +238,10 @@ class Client:
|
||||
def play_file(self, filename):
|
||||
"""
|
||||
Play an audio file and send it to the server for processing.
|
||||
|
||||
|
||||
Reads an audio file, plays it through the audio output, and simultaneously sends
|
||||
the audio data to the server for processing. It uses PyAudio to create an audio
|
||||
stream for playback. The audio data is read from the file in chunks, converted to
|
||||
the audio data to the server for processing. It uses PyAudio to create an audio
|
||||
stream for playback. The audio data is read from the file in chunks, converted to
|
||||
floating-point format, and sent to the server using WebSocket communication.
|
||||
This method is typically used when you want to process pre-recorded audio and send it
|
||||
to the server in real-time.
|
||||
@@ -310,7 +249,7 @@ class Client:
|
||||
Args:
|
||||
filename (str): The path to the audio file to be played and sent to the server.
|
||||
"""
|
||||
|
||||
|
||||
# read audio and create pyaudio stream
|
||||
with wave.open(filename, "rb") as wavfile:
|
||||
self.stream = self.p.open(
|
||||
@@ -356,7 +295,7 @@ class Client:
|
||||
"""
|
||||
Close the WebSocket connection and join the WebSocket thread.
|
||||
|
||||
First attempts to close the WebSocket connection using `self.client_socket.close()`. After
|
||||
First attempts to close the WebSocket connection using `self.client_socket.close()`. After
|
||||
closing the connection, it joins the WebSocket thread to ensure proper termination.
|
||||
|
||||
"""
|
||||
@@ -383,7 +322,7 @@ class Client:
|
||||
"""
|
||||
Write audio frames to a WAV file.
|
||||
|
||||
The WAV file is created or overwritten with the specified name. The audio frames should be
|
||||
The WAV file is created or overwritten with the specified name. The audio frames should be
|
||||
in the correct format and match the specified channel, sample width, and sample rate.
|
||||
|
||||
Args:
|
||||
@@ -433,7 +372,6 @@ class Client:
|
||||
|
||||
print("[INFO]: HLS stream processing finished.")
|
||||
|
||||
|
||||
def record(self, out_file="output_recording.wav"):
|
||||
"""
|
||||
Record audio data from the input stream and save it to a WAV file.
|
||||
@@ -444,11 +382,12 @@ class Client:
|
||||
|
||||
Audio data is saved in chunks to the "chunks" directory. Each chunk is saved as a separate WAV file.
|
||||
The recording will continue until the specified duration is reached or until the `RECORDING` flag is set to `False`.
|
||||
The recording process can be interrupted by sending a KeyboardInterrupt (e.g., pressing Ctrl+C). After recording,
|
||||
The recording process can be interrupted by sending a KeyboardInterrupt (e.g., pressing Ctrl+C). After recording,
|
||||
the method combines all the saved audio chunks into the specified `out_file`.
|
||||
|
||||
Args:
|
||||
out_file (str, optional): The name of the output WAV file to save the entire recording. Default is "output_recording.wav".
|
||||
out_file (str, optional): The name of the output WAV file to save the entire recording.
|
||||
Default is "output_recording.wav".
|
||||
|
||||
"""
|
||||
n_audio_file = 0
|
||||
@@ -458,7 +397,7 @@ class Client:
|
||||
for _ in range(0, int(self.rate / self.chunk * self.record_seconds)):
|
||||
if not self.recording:
|
||||
break
|
||||
data = self.stream.read(self.chunk, exception_on_overflow = False)
|
||||
data = self.stream.read(self.chunk, exception_on_overflow=False)
|
||||
self.frames += data
|
||||
|
||||
audio_array = Client.bytes_to_float_array(data)
|
||||
@@ -498,8 +437,8 @@ class Client:
|
||||
def write_output_recording(self, n_audio_file, out_file):
|
||||
"""
|
||||
Combine and save recorded audio chunks into a single WAV file.
|
||||
|
||||
The individual audio chunk files are expected to be located in the "chunks" directory. Reads each chunk
|
||||
|
||||
The individual audio chunk files are expected to be located in the "chunks" directory. Reads each chunk
|
||||
file, appends its audio data to the final recording, and then deletes the chunk file. After combining
|
||||
and saving, the final recording is stored in the specified `out_file`.
|
||||
|
||||
@@ -532,7 +471,7 @@ class Client:
|
||||
|
||||
def write_srt_file(self, output_path="output.srt"):
|
||||
self.transcript.append(self.last_segment)
|
||||
create_srt_file(self.transcript, output_path)
|
||||
utils.create_srt_file(self.transcript, output_path)
|
||||
|
||||
|
||||
class TranscriptionClient:
|
||||
@@ -558,13 +497,7 @@ class TranscriptionClient:
|
||||
transcription_client()
|
||||
```
|
||||
"""
|
||||
def __init__(self,
|
||||
host,
|
||||
port,
|
||||
lang=None,
|
||||
translate=False,
|
||||
model="small",
|
||||
):
|
||||
def __init__(self, host, port, lang=None, translate=False, model="small"):
|
||||
self.client = Client(host, port, lang, translate, model)
|
||||
|
||||
def __call__(self, audio=None, hls_url=None):
|
||||
@@ -572,12 +505,12 @@ class TranscriptionClient:
|
||||
Start the transcription process.
|
||||
|
||||
Initiates the transcription process by connecting to the server via a WebSocket. It waits for the server
|
||||
to be ready to receive audio data and then sends audio for transcription. If an audio file is provided, it
|
||||
to be ready to receive audio data and then sends audio for transcription. If an audio file is provided, it
|
||||
will be played and streamed to the server; otherwise, it will perform live recording.
|
||||
|
||||
Args:
|
||||
audio (str, optional): Path to an audio file for transcription. Default is None, which triggers live recording.
|
||||
|
||||
|
||||
"""
|
||||
print("[INFO]: Waiting for server ready ...")
|
||||
while not self.client.recording:
|
||||
@@ -589,7 +522,7 @@ class TranscriptionClient:
|
||||
if hls_url is not None:
|
||||
self.client.process_hls_stream(hls_url)
|
||||
elif audio is not None:
|
||||
resampled_file = resample(audio)
|
||||
resampled_file = utils.resample(audio)
|
||||
self.client.play_file(resampled_file)
|
||||
else:
|
||||
self.client.record()
|
||||
self.client.record()
|
||||
|
||||
+560
-412
File diff suppressed because it is too large
Load Diff
@@ -214,7 +214,7 @@ def store_transcripts(filename: Pathlike, texts: Iterable[Tuple[str, str,
|
||||
print(f"{cut_id}:\thyp={hyp}", file=f)
|
||||
|
||||
|
||||
def write_error_stats(
|
||||
def write_error_stats( # noqa: C901
|
||||
f: TextIO,
|
||||
test_set_name: str,
|
||||
results: List[Tuple[str, str]],
|
||||
@@ -362,4 +362,4 @@ def write_error_stats(
|
||||
hyp_count = corr + hyp_sub + ins
|
||||
|
||||
print(f"{word} {corr} {tot_errs} {ref_count} {hyp_count}", file=f)
|
||||
return float(tot_err_rate)
|
||||
return float(tot_err_rate)
|
||||
|
||||
@@ -400,7 +400,7 @@ class WhisperModel:
|
||||
|
||||
return segments, info
|
||||
|
||||
def generate_segments(
|
||||
def generate_segments( # noqa: C901
|
||||
self,
|
||||
features: np.ndarray,
|
||||
tokenizer: Tokenizer,
|
||||
@@ -425,7 +425,7 @@ class WhisperModel:
|
||||
all_segments = []
|
||||
while seek < content_frames:
|
||||
time_offset = seek * self.feature_extractor.time_per_frame
|
||||
segment = features[:, seek : seek + self.feature_extractor.nb_max_frames]
|
||||
segment = features[:, seek:seek + self.feature_extractor.nb_max_frames]
|
||||
segment_size = min(
|
||||
self.feature_extractor.nb_max_frames, content_frames - seek
|
||||
)
|
||||
@@ -749,7 +749,7 @@ class WhisperModel:
|
||||
|
||||
if previous_tokens:
|
||||
prompt.append(tokenizer.sot_prev)
|
||||
prompt.extend(previous_tokens[-(self.max_length // 2 - 1) :])
|
||||
prompt.extend(previous_tokens[-(self.max_length // 2 - 1):])
|
||||
|
||||
prompt.extend(tokenizer.sot_sequence)
|
||||
|
||||
@@ -766,7 +766,7 @@ class WhisperModel:
|
||||
|
||||
return prompt
|
||||
|
||||
def add_word_timestamps(
|
||||
def add_word_timestamps( # noqa: C901
|
||||
self,
|
||||
segments: List[dict],
|
||||
tokenizer: Tokenizer,
|
||||
|
||||
@@ -1,17 +1,14 @@
|
||||
import argparse
|
||||
import json
|
||||
import re
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
from pathlib import Path
|
||||
from typing import Dict, Iterable, List, Optional, TextIO, Tuple, Union
|
||||
from typing import Union
|
||||
|
||||
import torch
|
||||
import numpy as np
|
||||
import torch.nn.functional as F
|
||||
from whisper.tokenizer import get_tokenizer
|
||||
from whisper_live.tensorrt_utils import (mel_filters, store_transcripts,
|
||||
write_error_stats, load_audio_wav_format,
|
||||
pad_or_trim, load_audio)
|
||||
from whisper_live.tensorrt_utils import (mel_filters, load_audio_wav_format, pad_or_trim, load_audio)
|
||||
|
||||
import tensorrt_llm
|
||||
import tensorrt_llm.logger as logger
|
||||
@@ -38,8 +35,6 @@ class WhisperEncoding:
|
||||
with open(config_path, 'r') as f:
|
||||
config = json.load(f)
|
||||
|
||||
use_gpt_attention_plugin = config['plugin_config'][
|
||||
'gpt_attention_plugin']
|
||||
dtype = config['builder_config']['precision']
|
||||
n_mels = config['builder_config']['n_mels']
|
||||
num_languages = config['builder_config']['num_languages']
|
||||
@@ -176,16 +171,8 @@ class WhisperDecoding:
|
||||
|
||||
class WhisperTRTLLM(object):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
engine_dir,
|
||||
debug_mode=False,
|
||||
assets_dir=None,
|
||||
device=None,
|
||||
is_multilingual=False,
|
||||
language="en",
|
||||
task="transcribe"
|
||||
):
|
||||
def __init__(self, engine_dir, assets_dir=None, device=None, is_multilingual=False,
|
||||
language="en", task="transcribe"):
|
||||
world_size = 1
|
||||
runtime_rank = tensorrt_llm.mpi_rank()
|
||||
runtime_mapping = tensorrt_llm.Mapping(world_size, runtime_rank)
|
||||
@@ -212,7 +199,7 @@ class WhisperTRTLLM(object):
|
||||
self,
|
||||
audio: Union[str, np.ndarray, torch.Tensor],
|
||||
padding: int = 0,
|
||||
return_duration = True
|
||||
return_duration=True
|
||||
):
|
||||
"""
|
||||
Compute the log-Mel spectrogram of
|
||||
@@ -242,8 +229,7 @@ class WhisperTRTLLM(object):
|
||||
audio, _ = load_audio_wav_format(audio)
|
||||
else:
|
||||
audio = load_audio(audio)
|
||||
assert isinstance(audio,
|
||||
np.ndarray), f"Unsupported audio type: {type(audio)}"
|
||||
assert isinstance(audio, np.ndarray), f"Unsupported audio type: {type(audio)}"
|
||||
duration = audio.shape[-1] / SAMPLE_RATE
|
||||
audio = pad_or_trim(audio, N_SAMPLES)
|
||||
audio = audio.astype(np.float32)
|
||||
@@ -254,14 +240,9 @@ class WhisperTRTLLM(object):
|
||||
if padding > 0:
|
||||
audio = F.pad(audio, (0, padding))
|
||||
window = torch.hann_window(N_FFT).to(audio.device)
|
||||
stft = torch.stft(audio,
|
||||
N_FFT,
|
||||
HOP_LENGTH,
|
||||
window=window,
|
||||
return_complex=True)
|
||||
stft = torch.stft(audio, N_FFT, HOP_LENGTH, window=window, return_complex=True)
|
||||
magnitudes = stft[..., :-1].abs()**2
|
||||
|
||||
|
||||
mel_spec = self.filters @ magnitudes
|
||||
|
||||
log_spec = torch.clamp(mel_spec, min=1e-10).log10()
|
||||
@@ -272,7 +253,6 @@ class WhisperTRTLLM(object):
|
||||
else:
|
||||
return log_spec
|
||||
|
||||
|
||||
def process_batch(
|
||||
self,
|
||||
mel,
|
||||
@@ -296,7 +276,7 @@ class WhisperTRTLLM(object):
|
||||
text = self.tokenizer.decode(output_ids[i][0]).strip()
|
||||
texts.append(text)
|
||||
return texts
|
||||
|
||||
|
||||
def transcribe(
|
||||
self,
|
||||
mel,
|
||||
@@ -336,5 +316,5 @@ def decode_wav_file(
|
||||
prediction = re.sub(r'<\|.*?\|>', '', prediction)
|
||||
if normalizer:
|
||||
prediction = normalizer(prediction)
|
||||
|
||||
|
||||
return prediction.strip()
|
||||
|
||||
@@ -0,0 +1,71 @@
|
||||
import os
|
||||
import textwrap
|
||||
import scipy
|
||||
import ffmpeg
|
||||
import numpy as np
|
||||
|
||||
|
||||
def clear_screen():
|
||||
"""Clears the console screen."""
|
||||
os.system("cls" if os.name == "nt" else "clear")
|
||||
|
||||
|
||||
def print_transcript(text):
|
||||
"""Prints formatted transcript text."""
|
||||
wrapper = textwrap.TextWrapper(width=60)
|
||||
for line in wrapper.wrap(text="".join(text)):
|
||||
print(line)
|
||||
|
||||
|
||||
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):
|
||||
"""
|
||||
# https://github.com/openai/whisper/blob/7858aa9c08d98f75575035ecd6481f462d66ca27/whisper/audio.py#L22
|
||||
Open an audio file and read as mono waveform, resampling as necessary,
|
||||
save the resampled audio
|
||||
|
||||
Args:
|
||||
file (str): The audio file to open
|
||||
sr (int): The sample rate to resample the audio if necessary
|
||||
|
||||
Returns:
|
||||
resampled_file (str): The resampled audio file
|
||||
"""
|
||||
try:
|
||||
# This launches a subprocess to decode audio while down-mixing and resampling as necessary.
|
||||
# Requires the ffmpeg CLI and `ffmpeg-python` package to be installed.
|
||||
out, _ = (
|
||||
ffmpeg.input(file, threads=0)
|
||||
.output("-", format="s16le", acodec="pcm_s16le", ac=1, ar=sr)
|
||||
.run(cmd=["ffmpeg", "-nostdin"], capture_stdout=True, capture_stderr=True)
|
||||
)
|
||||
except ffmpeg.Error as e:
|
||||
raise RuntimeError(f"Failed to load audio: {e.stderr.decode()}") from e
|
||||
np_buffer = np.frombuffer(out, dtype=np.int16)
|
||||
|
||||
resampled_file = f"{file.split('.')[0]}_resampled.wav"
|
||||
scipy.io.wavfile.write(resampled_file, sr, np_buffer.astype(np.int16))
|
||||
return resampled_file
|
||||
+31
-2
@@ -34,7 +34,7 @@ class VoiceActivityDetection():
|
||||
|
||||
if sr != 16000 and (sr % 16000 == 0):
|
||||
step = sr // 16000
|
||||
x = x[:,::step]
|
||||
x = x[:, ::step]
|
||||
sr = 16000
|
||||
|
||||
if sr not in self.sample_rates:
|
||||
@@ -110,4 +110,33 @@ class VoiceActivityDetection():
|
||||
subprocess.run(["wget", "-O", model_filename, model_url], check=True)
|
||||
except subprocess.CalledProcessError:
|
||||
print("Failed to download the model using wget.")
|
||||
return model_filename
|
||||
return model_filename
|
||||
|
||||
|
||||
class VoiceActivityDetector:
|
||||
def __init__(self, threshold=0.5, frame_rate=16000):
|
||||
"""
|
||||
Initializes the VoiceActivityDetector with a voice activity detection model and a threshold.
|
||||
|
||||
Args:
|
||||
threshold (float, optional): The probability threshold for detecting voice activity. Defaults to 0.5.
|
||||
"""
|
||||
self.model = VoiceActivityDetection()
|
||||
self.threshold = threshold
|
||||
self.frame_rate = frame_rate
|
||||
|
||||
def __call__(self, audio_frame):
|
||||
"""
|
||||
Determines if the given audio frame contains speech by comparing the detected speech probability against
|
||||
the threshold.
|
||||
|
||||
Args:
|
||||
audio_frame (np.ndarray): The audio frame to be analyzed for voice activity. It is expected to be a
|
||||
NumPy array of audio samples.
|
||||
|
||||
Returns:
|
||||
bool: True if the speech probability exceeds the threshold, indicating the presence of voice activity;
|
||||
False otherwise.
|
||||
"""
|
||||
speech_prob = self.model(torch.from_numpy(audio_frame), self.frame_rate).item()
|
||||
return speech_prob > self.threshold
|
||||
|
||||
Reference in New Issue
Block a user