mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-04 09:35:16 -07:00
fix(backend): keep the pad-to-30s Whisper path for clips under 30s
With padding="longest" the encoder rejects inputs shorter than 3000 mel frames when generate() has to detect the language first, so every short clip without a forced language failed with "Whisper expects the mel input features to be of length 3000". Use the long-form feature-extractor and generate() options only when the audio exceeds one 30s window.
This commit is contained in:
committed by
capy-ai-staging[bot]
parent
7575a65e7f
commit
51882b5065
@@ -348,13 +348,24 @@ class PyTorchSTTBackend:
|
||||
# are required for long-form transcription. Without them the
|
||||
# feature extractor silently truncates to 30s (Whisper's native
|
||||
# window) and audio past that point is dropped.
|
||||
#
|
||||
# Only use them for audio longer than one 30s window. For shorter
|
||||
# clips the default pad-to-30s path is needed: with
|
||||
# padding="longest" the encoder rejects the shorter mel input
|
||||
# ("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
|
||||
processor_kwargs = (
|
||||
{"truncation": False, "padding": "longest", "return_attention_mask": True}
|
||||
if is_long_form
|
||||
else {}
|
||||
)
|
||||
inputs = self.processor(
|
||||
audio,
|
||||
sampling_rate=16000,
|
||||
return_tensors="pt",
|
||||
truncation=False,
|
||||
padding="longest",
|
||||
return_attention_mask=True,
|
||||
**processor_kwargs,
|
||||
)
|
||||
inputs = inputs.to(self.device)
|
||||
|
||||
@@ -369,12 +380,13 @@ class PyTorchSTTBackend:
|
||||
if language:
|
||||
generate_kwargs["language"] = language
|
||||
generate_kwargs["task"] = "transcribe"
|
||||
if is_long_form:
|
||||
generate_kwargs["attention_mask"] = inputs["attention_mask"]
|
||||
generate_kwargs["return_timestamps"] = True
|
||||
|
||||
with torch.no_grad():
|
||||
predicted_ids = self.model.generate(
|
||||
inputs["input_features"],
|
||||
attention_mask=inputs["attention_mask"],
|
||||
return_timestamps=True,
|
||||
**generate_kwargs,
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user