mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 17:15:19 -07:00
fix(backend): pass language/task directly to Whisper generate
This commit is contained in:
committed by
capy-ai-staging[bot]
parent
e61c85365b
commit
7575a65e7f
@@ -358,15 +358,17 @@ class PyTorchSTTBackend:
|
|||||||
)
|
)
|
||||||
inputs = inputs.to(self.device)
|
inputs = inputs.to(self.device)
|
||||||
|
|
||||||
# Generate transcription
|
# Generate transcription.
|
||||||
# If language is provided, force it; otherwise let Whisper auto-detect
|
# If language is provided, force it; otherwise let Whisper
|
||||||
|
# auto-detect. Pass language/task directly to generate() instead
|
||||||
|
# of building forced_decoder_ids — get_decoder_prompt_ids defaults
|
||||||
|
# to no_timestamps=True, which injects <|notimestamps|> and
|
||||||
|
# disables the timestamp tokens that return_timestamps=True (and
|
||||||
|
# therefore long-form decoding) depend on.
|
||||||
generate_kwargs = {}
|
generate_kwargs = {}
|
||||||
if language:
|
if language:
|
||||||
forced_decoder_ids = self.processor.get_decoder_prompt_ids(
|
generate_kwargs["language"] = language
|
||||||
language=language,
|
generate_kwargs["task"] = "transcribe"
|
||||||
task="transcribe",
|
|
||||||
)
|
|
||||||
generate_kwargs["forced_decoder_ids"] = forced_decoder_ids
|
|
||||||
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
predicted_ids = self.model.generate(
|
predicted_ids = self.model.generate(
|
||||||
|
|||||||
Reference in New Issue
Block a user