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:
Vineet Suryan
2026-07-06 11:37:35 +05:30
committed by GitHub
2 changed files with 150 additions and 17 deletions
+126 -1
View File
@@ -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()
+12 -4
View File
@@ -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,7 +187,7 @@ 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):]
@@ -193,7 +200,7 @@ class ServeClientBase(object):
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)