metal : add sparse FA - #28098
Conversation
cd3a488 to
6d68373
Compare
6d68373 to
30a41e3
Compare
|
/bot review |
Automated code reviewReview complete. Here is my review of the Metal sparse Flash Attention changes. Summary: the approach (gather finite mask entries into per-row index lists, then run the existing vec kernels over gathered rows) is sound and matches the shape of the CUDA sparse FA from #27970. The kernel gather logic in Blocking (point 1) Will slow the review (point 2) All sparse (point 3) The (point 4) Hardening: the gathered Nits (point 5) (point 6) (point 7) (point 8) Good to see the perf table in the description and the This review was generated automatically by pi coding agent using |
|
M1 Ultra results: |
- add kernel_flash_attn_ext_vec_idx: compacts finite mask entries into a per-row index list (Hillis-Steele scan, one threadgroup per row) - extend vec FA kernel with optional sparse index gathering (FC slot 5) - add host-side gate: sparse path when n_kv_max > 0, mask present, supported head sizes / KV types, n_kv_max <= 4096 - new buffer region extra_idx for the index list - pipeline getter extended with has_sparse param - add test cases: head sizes, quant types, nb>1, nr23 variants, sinks, ALiBi, softcap, permute, v_view_of_k, no-mask fallback Note: multi-row (nb*nr23[1] > 1) cases still failing - rid mapping in the store phase needs revisiting for the sparse path. Assisted-by: pi:llama.cpp/Qwen3.8-27B
- kernel_flash_attn_ext_vec_idx: mask param is half* but nb31 is a byte stride, so the per-row mask offset was scaled by 2x; cast to char* before applying the byte strides - kernel_flash_attn_ext_vec: sparse pidx param is char* so the per-row element offset was under-scaled by sizeof(int); scale it by sizeof(int) to get the correct byte offset - fixes the multi-row (nb*nr23[1] > 1) sparse flash attention failures Assisted-by: pi:llama.cpp/DeepSeek-v4-0731
The idx kernel previously read the mask row twice: once to count the finite entries (for the prefix scan) and again to recover their positions. Since the kernel is memory-bound, this doubled the mask traffic. Keep the finite positions in a per-thread register array during the count pass and write them out directly, avoiding the second mask read. A dense mask with more than NLOCAL finite entries in a slice falls back to re-reading the mask to write the remaining positions. Assisted-by: pi:llama.cpp/DeepSeek-v4-0731
Measure the sparse vec FA kernel across KV sizes, n_kv_max hints and batch
sizes. Run with:
./build/bin/test-backend-ops -b MTL0 -o FLASH_ATTN_EXT -p "n_kv_max=[1-9]" perf
Assisted-by: pi:llama.cpp/DeepSeek-v4-0731
688e333 to
1d818e8
Compare
|
FYI @ggerganov, for the qwen4exp TODO you added: enabling sparse FA there (passing
Build is master (de8656b) + #27836 (MTP). Small samples size of course, but happy to test anything further or raise a PR if it helps. AI usage disclosure: YES (just for testing this), claude:opus-5 *Edit: PR: #28349 |
* metal : support n_kv_max sparse mask hint in flash attention vec kernel
- add kernel_flash_attn_ext_vec_idx: compacts finite mask entries into
a per-row index list (Hillis-Steele scan, one threadgroup per row)
- extend vec FA kernel with optional sparse index gathering (FC slot 5)
- add host-side gate: sparse path when n_kv_max > 0, mask present,
supported head sizes / KV types, n_kv_max <= 4096
- new buffer region extra_idx for the index list
- pipeline getter extended with has_sparse param
- add test cases: head sizes, quant types, nb>1, nr23 variants,
sinks, ALiBi, softcap, permute, v_view_of_k, no-mask fallback
Note: multi-row (nb*nr23[1] > 1) cases still failing - rid mapping
in the store phase needs revisiting for the sparse path.
Assisted-by: pi:llama.cpp/Qwen3.8-27B
* metal : fix sparse flash attention row addressing
- kernel_flash_attn_ext_vec_idx: mask param is half* but nb31 is a byte
stride, so the per-row mask offset was scaled by 2x; cast to char*
before applying the byte strides
- kernel_flash_attn_ext_vec: sparse pidx param is char* so the per-row
element offset was under-scaled by sizeof(int); scale it by sizeof(int)
to get the correct byte offset
- fixes the multi-row (nb*nr23[1] > 1) sparse flash attention failures
Assisted-by: pi:llama.cpp/DeepSeek-v4-0731
* cont : use sparse vec FA for prefill
* metal : single-pass flash attention sparse index compaction
The idx kernel previously read the mask row twice: once to count the finite
entries (for the prefix scan) and again to recover their positions. Since the
kernel is memory-bound, this doubled the mask traffic.
Keep the finite positions in a per-thread register array during the count
pass and write them out directly, avoiding the second mask read. A dense
mask with more than NLOCAL finite entries in a slice falls back to re-reading
the mask to write the remaining positions.
Assisted-by: pi:llama.cpp/DeepSeek-v4-0731
* tests : add perf cases for sparse flash attention prefill
Measure the sparse vec FA kernel across KV sizes, n_kv_max hints and batch
sizes. Run with:
./build/bin/test-backend-ops -b MTL0 -o FLASH_ATTN_EXT -p "n_kv_max=[1-9]" perf
Assisted-by: pi:llama.cpp/DeepSeek-v4-0731
* qwen4 : enable sparse attention
* cont : adjust nsg
* cont : sync test-backend-ops
* cont : disable Qwen4 for now
* cont : clean-up + tests
Overview
cont #27970
Add Metal support for sparse Flash Attention.
Results for DSv4 on M2 Ultra:
Requirements