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:
nox
2026-10-04 00:01:34 +00:00
committed by capy-ai-staging[bot]
co-authored by Claude Opus 4.7
parent ae300c5316
commit e61c85365b
+10 -1
View File
@@ -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,
) )