Skip to content

Fix NaN in FlashInfer decode kernel for short KV caches (#156) - #158

Open
mmjerge wants to merge 3 commits into
awslabs:mainfrom
mmjerge:fix/156-flashinfer-decode
Open

mmjerge wants to merge 3 commits into
awslabs:mainfrom
mmjerge:fix/156-flashinfer-decode

Conversation

@mmjerge

@mmjerge mmjerge commented Oct 6, 2026

Copy link
Copy Markdown
Collaborator

The vendored decode kernel returns NaN output and attention weights when kv_len < TILE_SIZE (128 for head_size 128, 256 for head_size 64). It splits the KV axis over BDZ chunks, and a chunk that has not seen a valid key still has st.m == -inf. The online softmax rescale then computes exp2(-inf - -inf) = NaN, which poisons the cross-chunk merge.

This is not limited to toy shapes. Every decode right after a short prompt hits it. In RL rollouts on math prompts (a few hundred tokens) it gave NaN logits and a device-side assert in multinomial.

Fix: skip the rescale while the chunk is still empty. merge() already handles a partner with m == -inf.

Tests: TestDecodeKernelShortCache checks decode with attention weights against the eager reference for kv_len from 2 to 300, head_size 64 and 128, fp16 and bf16. Run on an L40S:

  • main: 10 of 18 fail (all kv_len below the tile size)
  • this PR: 18 of 18 pass, and test_flashinfer_wrapper.py + test_attn_weights.py pass in full (72). This includes test_property_3 and test_property_4, which fail on main for the same reason.

The first commit adds the tests alone, so you can see them fail before the fix.

Closes #156.


By submitting this pull request, I confirm that you can use, modify, copy, and redistribute this contribution, under the terms of your choice.

… sizes

Two failures in the vendored FlashInfer decode path, both on q_len = 1:
- kv_len < TILE_SIZE with attention weights returns NaN (awslabs#156)
- GQA group sizes outside {1, 2, 4, 8} without attention weights raise
  'Unsupported group_size' from FlashInfer's SingleDecodeWithKVCache

These tests fail on main; the fix follows in the next commit.
…GQA group sizes

- tiled_decode_attention_kernel: a tz chunk that has not seen a valid key
  still has st.m == -inf, and the online-softmax rescale computed
  exp2(-inf - -inf) = NaN, which poisoned the cross-BDZ merge. Hit for any
  kv_len < TILE_SIZE (128 for head_size 128, 256 for 64), i.e. the first
  decode steps after a short prompt. Skip the rescale while the chunk is
  empty. Fixes awslabs#156.
- FlashInfer's library decode only instantiates GQA group sizes
  1, 2, 3, 4, 6, 8 and throws for others. Qwen2.5-7B (28 / 4 = 7) crashed
  in every decode step without attention weights. Check the group size
  before dispatching and catch the exception, falling back to the vendored
  tiled kernel.
The tiled fallback gives outputs that do not match eager for group sizes
5 and 7, so falling back is worse than the current loud failure. Tracked
separately.
@mmjerge
mmjerge requested review from mseeger and vihangp October 6, 2026 23:07
@mseeger

mseeger commented Oct 7, 2026

Copy link
Copy Markdown
Contributor

Is this something we could contribute back, or does it only affect our use case?

@mmjerge

mmjerge commented Oct 7, 2026

Copy link
Copy Markdown
Collaborator Author

@mseeger do you mean contribute to the original source code?

@mmjerge

mmjerge commented Oct 7, 2026

Copy link
Copy Markdown
Collaborator Author

The bug is in our own decode kernel (tiled_decode_attention_kernel, code in keys_values/csrc), not in FlashInfer itself. It only uses FlashInfer's primitives.

@mseeger mseeger left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This looks good to me. But @vihangp should have a look.

mmjerge added a commit to mmjerge/keys_values that referenced this pull request Oct 7, 2026
Kept only the awslabs#156 NaN fix (same as PR awslabs#158). The
fallback to the tiled kernel for group sizes outside FlashInfer's
{1,2,3,4,6,8} produced outputs that do not match the eager reference for
group sizes 5 and 7, so it traded a loud crash for silently wrong
attention. Qwen2.5-7B (group 7) dense decode without attention weights
will raise again until the fallback kernel is fixed.
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.

FlashInfer decode kernel returns NaN attention weights for kv_len < 128

2 participants