Skip to content

Fix NonCausalAttnPress key-padding mask value 馃馃馃 - #296

Open
AmirF194 wants to merge 1 commit into
NVIDIA:mainfrom
AmirF194:fix/295-non-causal-mask-value
Open

AmirF194 wants to merge 1 commit into
NVIDIA:mainfrom
AmirF194:fix/295-non-causal-mask-value

Conversation

@AmirF194

@AmirF194 AmirF194 commented Sep 28, 2026 •

Copy link
Copy Markdown
Contributor

PR description

non_causal_chunked_attn masks padded key columns in the ragged last chunk with -1e-9 before the softmax (kvpress/presses/non_causal_attention_press.py:91). -1e-9 is 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 to S positions. This happens for every sequence length that is not an exact multiple of chunk_size, i.e. essentially always, and it understates the per-key importance score NonCausalAttnPress/CompactorPress use to pick which tokens to evict.

The fix masks with torch.finfo(dots.dtype).min instead, matching the pattern kvzip_press.py already uses for the same kind of key masking, so it stays correct across dtypes (fp32/fp16/bf16).

Fixes #295

Verification

  • Added 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 to chunk_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) and mypy . / flake8 / black / isort all clean, run in a CPU-only container since this box has no GPU; the Test workflow needs an linux-amd64-gpu-l4-latest-1 runner I could not reproduce locally.
  • Did not check whether this changes downstream generation quality on a real model/benchmark, only the masking math itself.
    馃馃馃

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>
@copy-pr-bot

copy-pr-bot Bot commented Sep 28, 2026

Copy link
Copy Markdown

This pull request requires additional validation before any workflows can run on NVIDIA's runners.

Pull request vetters can view their responsibilities here.

Contributors can view more details about this message here.

@AmirF194 AmirF194 changed the title Fix NonCausalAttnPress key-padding mask value 馃馃馃 Fix NonCausalAttnPress key-padding mask value Sep 29, 2026
@AmirF194 AmirF194 changed the title Fix NonCausalAttnPress key-padding mask value Fix NonCausalAttnPress key-padding mask value 馃馃馃 Sep 29, 2026
@SimJeg

SimJeg commented Sep 29, 2026

Copy link
Copy Markdown
Collaborator

Thanks @AmirF194, @vnchari what do you think ?

@SimJeg

SimJeg commented Oct 5, 2026

Copy link
Copy Markdown
Collaborator

@vnchari, do you have bandwidth to check this fix on the press you contributed ?

@SimJeg SimJeg self-assigned this Oct 5, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

NonCausalAttnPress masks padded keys with -1e-9 instead of a large negative value 馃馃馃

2 participants