mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 17:15:19 -07:00
fix(backend): enable long-form Whisper transcription on PyTorch path
The PyTorch Whisper transcribe() called the HF processor without truncation=False and model.generate() without return_timestamps=True. With those defaults, WhisperFeatureExtractor silently truncates inputs to 30 s (Whisper's native receptive field), so any dictation longer than ~30 s lost its tail. Setting truncation=False + padding="longest" + return_attention_mask=True on the processor, then forwarding the attention mask plus return_timestamps=True to generate(), flips HF Whisper into long-form mode: autoregressive decoding over rolling 30 s windows. Verified by round-tripping a 56.6 s Kokoro TTS sample through /transcribe — full text returned including content past the 30 s mark; previously the transcript was cut off roughly halfway through. MLX backend (mlx_backend.py) is intentionally unchanged: mlx_audio.stt's generate() already implements rolling-window long-form transcription with condition_on_previous_text in the upstream library, so it does not have the same bug. The HF-only kwargs added here would also break the MLX call signature. Co-Authored-By: Claude Opus 4.7 <[email protected]>
This commit is contained in:
committed by
capy-ai-staging[bot]
co-authored by
Claude Opus 4.7
parent
ae300c5316
commit
e61c85365b
@@ -343,11 +343,18 @@ class PyTorchSTTBackend:
|
|||||||
# state — forcing offline here (issue #462) broke online users
|
# state — forcing offline here (issue #462) broke online users
|
||||||
# whose `get_decoder_prompt_ids` / tokenizer calls issue
|
# whose `get_decoder_prompt_ids` / tokenizer calls issue
|
||||||
# legitimate metadata lookups.
|
# legitimate metadata lookups.
|
||||||
# Process audio
|
# Process audio.
|
||||||
|
# truncation=False + padding="longest" + return_attention_mask=True
|
||||||
|
# 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.
|
||||||
inputs = self.processor(
|
inputs = self.processor(
|
||||||
audio,
|
audio,
|
||||||
sampling_rate=16000,
|
sampling_rate=16000,
|
||||||
return_tensors="pt",
|
return_tensors="pt",
|
||||||
|
truncation=False,
|
||||||
|
padding="longest",
|
||||||
|
return_attention_mask=True,
|
||||||
)
|
)
|
||||||
inputs = inputs.to(self.device)
|
inputs = inputs.to(self.device)
|
||||||
|
|
||||||
@@ -364,6 +371,8 @@ class PyTorchSTTBackend:
|
|||||||
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