Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions app/main_window.py
Original file line number Diff line number Diff line change
Expand Up @@ -768,12 +768,14 @@ def _start_transcription(self, audio_path, session=None):
model_size = self.config.get("transcription", "model_size")
language = self.config.get("transcription", "language")
device = self.config.get("transcription", "device")
batch_size = self.config.get("transcription", "batch_size")

self._transcription_worker = TranscriptionWorker(
audio_path=audio_path,
model_size=model_size,
language=language,
device=device,
batch_size=batch_size,
)
self._transcription_worker.session = session
self._transcription_worker.progress.connect(self._on_transcription_progress)
Expand Down
30 changes: 23 additions & 7 deletions app/transcription/transcriber.py
Original file line number Diff line number Diff line change
Expand Up @@ -142,12 +142,17 @@ class TranscriptionWorker(QThread):

cancelled = pyqtSignal()

def __init__(self, audio_path, model_size="base", language=None, device="cpu"):
def __init__(self, audio_path, model_size="base", language=None, device="cpu",
batch_size=8):
super().__init__()
self.audio_path = audio_path
self.model_size = model_size
self.language = language
self.device = device
# batch_size > 1 uses faster-whisper's BatchedInferencePipeline (VAD-chunked
# parallel decode, typically several times faster). batch_size == 1 keeps the
# classic sequential path (which retains condition_on_previous_text).
self.batch_size = batch_size
self._cancel_requested = False

def cancel(self):
Expand Down Expand Up @@ -180,12 +185,23 @@ def run(self):
self.cancelled.emit()
return

self.progress.emit("Transcribing audio...")
segments_gen, info = model.transcribe(
self.audio_path,
language=self.language,
vad_filter=True,
)
if self.batch_size and self.batch_size > 1:
self.progress.emit(f"Transcribing audio (batched, batch size {self.batch_size})...")
from faster_whisper import BatchedInferencePipeline
pipeline = BatchedInferencePipeline(model=model)
segments_gen, info = pipeline.transcribe(
self.audio_path,
language=self.language,
vad_filter=True,
batch_size=self.batch_size,
)
else:
self.progress.emit("Transcribing audio...")
segments_gen, info = model.transcribe(
self.audio_path,
language=self.language,
vad_filter=True,
)

result = TranscriptResult(
language=info.language,
Expand Down
14 changes: 14 additions & 0 deletions app/ui/settings_dialog.py
Original file line number Diff line number Diff line change
Expand Up @@ -240,6 +240,16 @@ def _setup_ui(self):
)
whisper_form.addRow("Min duration to auto-transcribe:", self.min_duration_spin)

self.batch_size_spin = QSpinBox()
self.batch_size_spin.setRange(1, 16)
self.batch_size_spin.setSpecialValueText("1 (sequential / classic)")
self.batch_size_spin.setToolTip(
"Batched inference decodes VAD-chunked audio in parallel — typically\n"
"several times faster on long recordings. Higher values use more RAM.\n"
"Set to 1 for the classic sequential path (keeps cross-chunk context)."
)
whisper_form.addRow("Batch size:", self.batch_size_spin)

transcription_layout.addWidget(whisper_group)

# Diarization group
Expand Down Expand Up @@ -426,6 +436,9 @@ def _load_settings(self):
min_dur = self.config.get("transcription", "min_duration")
self.min_duration_spin.setValue(min_dur if min_dur else 0)

batch_size = self.config.get("transcription", "batch_size")
self.batch_size_spin.setValue(batch_size if batch_size else 8)

# Diarization
self.diarization_enabled.setChecked(self.config.get("diarization", "enabled"))
self.hf_token_edit.setText(self.config.get("diarization", "hf_token") or "")
Expand Down Expand Up @@ -484,6 +497,7 @@ def _save_and_close(self):
lang = self.language_edit.text().strip()
self.config.set("transcription", "language", lang if lang else None)
self.config.set("transcription", "min_duration", self.min_duration_spin.value())
self.config.set("transcription", "batch_size", self.batch_size_spin.value())

