refactor(stt): name the Whisper sample rate and window constants

This commit is contained in:
jamiepine
2026-10-04 00:01:34 +00:00
committed by capy-ai-staging[bot]
parent 51882b5065
commit 86dc46b930
+7 -3
View File
@@ -10,6 +10,10 @@ import numpy as np
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
WHISPER_SAMPLE_RATE = 16000
# Whisper encodes one 30s window per pass; longer audio needs sequential long-form decoding.
WHISPER_WINDOW_SAMPLES = 30 * WHISPER_SAMPLE_RATE
from . import TTSBackend, STTBackend, LANGUAGE_CODE_TO_NAME, WHISPER_HF_REPOS from . import TTSBackend, STTBackend, LANGUAGE_CODE_TO_NAME, WHISPER_HF_REPOS
from .base import ( from .base import (
is_model_cached, is_model_cached,
@@ -337,7 +341,7 @@ class PyTorchSTTBackend:
def _transcribe_sync(): def _transcribe_sync():
"""Run synchronous transcription in thread pool.""" """Run synchronous transcription in thread pool."""
# Load audio # Load audio
audio, _sr = load_audio(audio_path, sample_rate=16000) audio, _sr = load_audio(audio_path, sample_rate=WHISPER_SAMPLE_RATE)
# Inference runs with the process's default HF_HUB_OFFLINE # Inference runs with the process's default HF_HUB_OFFLINE
# state — forcing offline here (issue #462) broke online users # state — forcing offline here (issue #462) broke online users
@@ -355,7 +359,7 @@ class PyTorchSTTBackend:
# ("expects the mel input features to be of length 3000") when # ("expects the mel input features to be of length 3000") when
# generate() runs language detection, i.e. whenever no language # generate() runs language detection, i.e. whenever no language
# is forced. # is forced.
is_long_form = len(audio) > 30 * 16000 is_long_form = len(audio) > WHISPER_WINDOW_SAMPLES
processor_kwargs = ( processor_kwargs = (
{"truncation": False, "padding": "longest", "return_attention_mask": True} {"truncation": False, "padding": "longest", "return_attention_mask": True}
if is_long_form if is_long_form
@@ -363,7 +367,7 @@ class PyTorchSTTBackend:
) )
inputs = self.processor( inputs = self.processor(
audio, audio,
sampling_rate=16000, sampling_rate=WHISPER_SAMPLE_RATE,
return_tensors="pt", return_tensors="pt",
**processor_kwargs, **processor_kwargs,
) )