mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 17:15:19 -07:00
refactor(stt): name the Whisper sample rate and window constants
This commit is contained in:
committed by
capy-ai-staging[bot]
parent
51882b5065
commit
86dc46b930
@@ -10,6 +10,10 @@ import numpy as np
|
||||
|
||||
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 .base import (
|
||||
is_model_cached,
|
||||
@@ -337,7 +341,7 @@ class PyTorchSTTBackend:
|
||||
def _transcribe_sync():
|
||||
"""Run synchronous transcription in thread pool."""
|
||||
# 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
|
||||
# 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
|
||||
# generate() runs language detection, i.e. whenever no language
|
||||
# is forced.
|
||||
is_long_form = len(audio) > 30 * 16000
|
||||
is_long_form = len(audio) > WHISPER_WINDOW_SAMPLES
|
||||
processor_kwargs = (
|
||||
{"truncation": False, "padding": "longest", "return_attention_mask": True}
|
||||
if is_long_form
|
||||
@@ -363,7 +367,7 @@ class PyTorchSTTBackend:
|
||||
)
|
||||
inputs = self.processor(
|
||||
audio,
|
||||
sampling_rate=16000,
|
||||
sampling_rate=WHISPER_SAMPLE_RATE,
|
||||
return_tensors="pt",
|
||||
**processor_kwargs,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user