Merge pull request #422 from ianwh02/feature/batch-inference
Add cross-client GPU batch inference for faster_whisper backend
This commit is contained in:
@@ -66,6 +66,24 @@ if __name__ == "__main__":
|
|||||||
default=None,
|
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."
|
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()
|
args = parser.parse_args()
|
||||||
|
|
||||||
if args.backend == "tensorrt":
|
if args.backend == "tensorrt":
|
||||||
@@ -92,4 +110,7 @@ if __name__ == "__main__":
|
|||||||
rest_port=args.rest_port,
|
rest_port=args.rest_port,
|
||||||
enable_rest=args.enable_rest,
|
enable_rest=args.enable_rest,
|
||||||
cors_origins=args.cors_origins,
|
cors_origins=args.cors_origins,
|
||||||
|
batch_enabled=args.batch_inference,
|
||||||
|
batch_max_size=args.batch_max_size,
|
||||||
|
batch_window_ms=args.batch_window_ms,
|
||||||
)
|
)
|
||||||
@@ -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))
|
||||||
@@ -14,6 +14,7 @@ from whisper_live.backend.base import ServeClientBase
|
|||||||
class ServeClientFasterWhisper(ServeClientBase):
|
class ServeClientFasterWhisper(ServeClientBase):
|
||||||
SINGLE_MODEL = None
|
SINGLE_MODEL = None
|
||||||
SINGLE_MODEL_LOCK = threading.Lock()
|
SINGLE_MODEL_LOCK = threading.Lock()
|
||||||
|
BATCH_WORKER = None
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -202,6 +203,26 @@ class ServeClientFasterWhisper(ServeClientBase):
|
|||||||
depends on the implementation of the `transcriber.transcribe` method but typically
|
depends on the implementation of the `transcriber.transcribe` method but typically
|
||||||
includes the transcribed text.
|
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:
|
if ServeClientFasterWhisper.SINGLE_MODEL:
|
||||||
ServeClientFasterWhisper.SINGLE_MODEL_LOCK.acquire()
|
ServeClientFasterWhisper.SINGLE_MODEL_LOCK.acquire()
|
||||||
result, info = self.transcriber.transcribe(
|
result, info = self.transcriber.transcribe(
|
||||||
|
|||||||
@@ -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,
|
||||||
|
)
|
||||||
+38
-1
@@ -159,6 +159,7 @@ class TranscriptionServer:
|
|||||||
self.no_voice_activity_chunks = 0
|
self.no_voice_activity_chunks = 0
|
||||||
self.use_vad = True
|
self.use_vad = True
|
||||||
self.single_model = False
|
self.single_model = False
|
||||||
|
self.batch_config = None
|
||||||
|
|
||||||
def initialize_client(
|
def initialize_client(
|
||||||
self, websocket, options, faster_whisper_custom_model_path,
|
self, websocket, options, faster_whisper_custom_model_path,
|
||||||
@@ -277,6 +278,18 @@ class TranscriptionServer:
|
|||||||
)
|
)
|
||||||
|
|
||||||
logging.info("Running faster_whisper backend.")
|
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:
|
except Exception as e:
|
||||||
logging.error(e)
|
logging.error(e)
|
||||||
return
|
return
|
||||||
@@ -415,13 +428,25 @@ class TranscriptionServer:
|
|||||||
cache_path="~/.cache/whisper-live/",
|
cache_path="~/.cache/whisper-live/",
|
||||||
rest_port=8000,
|
rest_port=8000,
|
||||||
enable_rest=False,
|
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.
|
Run the transcription server.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
host (str): The host address to bind the server.
|
host (str): The host address to bind the server.
|
||||||
port (int): The port number 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.cache_path = cache_path
|
||||||
self.client_manager = ClientManager(max_clients, max_connection_time)
|
self.client_manager = ClientManager(max_clients, max_connection_time)
|
||||||
@@ -430,6 +455,18 @@ class TranscriptionServer:
|
|||||||
raise ValueError(f"Custom faster_whisper model '{faster_whisper_custom_model_path}' is not a valid path or HuggingFace model.")
|
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):
|
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.")
|
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 single_model:
|
||||||
if faster_whisper_custom_model_path or whisper_tensorrt_path:
|
if faster_whisper_custom_model_path or whisper_tensorrt_path:
|
||||||
logging.info("Custom model option was provided. Switching to single model mode.")
|
logging.info("Custom model option was provided. Switching to single model mode.")
|
||||||
|
|||||||
Reference in New Issue
Block a user