mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 17:15:19 -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
|
# are required for long-form transcription. Without them the
|
||||||
# feature extractor silently truncates to 30s (Whisper's native
|
# feature extractor silently truncates to 30s (Whisper's native
|
||||||
# window) and audio past that point is dropped.
|
# 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(
|
inputs = self.processor(
|
||||||
audio,
|
audio,
|
||||||
sampling_rate=16000,
|
sampling_rate=16000,
|
||||||
return_tensors="pt",
|
return_tensors="pt",
|
||||||
truncation=False,
|
**processor_kwargs,
|
||||||
padding="longest",
|
|
||||||
return_attention_mask=True,
|
|
||||||
)
|
)
|
||||||
inputs = inputs.to(self.device)
|
inputs = inputs.to(self.device)
|
||||||
|
|
||||||
@@ -369,12 +380,13 @@ class PyTorchSTTBackend:
|
|||||||
if language:
|
if language:
|
||||||
generate_kwargs["language"] = language
|
generate_kwargs["language"] = language
|
||||||
generate_kwargs["task"] = "transcribe"
|
generate_kwargs["task"] = "transcribe"
|
||||||
|
if is_long_form:
|
||||||
|
generate_kwargs["attention_mask"] = inputs["attention_mask"]
|
||||||
|
generate_kwargs["return_timestamps"] = True
|
||||||
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
predicted_ids = self.model.generate(
|
predicted_ids = self.model.generate(
|
||||||
inputs["input_features"],
|
inputs["input_features"],
|
||||||
attention_mask=inputs["attention_mask"],
|
|
||||||
return_timestamps=True,
|
|
||||||
**generate_kwargs,
|
**generate_kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user