🔨 refactor whisper_live according to flake8

This commit is contained in:
makaveli10
2024-02-09 13:45:13 +05:30
parent b4abe95fc6
commit 9fbff47126
13 changed files with 549 additions and 540 deletions
+8 -7
View File
@@ -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)
+20 -22
View File
@@ -14,32 +14,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 +49,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 +56,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 +69,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 +82,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)
+1 -2
View File
@@ -1,7 +1,6 @@
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
@@ -25,4 +24,4 @@ class TestVoiceActivityDetection(unittest.TestCase):
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.")
self.assertGreater(speech_prob, 0.5, "VAD failed to identify speech segment.")