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__)
|
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,
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user