Repository navigation
Conversation
non_causal_chunked_attn masked padded key columns in the ragged last chunk with -1e-9, which is negligible next to typical dot-product logits and barely changes the softmax. Real queries in that chunk leaked attention mass onto the phantom zero-vector padding keys instead of concentrating it on the real keys, which then get trimmed away, so the per-key importance score used for eviction came back systematically under-counted whenever the sequence length was not an exact multiple of chunk_size. Mask with torch.finfo(dots.dtype).min instead, matching the pattern kvzip_press.py already uses for the same kind of key masking. Fixes NVIDIA#295 馃馃馃 Signed-off-by: Amir Fathi <amirfathi.me@gmail.com>
Collaborator
Collaborator
|
@vnchari, do you have bandwidth to check this fix on the press you contributed ? |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
PR description
non_causal_chunked_attnmasks padded key columns in the ragged last chunk with-1e-9before the softmax (kvpress/presses/non_causal_attention_press.py:91).-1e-9is close to zero, so it barely suppresses the padded columns: a real query in the last chunk leaks attention mass onto the phantom zero-vector padding keys instead of concentrating it on the real keys, and that leaked mass is then discarded when the result is trimmed back toSpositions. This happens for every sequence length that is not an exact multiple ofchunk_size, i.e. essentially always, and it understates the per-key importance scoreNonCausalAttnPress/CompactorPressuse to pick which tokens to evict.The fix masks with
torch.finfo(dots.dtype).mininstead, matching the patternkvzip_press.pyalready uses for the same kind of key masking, so it stays correct across dtypes (fp32/fp16/bf16).Fixes #295
Verification
test_non_causal_chunked_attn_does_not_leak_mass_to_padding_keys: a 4-query ragged chunk (2 real, 2 padded) must return per-key scores summing tochunk_size(every query's full softmax unit of mass lands on the 2 real keys). Fails on main (returns ~2.3), passes on this branch.pytest tests/presses/(729 passed, 3 GPU-only skipped) andmypy ./flake8/black/isortall clean, run in a CPU-only container since this box has no GPU; theTestworkflow needs anlinux-amd64-gpu-l4-latest-1runner I could not reproduce locally.馃馃馃