Merge pull request #517 from dmaier-ef/fix/idle-transcription-thread-cpu-contention
Fix idle-client busy-wait before first audio frame
This commit is contained in:
+126
-1
@@ -24,6 +24,24 @@ class ConcreteServeClient(ServeClientBase):
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class WaitTrackingEvent:
|
||||||
|
"""Threading event that records when wait() is entered."""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self._event = threading.Event()
|
||||||
|
self.wait_started = threading.Event()
|
||||||
|
|
||||||
|
def wait(self, timeout=None):
|
||||||
|
self.wait_started.set()
|
||||||
|
return self._event.wait(timeout)
|
||||||
|
|
||||||
|
def set(self):
|
||||||
|
self._event.set()
|
||||||
|
|
||||||
|
def __getattr__(self, name):
|
||||||
|
return getattr(self._event, name)
|
||||||
|
|
||||||
|
|
||||||
class TestServeClientBaseInit(unittest.TestCase):
|
class TestServeClientBaseInit(unittest.TestCase):
|
||||||
def test_default_values(self):
|
def test_default_values(self):
|
||||||
ws = MagicMock()
|
ws = MagicMock()
|
||||||
@@ -91,7 +109,6 @@ class TestAddFrames(unittest.TestCase):
|
|||||||
# timestamp_offset should be bumped to at least frames_offset
|
# timestamp_offset should be bumped to at least frames_offset
|
||||||
self.assertGreaterEqual(self.client.timestamp_offset, self.client.frames_offset)
|
self.assertGreaterEqual(self.client.timestamp_offset, self.client.frames_offset)
|
||||||
|
|
||||||
|
|
||||||
class TestAddFramesThreadSafety(unittest.TestCase):
|
class TestAddFramesThreadSafety(unittest.TestCase):
|
||||||
def test_concurrent_add_frames(self):
|
def test_concurrent_add_frames(self):
|
||||||
ws = MagicMock()
|
ws = MagicMock()
|
||||||
@@ -114,6 +131,18 @@ class TestAddFramesThreadSafety(unittest.TestCase):
|
|||||||
self.assertEqual(errors, [])
|
self.assertEqual(errors, [])
|
||||||
self.assertIsNotNone(client.frames_np)
|
self.assertIsNotNone(client.frames_np)
|
||||||
|
|
||||||
|
def test_exception_releases_lock_without_signaling_frames_ready(self):
|
||||||
|
ws = MagicMock()
|
||||||
|
client = ConcreteServeClient(client_uid="test", websocket=ws)
|
||||||
|
client.frames_np = np.array([0.1], dtype=np.float32)
|
||||||
|
|
||||||
|
with patch("whisper_live.backend.base.np.concatenate", side_effect=RuntimeError("boom")):
|
||||||
|
with self.assertRaisesRegex(RuntimeError, "boom"):
|
||||||
|
client.add_frames(np.array([0.2], dtype=np.float32))
|
||||||
|
|
||||||
|
self.assertFalse(client.lock.locked())
|
||||||
|
self.assertFalse(client.frames_ready.is_set())
|
||||||
|
|
||||||
|
|
||||||
class TestGetAudioChunkForProcessing(unittest.TestCase):
|
class TestGetAudioChunkForProcessing(unittest.TestCase):
|
||||||
def setUp(self):
|
def setUp(self):
|
||||||
@@ -266,6 +295,102 @@ class TestCleanup(unittest.TestCase):
|
|||||||
self.assertTrue(client.exit)
|
self.assertTrue(client.exit)
|
||||||
|
|
||||||
|
|
||||||
|
def _supports_thread_time():
|
||||||
|
thread_time = getattr(time, "thread_time", None)
|
||||||
|
if thread_time is None:
|
||||||
|
return False
|
||||||
|
try:
|
||||||
|
thread_time()
|
||||||
|
except NotImplementedError:
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
class TestSpeechToTextWaitingBehavior(unittest.TestCase):
|
||||||
|
"""Tests the first-frame wait behavior in speech_to_text()."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
self.ws = MagicMock()
|
||||||
|
self.client = ConcreteServeClient(client_uid="test", websocket=self.ws)
|
||||||
|
self.client.frames_ready = WaitTrackingEvent()
|
||||||
|
self.transcribe_called = threading.Event()
|
||||||
|
self.thread_started = threading.Event()
|
||||||
|
self.cpu_used = None
|
||||||
|
self.thread = None
|
||||||
|
|
||||||
|
def tearDown(self):
|
||||||
|
if self.thread is not None and self.thread.is_alive():
|
||||||
|
self.client.exit = True
|
||||||
|
# release wait() directly so a broken cleanup() cannot hang the test process
|
||||||
|
self.client.frames_ready.set()
|
||||||
|
self.thread.join(timeout=1.0)
|
||||||
|
|
||||||
|
def _start_speech_thread(self, target=None):
|
||||||
|
self.thread = threading.Thread(target=target or self.client.speech_to_text)
|
||||||
|
self.thread.start()
|
||||||
|
return self.thread
|
||||||
|
|
||||||
|
def _join_speech_thread(self):
|
||||||
|
self.thread.join(timeout=1.0)
|
||||||
|
return not self.thread.is_alive()
|
||||||
|
|
||||||
|
def _transcribe_once(self, input_sample):
|
||||||
|
# mark the first processing step after wait and stop the loop
|
||||||
|
self.transcribe_called.set()
|
||||||
|
self.client.exit = True
|
||||||
|
return []
|
||||||
|
|
||||||
|
def _measure_waiting_cpu(self):
|
||||||
|
# measure CPU consumed by speech_to_text loop while it waits for the first frame
|
||||||
|
self.thread_started.set()
|
||||||
|
start_cpu = time.thread_time()
|
||||||
|
self.client.speech_to_text()
|
||||||
|
self.cpu_used = time.thread_time() - start_cpu
|
||||||
|
|
||||||
|
def test_waits_for_first_frame_before_transcribing(self):
|
||||||
|
self.client.transcribe_audio = MagicMock(side_effect=self._transcribe_once)
|
||||||
|
self._start_speech_thread()
|
||||||
|
self.assertTrue(self.client.frames_ready.wait_started.wait(timeout=1.0))
|
||||||
|
self.assertFalse(self.transcribe_called.is_set())
|
||||||
|
|
||||||
|
self.client.add_frames(np.zeros(self.client.RATE, dtype=np.float32))
|
||||||
|
|
||||||
|
self.assertTrue(self.transcribe_called.wait(timeout=1.0))
|
||||||
|
self.assertTrue(self._join_speech_thread())
|
||||||
|
|
||||||
|
def test_cleanup_unblocks_waiting_thread_without_audio(self):
|
||||||
|
self.client.transcribe_audio = MagicMock()
|
||||||
|
self._start_speech_thread()
|
||||||
|
self.assertTrue(self.client.frames_ready.wait_started.wait(timeout=1.0))
|
||||||
|
|
||||||
|
self.client.cleanup()
|
||||||
|
|
||||||
|
self.assertTrue(self._join_speech_thread())
|
||||||
|
self.assertTrue(self.client.exit)
|
||||||
|
self.client.transcribe_audio.assert_not_called()
|
||||||
|
|
||||||
|
def test_exit_flag_unblocks_waiting_thread_without_signal(self):
|
||||||
|
self.client.transcribe_audio = MagicMock()
|
||||||
|
self._start_speech_thread()
|
||||||
|
self.assertTrue(self.client.frames_ready.wait_started.wait(timeout=1.0))
|
||||||
|
|
||||||
|
self.client.exit = True
|
||||||
|
|
||||||
|
self.assertTrue(self._join_speech_thread())
|
||||||
|
self.client.transcribe_audio.assert_not_called()
|
||||||
|
|
||||||
|
@unittest.skipUnless(_supports_thread_time(), "time.thread_time() not supported")
|
||||||
|
def test_waiting_for_first_frame_uses_negligible_thread_cpu(self):
|
||||||
|
self._start_speech_thread(target=self._measure_waiting_cpu)
|
||||||
|
self.assertTrue(self.thread_started.wait(timeout=1.0))
|
||||||
|
|
||||||
|
time.sleep(0.25)
|
||||||
|
self.client.cleanup()
|
||||||
|
|
||||||
|
self.assertTrue(self._join_speech_thread())
|
||||||
|
self.assertIsNotNone(self.cpu_used)
|
||||||
|
self.assertLess(self.cpu_used, 0.1)
|
||||||
|
|
||||||
|
|
||||||
class TestTrimTranscript(unittest.TestCase):
|
class TestTrimTranscript(unittest.TestCase):
|
||||||
def setUp(self):
|
def setUp(self):
|
||||||
self.ws = MagicMock()
|
self.ws = MagicMock()
|
||||||
|
|||||||
@@ -21,6 +21,8 @@ class ServeClientBase(object):
|
|||||||
"""Duration threshold in seconds for clipping audio with no valid segments."""
|
"""Duration threshold in seconds for clipping audio with no valid segments."""
|
||||||
CLIP_TAIL_DURATION_S = 5
|
CLIP_TAIL_DURATION_S = 5
|
||||||
"""Duration in seconds of audio to keep after clipping."""
|
"""Duration in seconds of audio to keep after clipping."""
|
||||||
|
FIRST_FRAME_WAIT_TIMEOUT_S = 0.1
|
||||||
|
"""Interval in seconds for re-checking exit while waiting for the first audio frame."""
|
||||||
|
|
||||||
client_uid: str
|
client_uid: str
|
||||||
"""A unique identifier for the client."""
|
"""A unique identifier for the client."""
|
||||||
@@ -81,13 +83,15 @@ class ServeClientBase(object):
|
|||||||
|
|
||||||
# threading
|
# threading
|
||||||
self.lock = threading.Lock()
|
self.lock = threading.Lock()
|
||||||
|
self.frames_ready = threading.Event()
|
||||||
|
|
||||||
def speech_to_text(self):
|
def speech_to_text(self):
|
||||||
"""
|
"""
|
||||||
Process an audio stream in an infinite loop, continuously transcribing the speech.
|
Process an audio stream in an infinite loop, continuously transcribing the speech.
|
||||||
|
|
||||||
This method continuously receives audio frames, performs real-time transcription, and sends
|
This method continuously receives audio frames, performs real-time transcription, and sends
|
||||||
transcribed segments to the client via a WebSocket connection.
|
transcribed segments to the client via a WebSocket connection. The loop blocks until the first
|
||||||
|
audio frame arrives when a client is connected but still idle.
|
||||||
|
|
||||||
If the client's language is not detected, it waits for 30 seconds of audio input to make a language prediction.
|
If the client's language is not detected, it waits for 30 seconds of audio input to make a language prediction.
|
||||||
It utilizes the Whisper ASR model to transcribe the audio, continuously processing and streaming results. Segments
|
It utilizes the Whisper ASR model to transcribe the audio, continuously processing and streaming results. Segments
|
||||||
@@ -103,6 +107,8 @@ class ServeClientBase(object):
|
|||||||
break
|
break
|
||||||
|
|
||||||
if self.frames_np is None:
|
if self.frames_np is None:
|
||||||
|
while self.frames_np is None and not self.exit:
|
||||||
|
self.frames_ready.wait(timeout=self.FIRST_FRAME_WAIT_TIMEOUT_S)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
if self.clip_audio:
|
if self.clip_audio:
|
||||||
@@ -170,7 +176,8 @@ class ServeClientBase(object):
|
|||||||
|
|
||||||
This method is responsible for maintaining the audio stream buffer, allowing the continuous addition
|
This method is responsible for maintaining the audio stream buffer, allowing the continuous addition
|
||||||
of audio frames as they are received. It also ensures that the buffer does not exceed a specified size
|
of audio frames as they are received. It also ensures that the buffer does not exceed a specified size
|
||||||
to prevent excessive memory usage.
|
to prevent excessive memory usage. When the first frame arrives, it also wakes the transcription
|
||||||
|
thread so processing can begin.
|
||||||
|
|
||||||
If the buffer size exceeds a threshold (45 seconds of audio data), it discards the oldest 30 seconds
|
If the buffer size exceeds a threshold (45 seconds of audio data), it discards the oldest 30 seconds
|
||||||
of audio data to maintain a reasonable buffer size. If the buffer is empty, it initializes it with the provided
|
of audio data to maintain a reasonable buffer size. If the buffer is empty, it initializes it with the provided
|
||||||
@@ -180,20 +187,20 @@ class ServeClientBase(object):
|
|||||||
frame_np (numpy.ndarray): The audio frame data as a NumPy array.
|
frame_np (numpy.ndarray): The audio frame data as a NumPy array.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
self.lock.acquire()
|
with self.lock:
|
||||||
if self.frames_np is not None and self.frames_np.shape[0] > self.MAX_BUFFER_DURATION_S*self.RATE:
|
if self.frames_np is not None and self.frames_np.shape[0] > self.MAX_BUFFER_DURATION_S*self.RATE:
|
||||||
self.frames_offset += float(self.BUFFER_TRIM_DURATION_S)
|
self.frames_offset += float(self.BUFFER_TRIM_DURATION_S)
|
||||||
self.frames_np = self.frames_np[int(self.BUFFER_TRIM_DURATION_S*self.RATE):]
|
self.frames_np = self.frames_np[int(self.BUFFER_TRIM_DURATION_S*self.RATE):]
|
||||||
# check timestamp offset(should be >= self.frame_offset)
|
# check timestamp offset(should be >= self.frame_offset)
|
||||||
# this basically means that there is no speech as timestamp offset hasnt updated
|
# this basically means that there is no speech as timestamp offset hasnt updated
|
||||||
# and is less than frame_offset
|
# and is less than frame_offset
|
||||||
if self.timestamp_offset < self.frames_offset:
|
if self.timestamp_offset < self.frames_offset:
|
||||||
self.timestamp_offset = self.frames_offset
|
self.timestamp_offset = self.frames_offset
|
||||||
if self.frames_np is None:
|
if self.frames_np is None:
|
||||||
self.frames_np = frame_np.copy()
|
self.frames_np = frame_np.copy()
|
||||||
else:
|
else:
|
||||||
self.frames_np = np.concatenate((self.frames_np, frame_np), axis=0)
|
self.frames_np = np.concatenate((self.frames_np, frame_np), axis=0)
|
||||||
self.lock.release()
|
self.frames_ready.set()
|
||||||
|
|
||||||
def clip_audio_if_no_valid_segment(self):
|
def clip_audio_if_no_valid_segment(self):
|
||||||
"""
|
"""
|
||||||
@@ -323,6 +330,7 @@ class ServeClientBase(object):
|
|||||||
"""
|
"""
|
||||||
logging.info("Cleaning up.")
|
logging.info("Cleaning up.")
|
||||||
self.exit = True
|
self.exit = True
|
||||||
|
self.frames_ready.set()
|
||||||
|
|
||||||
def get_segment_no_speech_prob(self, segment):
|
def get_segment_no_speech_prob(self, segment):
|
||||||
return getattr(segment, "no_speech_prob", 0)
|
return getattr(segment, "no_speech_prob", 0)
|
||||||
|
|||||||
Reference in New Issue
Block a user