self.config.set("diarization", "enabled", self.diarization_enabled.isChecked())
self.config.set("diarization", "hf_token", self.hf_token_edit.text().strip())
Expand Down
1 change: 1 addition & 0 deletions app/utils/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@
"language": None,
"device": "cpu",
"min_duration": 10,
"batch_size": 8,
},
"diarization": {
"enabled": True,
Expand Down
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@ requires-python = ">=3.10"
# current locked versions all sit under these caps. (issue #4)
dependencies = [
"comtypes>=1.2.0",
"faster-whisper>=1.0.0,<2",
"faster-whisper>=1.1.0,<2",
"numpy>=1.24.0,<3",
"psutil>=5.9.0",
"pyannote-audio>=4.0.0,<5",
Expand Down
2 changes: 1 addition & 1 deletion requirements.txt
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@ PyQt6>=6.6.0
sounddevice>=0.4.6
PyAudioWPatch>=0.2.12
numpy>=1.24.0,<3
faster-whisper>=1.0.0,<2
faster-whisper>=1.1.0,<2
pyannote.audio>=4.0.0,<5
torch>=2.0.0,<3
torchaudio>=2.0.0,<3
Expand Down
57 changes: 54 additions & 3 deletions tests/test_transcriber.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,14 +8,23 @@


class _FwMocks:
"""Build a mocked faster_whisper module returning given segments."""
"""Build a mocked faster_whisper module returning given segments.

Wires both the sequential path (``WhisperModel.transcribe``) and the
batched path (``BatchedInferencePipeline(model=...).transcribe``).
"""

def __init__(self, segments=(), duration=5.0):
segs = list(segments)
self.module = MagicMock()
self.model = MagicMock()
self.module.WhisperModel.return_value = self.model
info = MagicMock(language="en", duration=duration)
self.model.transcribe.return_value = (iter(list(segments)), info)
self.model.transcribe.return_value = (iter(segs), info)
# Batched pipeline: BatchedInferencePipeline(model=...).transcribe(...)
self.pipeline = MagicMock()
self.module.BatchedInferencePipeline.return_value = self.pipeline
self.pipeline.transcribe.return_value = (iter(segs), info)


class TestWhisperModelCache(unittest.TestCase):
Expand All @@ -41,12 +50,13 @@ def test_different_params_create_new_model(self):

class TestRunSegmentMapping(unittest.TestCase):
def _run_worker(self, segments):
# batch_size=1 exercises the classic sequential path (model.transcribe).
fw = _FwMocks(segments=segments)
with patch.dict(sys.modules, {"faster_whisper": fw.module}):
import app.transcription.transcriber as tr
tr._MODEL_CACHE.clear()
worker = tr.TranscriptionWorker(
"a.wav", model_size="base", device="cpu"
"a.wav", model_size="base", device="cpu", batch_size=1
)
results = []
worker.finished.connect(results.append)
Expand All @@ -67,6 +77,47 @@ def test_word_timestamps_not_requested(self):
self.assertNotIn("word_timestamps", kwargs)


class TestBatchedInference(unittest.TestCase):
"""batch_size selects BatchedInferencePipeline (>1) vs the sequential path (1)."""

def _run(self, batch_size, segments=()):
fw = _FwMocks(segments=segments)
with patch.dict(sys.modules, {"faster_whisper": fw.module}):
import app.transcription.transcriber as tr
tr._MODEL_CACHE.clear()
worker = tr.TranscriptionWorker(
"a.wav", model_size="base", device="cpu", batch_size=batch_size
)
results = []
worker.finished.connect(results.append)
worker.run()
return results, fw

def test_batched_path_used_when_batch_size_gt_1(self):
_, fw = self._run(8)
fw.module.BatchedInferencePipeline.assert_called_once_with(model=fw.model)
self.assertEqual(fw.pipeline.transcribe.call_args.kwargs.get("batch_size"), 8)
fw.model.transcribe.assert_not_called()

def test_batched_transcribe_enables_vad(self):
_, fw = self._run(4)
self.assertTrue(fw.pipeline.transcribe.call_args.kwargs.get("vad_filter"))

def test_sequential_path_when_batch_size_1(self):
_, fw = self._run(1)
fw.module.BatchedInferencePipeline.assert_not_called()
fw.model.transcribe.assert_called_once()

def test_batched_segments_mapped_to_result(self):
seg = MagicMock(start=0.0, end=2.0, text=" hi ", avg_logprob=-0.2)
results, _ = self._run(8, segments=[seg])
self.assertEqual(len(results), 1)
self.assertEqual(results[0].segments[0].text, "hi")
self.assertAlmostEqual(
results[0].segments[0].confidence, math.exp(-0.2), places=5
)


class TestTranscriptSegment(unittest.TestCase):

def test_to_dict_without_original_text(self):
Expand Down
2 changes: 1 addition & 1 deletion uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.