CUDA: enable sparse-fa for dsv4 prefill (again) - #29298
Conversation
The query loop of flash_attn_mask_to_sparse_indices has a runtime trip count, which keeps the unrolled scan over the values of a lane from issuing its loads together. Template the kernel on ncols1 so the loop is bounded at compile time: batch one decodes compile to straight line code and the scan drops from 46 to 17 us at 49k columns on sparse decode shapes.
|
I pushed the compile time unroll of the mask scan query loop and retested on the PR, before/after: sparse decode goes from 54.75 to 25.56 us at 49k |
| #pragma unroll | ||
| for (int q = 0; q < ncols1 && q < q1 - q0 && !selected; ++q) { |
There was a problem hiding this comment.
This loop cannot actually be unrolled beyond the first iteration. I would recommend checking in host code whether or not an out-of-bounds check is needed and then launching a template specialization with/without one accordingly.
There was a problem hiding this comment.
Right, it only unrolled for ncols1 == 1: the kernel is now templated on whether the last query group is partial, chosen in host code, with the column bound hoisted out of the loop, which brings the ncols1 == 8 scan from 110 to 14 branches and halves the batched sparse op at 49k context (586 -> 244 us), decode unchanged.
The query loop of the ncols1 == 8 scan keeps a runtime bound and an early exit, so it does not unroll past its first iteration. Template the kernel on whether the last group of queries is partial, decided on the host from n_queries, and hoist the column bound out of the loop: the loop becomes straight line code and the batched sparse op at 49k context drops from 586 to 244 us.
|
@ServeurpersoCom Can you confirm this change fixes the issue? |
Yes, confirmed on my RTX PRO 6000: sparse decode is back to the b11047 level (54.89 -> 26.01 us at 49k context), the batched sparse op drops from 603 to 237 us, and FLASH_ATTN_EXT passes 3982/3982. |
* CUDA: enable sparse-fa for dsv4 prefill (again) * CUDA: unroll the query loop of the sparse mask scan The query loop of flash_attn_mask_to_sparse_indices has a runtime trip count, which keeps the unrolled scan over the values of a lane from issuing its loads together. Template the kernel on ncols1 so the loop is bounded at compile time: batch one decodes compile to straight line code and the scan drops from 46 to 17 us at 49k columns on sparse decode shapes. * CUDA: pick the out of bounds check of the sparse mask scan in host code The query loop of the ncols1 == 8 scan keeps a runtime bound and an early exit, so it does not unroll past its first iteration. Template the kernel on whether the last group of queries is partial, decided on the host from n_queries, and hoist the column bound out of the loop: the loop becomes straight line code and the batched sparse op at 49k context drops from 586 to 244 us. --------- Co-authored-by: Pascal <admin@serveurperso.com> (cherry picked from commit dc9879c)
* CUDA: enable sparse-fa for dsv4 prefill (again) * CUDA: unroll the query loop of the sparse mask scan The query loop of flash_attn_mask_to_sparse_indices has a runtime trip count, which keeps the unrolled scan over the values of a lane from issuing its loads together. Template the kernel on ncols1 so the loop is bounded at compile time: batch one decodes compile to straight line code and the scan drops from 46 to 17 us at 49k columns on sparse decode shapes. * CUDA: pick the out of bounds check of the sparse mask scan in host code The query loop of the ncols1 == 8 scan keeps a runtime bound and an early exit, so it does not unroll past its first iteration. Template the kernel on whether the last group of queries is partial, decided on the host from n_queries, and hoist the column bound out of the loop: the loop becomes straight line code and the batched sparse op at 49k context drops from 586 to 244 us. --------- Co-authored-by: Pascal <admin@serveurperso.com> (cherry picked from commit dc9879c) (cherry picked from commit e59b5a7)
Overview
#28770 switched off sparse-fa for DSV4/GLM because of the batch-dependent gate. This PR adds a clamp on the condition that allows DSV4 prefill to activate again (no change for decode). Also adds a preference for the wide kernel (i.e. 8 ncols1) if it exists at large batch sizes.
Performance on DGX spark
Additional information
Requirements