From e8bd4fd5320dd7c8aed9d791f08779277c9d2fda Mon Sep 17 00:00:00 2001 From: ianwh02 Date: Thu, 26 Feb 2026 00:00:51 +0000 Subject: [PATCH 1/3] Add cross-client GPU batch inference for faster_whisper backend --- run_server.py | 21 + .../backend/faster_whisper_backend.py | 21 + whisper_live/batch_inference.py | 397 ++++++++++++++++++ whisper_live/server.py | 38 +- 4 files changed, 476 insertions(+), 1 deletion(-) create mode 100644 whisper_live/batch_inference.py diff --git a/run_server.py b/run_server.py index 67c8af5..d45a4c8 100644 --- a/run_server.py +++ b/run_server.py @@ -66,6 +66,24 @@ if __name__ == "__main__": default=None, help="Comma-separated list of allowed CORS origins (e.g., 'http://localhost:3000,http://example.com'). Defaults to localhost/127.0.0.1 on the WebSocket port." ) + parser.add_argument( + '--batch_inference', + action='store_true', + help='Enable batched GPU inference for concurrent sessions. ' + 'Batches multiple sessions into a single GPU call for higher throughput.' + ) + parser.add_argument( + '--batch_max_size', + type=int, + default=8, + help='Maximum batch size for batched inference (default: 8).' + ) + parser.add_argument( + '--batch_window_ms', + type=int, + default=50, + help='Maximum time in ms to wait for batch to fill (default: 50).' + ) args = parser.parse_args() if args.backend == "tensorrt": @@ -92,4 +110,7 @@ if __name__ == "__main__": rest_port=args.rest_port, enable_rest=args.enable_rest, cors_origins=args.cors_origins, + batch_enabled=args.batch_inference, + batch_max_size=args.batch_max_size, + batch_window_ms=args.batch_window_ms, ) \ No newline at end of file diff --git a/whisper_live/backend/faster_whisper_backend.py b/whisper_live/backend/faster_whisper_backend.py index 60cf3cd..67463ca 100644 --- a/whisper_live/backend/faster_whisper_backend.py +++ b/whisper_live/backend/faster_whisper_backend.py @@ -14,6 +14,7 @@ from whisper_live.backend.base import ServeClientBase class ServeClientFasterWhisper(ServeClientBase): SINGLE_MODEL = None SINGLE_MODEL_LOCK = threading.Lock() + BATCH_WORKER = None def __init__( self, @@ -202,6 +203,26 @@ class ServeClientFasterWhisper(ServeClientBase): depends on the implementation of the `transcriber.transcribe` method but typically includes the transcribed text. """ + # Batch inference path: submit to central queue and wait + if ServeClientFasterWhisper.BATCH_WORKER is not None: + from whisper_live.batch_inference import BatchRequest + request = BatchRequest( + audio=input_sample, + language=self.language, + task=self.task, + initial_prompt=self.initial_prompt, + use_vad=self.use_vad, + vad_parameters=self.vad_parameters if self.use_vad else None, + ) + ServeClientFasterWhisper.BATCH_WORKER.submit(request) + request.future.wait(timeout=30) + if request.error: + raise request.error + if self.language is None and request.info is not None: + self.set_language(request.info) + return request.result + + # Original lock-based path (backward compatible) if ServeClientFasterWhisper.SINGLE_MODEL: ServeClientFasterWhisper.SINGLE_MODEL_LOCK.acquire() result, info = self.transcriber.transcribe( diff --git a/whisper_live/batch_inference.py b/whisper_live/batch_inference.py new file mode 100644 index 0000000..0e74180 --- /dev/null +++ b/whisper_live/batch_inference.py @@ -0,0 +1,397 @@ +""" +Batch inference scheduler for WhisperLive. + +Replaces the per-session SINGLE_MODEL_LOCK with a queue-based batch system. +Multiple sessions submit audio to a central queue; a single dedicated thread +collects pending requests and runs them as a GPU batch via CTranslate2's +batched encode() + generate() API. + +For batch_size=1, falls back to standard transcriber.transcribe() for +identical behavior to the non-batched path. + +Usage: + Enable via ``--batch_inference`` CLI flag. The batch worker is lazily + started after the first client connects and the shared model is loaded. + +Thread safety: + - ``queue.Queue`` is stdlib thread-safe. + - Each ``BatchRequest.future`` (``threading.Event``) is written by the + batch worker BEFORE ``.set()``, read by the session thread AFTER + ``.wait()`` — no data race. + - Only the batch worker thread touches the GPU model — zero lock + contention between session threads. +""" + +import logging +import queue +import threading +import time +from dataclasses import dataclass, field +from math import ceil +from typing import Any, Dict, List, Optional + +import numpy as np + +from faster_whisper.audio import pad_or_trim +from faster_whisper.tokenizer import Tokenizer +from faster_whisper.vad import ( + VadOptions, + collect_chunks, + get_speech_timestamps, +) + +from whisper_live.transcriber.transcriber_faster_whisper import ( + Segment, + TranscriptionInfo, + get_compression_ratio, + get_suppressed_tokens, +) + + +@dataclass +class BatchRequest: + """A single inference request submitted by a session thread. + + The session thread creates this, calls ``BatchInferenceWorker.submit()``, + then blocks on ``future.wait()``. The batch worker fills ``result`` + and/or ``error``, then signals ``future.set()``. + + Attributes: + audio: Raw audio samples (float32, 16 kHz mono). + language: ISO language code or None for auto-detection. + task: ``"transcribe"`` or ``"translate"``. + initial_prompt: Optional prompt for Whisper conditioning. + use_vad: Whether to apply Voice Activity Detection. + vad_parameters: Parameters forwarded to ``VadOptions``. + future: Event signaled when the result is ready. + result: List of ``Segment`` objects (filled by worker). + info: ``TranscriptionInfo`` metadata (filled by worker). + error: Exception instance if processing failed. + """ + audio: np.ndarray + language: Optional[str] = None + task: str = "transcribe" + initial_prompt: Optional[str] = None + use_vad: bool = True + vad_parameters: Optional[Dict] = None + # Signaling + future: threading.Event = field(default_factory=threading.Event) + # Results (filled by batch worker) + result: Optional[Any] = None + info: Optional[Any] = None + error: Optional[Exception] = None + + +class BatchInferenceWorker: + """Central batch inference scheduler for the faster_whisper backend. + + Owns a single daemon thread that is the **only** thread touching the GPU + model. Per-session transcription threads submit ``BatchRequest`` objects + and block on ``future.wait()`` instead of competing for + ``SINGLE_MODEL_LOCK``. + + The worker loop: + + 1. Blocks until the first request arrives from the queue. + 2. Waits up to ``batch_window_ms`` for additional requests (up to + ``max_batch_size``). + 3. Processes the collected batch: + - **batch_size == 1**: delegates to ``transcriber.transcribe()`` for + identical behavior to the non-batched path. + - **batch_size > 1**: runs a custom batched GPU path using + CTranslate2's ``encode()`` + ``generate()`` APIs. + + Args: + transcriber: The shared ``WhisperModel`` instance. + max_batch_size: Maximum number of requests per batch. + batch_window_ms: Maximum time (ms) to wait for the batch to fill + after the first request arrives. + """ + + def __init__( + self, + transcriber, + max_batch_size: int = 8, + batch_window_ms: int = 50, + ): + self.transcriber = transcriber + self.max_batch_size = max_batch_size + self.batch_window_ms = batch_window_ms + self._queue: queue.Queue = queue.Queue() + self._stop_event = threading.Event() + self._thread: Optional[threading.Thread] = None + + def start(self): + """Start the background batch worker thread.""" + self._thread = threading.Thread(target=self._worker_loop, daemon=True) + self._thread.start() + logging.info( + f"[BatchInference] Started (max_batch={self.max_batch_size}, " + f"window={self.batch_window_ms}ms)" + ) + + def stop(self): + """Signal the worker to stop and wait for it to finish.""" + self._stop_event.set() + if self._thread: + self._thread.join(timeout=5) + + def submit(self, request: BatchRequest): + """Submit an inference request to the batch queue. + + Args: + request: The ``BatchRequest`` to enqueue. The caller should + then call ``request.future.wait()`` to block until the + result is ready. + """ + self._queue.put(request) + + # ------------------------------------------------------------------------- + # Worker loop + # ------------------------------------------------------------------------- + + def _worker_loop(self): + """Main loop: collect requests into batches and process them.""" + while not self._stop_event.is_set(): + batch: List[BatchRequest] = [] + + # Block until first request arrives + try: + first = self._queue.get(timeout=0.5) + batch.append(first) + except queue.Empty: + continue + + # Collect more requests within the batch window + deadline = time.monotonic() + (self.batch_window_ms / 1000.0) + while len(batch) < self.max_batch_size: + remaining = deadline - time.monotonic() + if remaining <= 0: + break + try: + item = self._queue.get(timeout=remaining) + batch.append(item) + except queue.Empty: + break + + # Process the collected batch + try: + self._process_batch(batch) + except Exception as e: + logging.error(f"[BatchInference] Batch processing error: {e}") + for req in batch: + if not req.future.is_set(): + req.error = e + req.future.set() + + # ------------------------------------------------------------------------- + # Batch processing + # ------------------------------------------------------------------------- + + def _process_batch(self, batch: List[BatchRequest]): + """Dispatch to single or multi-item processing.""" + if len(batch) == 1: + self._process_single(batch[0]) + return + + logging.info(f"[BatchInference] Processing batch of {len(batch)}") + self._process_multi(batch) + + def _process_single(self, req: BatchRequest): + """Process a single request using standard ``transcriber.transcribe()``. + + This path is used when only one request is available in the batch + window, ensuring identical behavior to the non-batched code path. + """ + try: + result, info = self.transcriber.transcribe( + req.audio, + language=req.language, + task=req.task, + initial_prompt=req.initial_prompt, + vad_filter=req.use_vad, + vad_parameters=req.vad_parameters if req.use_vad else None, + ) + # Materialize the generator into a list + req.result = list(result) + req.info = info + except Exception as e: + req.error = e + finally: + req.future.set() + + def _process_multi(self, batch: List[BatchRequest]): + """Batched GPU path: encode + generate for multiple sessions at once. + + Pipeline: + 1. Per-item CPU preprocessing (VAD filtering + mel feature extraction) + 2. Batch GPU encode — single ``transcriber.encode()`` call + 3. Per-item prompt construction (handles different languages/tasks) + 4. Batch GPU generate — single ``transcriber.model.generate()`` call + 5. Per-item segment parsing and result dispatch + """ + # Step 1: Per-item CPU preprocessing (VAD + feature extraction) + preprocessed = [] + for req in batch: + try: + audio = req.audio + speech_chunks = None + + if req.use_vad: + vad_params = req.vad_parameters or {} + vad_opts = VadOptions(**vad_params) if isinstance(vad_params, dict) else vad_params + speech_chunks = get_speech_timestamps(audio, vad_opts) + if speech_chunks: + audio_chunks, _ = collect_chunks(audio, speech_chunks) + audio = np.concatenate(audio_chunks, axis=0) if audio_chunks else audio + + if audio.shape[0] == 0: + # No speech detected — return empty result immediately + req.result = [] + req.info = self._make_info(req, 0.0, 0.0) + req.future.set() + continue + + duration = audio.shape[0] / self.transcriber.feature_extractor.sampling_rate + features = self.transcriber.feature_extractor(audio) + features = pad_or_trim(features) # -> [n_mels, 3000] + preprocessed.append((req, features, audio, duration, speech_chunks)) + except Exception as e: + req.error = e + req.future.set() + + if not preprocessed: + return + + try: + # Step 2: Batch GPU encode + feature_batch = np.stack([p[1] for p in preprocessed]) # [B, n_mels, 3000] + encoder_output = self.transcriber.encode(feature_batch) + + # Step 3: Build per-item prompts (handles different languages/tasks) + tokenizers_list = [] + prompts = [] + resolved_languages = [] + + for i, (req, features, audio, duration, speech_chunks) in enumerate(preprocessed): + lang = req.language + # If language unknown, detect from encoder output + if lang is None: + try: + lang_results = self.transcriber.model.detect_language(encoder_output) + if lang_results and len(lang_results) > i: + detected = lang_results[i] + if detected: + lang = detected[0][0].strip("<|>") + except Exception: + lang = "en" # fallback + + resolved_languages.append(lang or "en") + + tokenizer = Tokenizer( + self.transcriber.hf_tokenizer, + self.transcriber.model.is_multilingual, + task=req.task, + language=lang or "en", + ) + + previous_tokens = [] + if req.initial_prompt: + previous_tokens = tokenizer.encode(" " + req.initial_prompt.strip()) + + prompt = self.transcriber.get_prompt( + tokenizer, + previous_tokens=previous_tokens, + without_timestamps=False, + ) + tokenizers_list.append(tokenizer) + prompts.append(prompt) + + # Step 4: Batch GPU generate + suppress_tokens = get_suppressed_tokens(tokenizers_list[0], [-1]) + + results = self.transcriber.model.generate( + encoder_output, + prompts, + beam_size=5, + patience=1, + length_penalty=1, + max_length=self.transcriber.max_length, + suppress_blank=True, + suppress_tokens=suppress_tokens, + return_scores=True, + return_no_speech_prob=True, + sampling_temperature=0.0, + repetition_penalty=1, + no_repeat_ngram_size=0, + ) + + # Step 5: Per-item segment parsing and result dispatch + for i, (req, features, audio, duration, speech_chunks) in enumerate(preprocessed): + try: + tokenizer = tokenizers_list[i] + gen_result = results[i] + + tokens = gen_result.sequences_ids[0] + seq_len = len(tokens) + cum_logprob = gen_result.scores[0] * seq_len + avg_logprob = cum_logprob / (seq_len + 1) if seq_len > 0 else 0.0 + + segment_size = int(ceil(duration) * self.transcriber.frames_per_second) + + subsegments, _, _ = self.transcriber._split_segments_by_timestamps( + tokenizer=tokenizer, + tokens=tokens, + time_offset=0, + segment_size=segment_size, + segment_duration=duration, + seek=0, + ) + + segments = [] + for seg_idx, subseg in enumerate(subsegments): + text = tokenizer.decode(subseg["tokens"]).strip() + if not text: + continue + segments.append(Segment( + id=seg_idx, + seek=subseg.get("seek", 0), + start=subseg["start"], + end=subseg["end"], + text=text, + tokens=subseg["tokens"], + avg_logprob=avg_logprob, + compression_ratio=get_compression_ratio(text), + no_speech_prob=gen_result.no_speech_prob, + words=None, + temperature=0.0, + )) + + req.result = segments + req.info = self._make_info( + req, duration, duration, + language=resolved_languages[i], + ) + except Exception as e: + req.error = e + finally: + req.future.set() + + except Exception as e: + logging.error(f"[BatchInference] GPU batch error: {e}") + for req, *_ in preprocessed: + if not req.future.is_set(): + req.error = e + req.future.set() + + def _make_info(self, req, duration, duration_after_vad, language=None): + """Build a ``TranscriptionInfo`` for the given request.""" + return TranscriptionInfo( + language=language or req.language or "en", + language_probability=1.0, + duration=duration, + duration_after_vad=duration_after_vad, + all_language_probs=None, + transcription_options=None, + vad_options=None, + ) diff --git a/whisper_live/server.py b/whisper_live/server.py index 02ca424..1654afb 100644 --- a/whisper_live/server.py +++ b/whisper_live/server.py @@ -277,6 +277,18 @@ class TranscriptionServer: ) logging.info("Running faster_whisper backend.") + + # Start batch inference worker on first client (after model is loaded) + if (self.batch_config is not None + and ServeClientFasterWhisper.BATCH_WORKER is None + and ServeClientFasterWhisper.SINGLE_MODEL is not None): + from whisper_live.batch_inference import BatchInferenceWorker + worker = BatchInferenceWorker( + transcriber=ServeClientFasterWhisper.SINGLE_MODEL, + **self.batch_config, + ) + worker.start() + ServeClientFasterWhisper.BATCH_WORKER = worker except Exception as e: logging.error(e) return @@ -415,13 +427,25 @@ class TranscriptionServer: cache_path="~/.cache/whisper-live/", rest_port=8000, enable_rest=False, - cors_origins: Optional[str] = None): + cors_origins: Optional[str] = None, + batch_enabled=False, + batch_max_size=8, + batch_window_ms=50): """ Run the transcription server. Args: host (str): The host address to bind the server. port (int): The port number to bind the server. + batch_enabled (bool): Enable cross-client GPU batch inference for + the faster_whisper backend. When enabled, ``single_model`` is + forced to True and a ``BatchInferenceWorker`` is started after + the first client connects. Defaults to False. + batch_max_size (int): Maximum number of requests per GPU batch. + Defaults to 8. + batch_window_ms (int): Maximum time in milliseconds to wait for + the batch to fill after the first request arrives. Defaults + to 50. """ self.cache_path = cache_path self.client_manager = ClientManager(max_clients, max_connection_time) @@ -430,6 +454,18 @@ class TranscriptionServer: raise ValueError(f"Custom faster_whisper model '{faster_whisper_custom_model_path}' is not a valid path or HuggingFace model.") if whisper_tensorrt_path is not None and not os.path.exists(whisper_tensorrt_path): raise ValueError(f"TensorRT model '{whisper_tensorrt_path}' is not a valid path.") + + # Batch inference config + if batch_enabled: + single_model = True # Batch mode requires shared model + self.batch_config = { + 'max_batch_size': batch_max_size, + 'batch_window_ms': batch_window_ms, + } + logging.info(f"Batch inference enabled (max_batch={batch_max_size}, window={batch_window_ms}ms)") + else: + self.batch_config = None + if single_model: if faster_whisper_custom_model_path or whisper_tensorrt_path: logging.info("Custom model option was provided. Switching to single model mode.") From 3508b39584e8b95692192156e8bd776c5b832d1f Mon Sep 17 00:00:00 2001 From: ianwh02 Date: Mon, 9 Mar 2026 11:42:15 +0000 Subject: [PATCH 2/3] Fix missing batch_config init causing CI test hang --- whisper_live/server.py | 1 + 1 file changed, 1 insertion(+) diff --git a/whisper_live/server.py b/whisper_live/server.py index 1654afb..c927990 100644 --- a/whisper_live/server.py +++ b/whisper_live/server.py @@ -159,6 +159,7 @@ class TranscriptionServer: self.no_voice_activity_chunks = 0 self.use_vad = True self.single_model = False + self.batch_config = None def initialize_client( self, websocket, options, faster_whisper_custom_model_path, From e7e78a71514f7886f6d83afd787b5f588a3cded3 Mon Sep 17 00:00:00 2001 From: ianwh02 Date: Mon, 9 Mar 2026 16:18:36 +0000 Subject: [PATCH 3/3] Add unit tests for BatchInferenceWorker --- tests/test_batch_inference.py | 163 ++++++++++++++++++++++++++++++++++ 1 file changed, 163 insertions(+) create mode 100644 tests/test_batch_inference.py diff --git a/tests/test_batch_inference.py b/tests/test_batch_inference.py new file mode 100644 index 0000000..0bbdc76 --- /dev/null +++ b/tests/test_batch_inference.py @@ -0,0 +1,163 @@ +import time +import unittest +from unittest import mock +from unittest.mock import MagicMock + +import numpy as np + +from whisper_live.batch_inference import BatchInferenceWorker, BatchRequest + + +class TestBatchInferenceWorker(unittest.TestCase): + def setUp(self): + self.mock_transcriber = MagicMock() + self.worker = BatchInferenceWorker( + transcriber=self.mock_transcriber, + max_batch_size=8, + batch_window_ms=200, + ) + self.worker.start() + + def tearDown(self): + self.worker.stop() + + def _make_audio(self, duration_s=1.0): + return np.random.randn(int(16000 * duration_s)).astype(np.float32) + + def test_single_request_uses_transcribe(self): + """Single request should fall back to transcriber.transcribe().""" + fake_segment = MagicMock() + fake_info = MagicMock() + self.mock_transcriber.transcribe.return_value = ([fake_segment], fake_info) + + req = BatchRequest(audio=self._make_audio(), language="en", use_vad=False) + self.worker.submit(req) + req.future.wait(timeout=5) + + self.assertTrue(req.future.is_set()) + self.assertIsNone(req.error) + self.assertEqual(req.result, [fake_segment]) + self.assertEqual(req.info, fake_info) + self.mock_transcriber.transcribe.assert_called_once() + + @mock.patch('whisper_live.batch_inference.get_suppressed_tokens', return_value=[-1]) + @mock.patch('whisper_live.batch_inference.Tokenizer') + def test_multiple_requests_batched(self, mock_tokenizer_cls, mock_suppress): + """Multiple concurrent requests should go through the batched GPU path.""" + # Mock tokenizer + mock_tok = MagicMock() + mock_tok.decode.return_value = "hello world" + mock_tokenizer_cls.return_value = mock_tok + + # Mock feature extractor + self.mock_transcriber.feature_extractor.return_value = np.zeros( + (80, 3000), dtype=np.float32 + ) + self.mock_transcriber.feature_extractor.sampling_rate = 16000 + + # Mock encode + self.mock_transcriber.encode.return_value = np.zeros( + (3, 1500, 512), dtype=np.float32 + ) + + # Mock model.generate — one result per item + gen_result = MagicMock() + gen_result.sequences_ids = [[50257, 50362, 1234, 50256]] + gen_result.scores = [np.float32(-1.0)] + gen_result.no_speech_prob = 0.1 + self.mock_transcriber.model.generate.return_value = [gen_result] * 3 + + # Mock remaining model attributes + self.mock_transcriber.model.is_multilingual = False + self.mock_transcriber.max_length = 448 + self.mock_transcriber.frames_per_second = 50 + self.mock_transcriber.get_prompt.return_value = [50258] + self.mock_transcriber._split_segments_by_timestamps.return_value = ( + [{"start": 0.0, "end": 1.0, "tokens": [1234], "seek": 0}], + None, + None, + ) + + requests = [ + BatchRequest(audio=self._make_audio(), language="en", use_vad=False) + for _ in range(3) + ] + for req in requests: + self.worker.submit(req) + for req in requests: + req.future.wait(timeout=5) + + for req in requests: + self.assertTrue(req.future.is_set()) + self.assertIsNone(req.error) + self.assertIsNotNone(req.result) + + # Verify the batched encode path was used (not transcribe) + self.mock_transcriber.encode.assert_called() + self.mock_transcriber.transcribe.assert_not_called() + + def test_error_propagation(self): + """Transcriber errors should propagate to the request without crashing the worker.""" + self.mock_transcriber.transcribe.side_effect = RuntimeError("GPU OOM") + + req = BatchRequest(audio=self._make_audio(), language="en", use_vad=False) + self.worker.submit(req) + req.future.wait(timeout=5) + + self.assertTrue(req.future.is_set()) + self.assertIsInstance(req.error, RuntimeError) + self.assertIn("GPU OOM", str(req.error)) + + # Worker should still be alive — submit another request + self.mock_transcriber.transcribe.side_effect = None + self.mock_transcriber.transcribe.return_value = ([MagicMock()], MagicMock()) + + req2 = BatchRequest(audio=self._make_audio(), language="en", use_vad=False) + self.worker.submit(req2) + req2.future.wait(timeout=5) + + self.assertIsNone(req2.error) + self.assertIsNotNone(req2.result) + + def test_worker_stop(self): + """Worker thread should exit cleanly when stop() is called.""" + self.assertTrue(self.worker._thread.is_alive()) + self.worker.stop() + self.assertFalse(self.worker._thread.is_alive()) + + def test_batch_respects_max_size(self): + """Batches should not exceed max_batch_size.""" + self.worker.stop() # Stop the default worker + + observed_batch_sizes = [] + original_process = BatchInferenceWorker._process_batch + + def tracking_process(self_inner, batch): + observed_batch_sizes.append(len(batch)) + original_process(self_inner, batch) + + self.worker = BatchInferenceWorker( + transcriber=self.mock_transcriber, + max_batch_size=2, + batch_window_ms=100, + ) + + self.mock_transcriber.transcribe.return_value = ([MagicMock()], MagicMock()) + + with mock.patch.object( + BatchInferenceWorker, '_process_batch', tracking_process + ): + self.worker.start() + + requests = [ + BatchRequest(audio=self._make_audio(), language="en", use_vad=False) + for _ in range(4) + ] + for req in requests: + self.worker.submit(req) + for req in requests: + req.future.wait(timeout=5) + + for size in observed_batch_sizes: + self.assertLessEqual(size, 2) + self.assertTrue(all(req.future.is_set() for req in requests))