Repository navigation
Fix Whisper beam_indices so compute_transition_scores does not select the beam twice - #48623
speedsharmaai wants to merge 2 commits into
Conversation
… the beam twice
_postprocess_outputs already gathers scores from the beam each step was
generated on, so the stacked scores hold one row per returned sequence
rather than one row per beam. The returned beam_indices still carried the
beam ancestry, so compute_transition_scores multiplied it by vocab_size
and gathered out of bounds:
RuntimeError: index 156073 is out of bounds for dimension 0 with size 51865
It only raised when the winning beam's ancestry passed through a beam
other than 0; sequences that stayed on beam 0 gathered in range and
returned plausible but arbitrary scores.
Point beam_indices at the rows of the scores that are actually returned,
keeping -1 for steps past the end of a sequence so they are still masked.
|
Thank you for your contribution 🤗! CI Security Gate — automatic approval blockedThis PR was not automatically approved for CI because the security gate failed. Possible reasons:
See the workflow run for the exact violations. A maintainer can review and manually approve CI if a finding is a false positive. |
|
For whoever triages the security gate: none of the listed reasons apply here — the PR changes two files, The so the bandit step never ran. |
vasqu
left a comment
There was a problem hiding this comment.
Some small initial comments from my side. Whisper is unique so this is totally fine imo, the comments are more on style
| if key in ["sequences", "beam_indices", "token_timestamps"]: | ||
| outputs[key] = torch.stack([v[key] for v in seek_outputs], dim=0).to(device) | ||
| if key == "beam_indices" and "scores" in seek_outputs[0]: | ||
| # `split_by_batch_index` already gathered `scores` from the beam each step was generated on, | ||
| # so the stacked scores hold one row per returned sequence instead of one row per beam. Point | ||
| # `beam_indices` at those rows, otherwise `compute_transition_scores` selects the beam a second | ||
| # time and gathers out of bounds. `-1` marks steps past the end of a sequence and is kept. | ||
| beam_indices = outputs[key] | ||
| rows = torch.arange(beam_indices.shape[0], device=device).unsqueeze(1).expand_as(beam_indices) | ||
| outputs[key] = torch.where(beam_indices == -1, beam_indices, rows) |
There was a problem hiding this comment.
very nitpicky but could we split this if then for beam as its own case + maybe reference the issue
There was a problem hiding this comment.
feel like we could potentially also shorten the comment, im getting claude vibes (sorry if thats not the case!)
| self.assertEqual(output.beam_indices.shape[0], input_features.shape[0] * 3) | ||
| self.assertEqual(output.sequences_scores.shape[0], input_features.shape[0] * 3) | ||
|
|
||
| def test_beam_search_transition_scores(self): |
There was a problem hiding this comment.
Would rather like a style of test_model_forward
| self.assertEqual(output.sequences_scores.shape[0], input_features.shape[0] * 3) | ||
|
|
||
| def test_beam_search_transition_scores(self): | ||
| config, input_dict = self.model_tester.prepare_config_and_inputs() |
There was a problem hiding this comment.
missing docstring and reference to the issue/PR
|
The bandit is wonky, it might be also due to your branch being not up to date with main. But anyways false flag |
|
We filed #48621, and this is the hardware where the bad index doesn't raise but hangs the machine, so we ran your patch against it. transformers 5.14.1,
Before the patch this path leaves the HIP stream un-drained and the next device-to-host copy spins forever — the process can't be killed and the machine needs a hard power-off (pytorch/pytorch#196377). After it, One thing worth checking in the new test: with Possibly out of scope: |
A run whose ancestry never leaves beam 0 gathers in bounds even without the fix, so the assertions would pass on unpatched code. Capture the ancestry before _stack_split_outputs rewrites it and require at least one index above 0.
|
Thanks for running it on the hardware that hangs — that ROCm detail is worse than what I could see locally. You're right about the test. It did happen to fail on unpatched code here (the tiny random tester model's ancestry does leave beam 0), but nothing in the test said so, and a model or seed change could have made it vacuous without anyone noticing. Pushed 2dd2996: it now captures the ancestry before On longform: agreed, the multi-segment path returns @vasqu on the bandit flag — I left the branch on its original base for now. Rebasing onto main pulls in a workflow file change that my token can't push. |
|
[For maintainers] Suggested jobs to run (before merge) run-slow: whisper |
CI recapDashboard: View test results in Grafana |
What does this PR do?
Fixes #48621.
compute_transition_scores()raises for Whisper withnum_beams > 1:The beam axis is selected twice.
_postprocess_outputsalready gathersscoresfrom the beam each step was generated on (added in #32336, fixing #32246):So the stacked
scores[t]hold one row per returned sequence, not one row per beam. The returnedbeam_indicesstill carried the beam ancestry, andcompute_transition_scoresmultiplies it byvocab_sizeto index into the flattened scores — reaching rows that no longer exist.It only raises when the winning beam's ancestry passes through a beam other than 0. Sequences that stayed on beam 0 gather
0 * vocab_size + token, which is in range, and return plausible-looking but arbitrary scores. That is what made this look data-dependent.Fix
Point
beam_indicesat the rows of thescoresthat are actually returned, so the two outputs agree.-1is preserved, so steps past the end of a sequence are still masked bycompute_transition_scores.With this, the usual beam-search invariant holds again —
transition_scores.sum(-1)reproducessequences_scoresexactly:The remap is guarded on
scoresbeing present, sooutput_scores=Falsestill returns the raw ancestry.On the alternative
@vasqu noted on the issue that the general rule is to follow what core
generatedoes. The other way to fix this is to stop gatheringscoresby beam in_postprocess_outputsand return the plain(batch_size * num_beams, vocab_size)rows that core beam search returns, which would makebeam_indicescorrect as-is.I did not go that way because it reverts the deliberate behaviour from #32336 and changes what
outputs.scoresmeans for every Whisper user since v4.45 — that felt like a call for maintainers rather than a drive-by fix. Happy to redo it that way if you prefer the core semantics.Tests
tests/models/whisper/test_modeling_whisper.py::WhisperModelTest::test_beam_search_transition_scoresasserts that each transition score comes from the row of its own sequence, and that the scores sum back tosequences_scores. Onmainit fails with the reportedRuntimeError.tests/models/whisper/test_modeling_whisper.pypasses: 370 passed, 346 skipped.Who can review?
@eustlb @vasqu (Whisper / generation)