diff --git a/docs/development/DSV4-vulkan-lightning-indexer-progress.md b/docs/development/DSV4-vulkan-lightning-indexer-progress.md new file mode 100644 index 000000000000..d19fbd6903aa --- /dev/null +++ b/docs/development/DSV4-vulkan-lightning-indexer-progress.md @@ -0,0 +1,158 @@ +# DeepSeek V4 Vulkan Lightning Indexer progress + +This is a restart note for the Strix Halo Lightning Indexer optimization. It is a development scratch pad and can be removed before the final PR. + +## Repository state + +- Main repository: `/home/jaap/Projects/git/llama.cpp` +- Optimization worktree: `/tmp/llama-strix-beta-bench` +- Branch: `strix-halo-vulkan-lightning-indexer` +- Base commit: `316c72ee9eab590f5891089d3b6bfc0d01d00d19` +- Base branch: Nathan's `strix-halo-vulkan-beta` +- Decode microbench work is stored in the main worktree as `stash@{0}: On strix-halo-vulkan: wip: DSV4 decode microbench depth matrix`. +- The Indexer changes are uncommitted. Do not commit without explicit user approval. An assisted commit needs an `Assisted-by:` trailer. +- Do not run builds and GPU benchmarks together. The APU shares its power and memory-bandwidth budget. +- GPU commands need sandbox escalation. + +## Objective and result + +After sparse prefill attention was flattened, the context-dependent Lightning Indexer became the next prefill bottleneck. The old cooperative-matrix shader used one wave64 subgroup per workgroup, processed one 16-key tile, and loaded one query head at a time. + +The new wide pipeline uses eight wave64 subgroups per workgroup. Each subgroup processes a separate 16-key tile, so one workgroup covers 128 keys. It stages four query heads and their weights together, reuses them across all eight subgroups, and uses subgroup-scoped synchronization between cooperative-matrix result stores. A workgroup barrier remains between four-head groups because all subgroups reuse the shared query storage. + +The optimized shader requires 512 workgroup invocations and 64 KiB shared memory. Pipeline creation is capability-based. Devices without those limits use a one-wave, one-head cooperative-matrix specialization. The scalar implementation remains the fallback when cooperative matrices are unavailable. The decode-specific cooperative-matrix pipeline is unchanged. + +At the 32k-equivalent prefill microbench shape: + +| Version | Time per layer | Throughput | +| --- | ---: | ---: | +| Baseline | 50.51 ms | 5.83 TFLOPS | +| Optimized | 31.81 ms | 9.25 TFLOPS | + +This is a 37.0% reduction in Lightning Indexer kernel time. + +The canonical 32k llama-bench improved from 209.45 to 216.32 tokens/s. Total Vulkan time fell from 9.73729 to 9.42448 seconds. Total Lightning Indexer time fell from 1.11860 to 0.697276 seconds. Sparse attention and top-K were effectively unchanged. + +## Changed files + +- `ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp`: parameterizes the shader, adds the eight-wave four-head implementation, and remains usable for the small fallback. +- `ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp`: generates wide `N_WAVES=8`, `HEADS_PER_TILE=4` and small `N_WAVES=1`, `HEADS_PER_TILE=1` variants. +- `ggml/src/ggml-vulkan/ggml-vulkan.cpp`: creates and selects the capability-gated wide pipeline and the small cooperative-matrix fallback. +- `tests/test-backend-ops.cpp`: adds 127, 128, and 129-key correctness boundaries and PP2048 performance shapes through 512k simulated source context. + +## Performance data + +The performance rows model the actual PP2048 Indexer shapes after the source context is filled in 2048-token batches. `kv=8704` is the measured shape near 32k source context. It differs from 32768 because the Indexer compresses source tokens into rows. + +| Source depth | `kv` | Baseline | Optimized | Reduction | +| ---: | ---: | ---: | ---: | ---: | +| 0 | 512 | 4.23 ms | 2.14 ms | 49.5% | +| 8k | 2560 | 17.10 ms | 10.09 ms | 41.0% | +| 16k | 4608 | 29.97 ms | 17.62 ms | 41.2% | +| 32k | 8704 | 50.51 ms | 32.94 ms | 34.8% | +| 64k | 16896 | 94.36 ms | 64.87 ms | 31.3% | +| 128k | 33280 | 174.96 ms | 128.76 ms | 26.4% | +| 256k | 66048 | 352.73 ms | 250.88 ms | 28.9% | +| 512k | 131584 | 696.79 ms | 495.90 ms | 28.8% | + +The final isolated 32k run after cleanup measured 31.81159 ms. Small matrix differences are normal laptop GPU clock variation. + +| Canonical 32k metric | Baseline | Optimized | +| --- | ---: | ---: | +| PP2048 | 209.45 tokens/s | 216.32 tokens/s | +| Total Vulkan | 9.73729 s | 9.42448 s | +| Lightning Indexer | 1.11860 s | 0.697276 s | +| Sparse FA raw | 0.198438 s | 0.195367 s | +| Sparse FA selected | 0.835110 s | 0.843385 s | +| Sparse FA reduce | 0.088463 s | 0.086678 s | +| TOP_K | 0.075473 s | 0.075350 s | + +Logs: + +- Baseline matrix: `/tmp/dsv4-lightning-prefill-baseline.log` +- Optimized matrix: `/tmp/dsv4-lightning-four-head-matrix.log` +- Final selected-pipeline 32k microbench: `/tmp/dsv4-lightning-final-selected-32k.log` +- Baseline canonical 32k llama-bench: `/tmp/dsv4-nathan-beta-32k-rerun-new-first.log` +- Optimized canonical 32k llama-bench: `/tmp/dsv4-lightning-final-32k-llama-bench.log` + +## Correctness and resources + +- Wide pipeline: all 20 focused F16 cases passed, including 127, 128, and 129-key boundaries. +- Small cooperative-matrix fallback: temporarily forced and all the same 20 cases passed. +- Wide shader on gfx1151: 168 VGPRs, 63,488 bytes LDS, no spills, eight subgroups per SIMD. +- The final pipeline-statistics run confirmed `lightning_indexer_cm_f16` was selected. +- No NaN, Inf, or comparison failures were reported. + +Correctness logs: + +- Wide: `/tmp/dsv4-lightning-consolidated-wide-correctness.log` +- Small fallback: `/tmp/dsv4-lightning-consolidated-small-correctness.log` + +## Experiments and decisions + +- Four waves improved the 32k shape about 3% and became worse at deep simulated contexts. +- Eight waves improved it about 8% before the other changes. +- Subgroup-scoped synchronization after cooperative-matrix stores improved the eight-wave version. +- Staging four query heads and weights produced the large gain by reducing redundant loads and barriers. +- Using only a subgroup barrier between head groups failed four boundary tests. One subgroup could overwrite shared query data while another still read it. A workgroup barrier is required there. +- A separate fallback shader source was avoided. Generator definitions create both variants from one file. + +## Commands + +Build only, with no GPU benchmark running: + +```sh +cd /tmp/llama-strix-beta-bench +git diff --check +cmake --build build --config Release --target test-backend-ops llama-bench -j "$(nproc)" +``` + +Focused correctness: + +```sh +cd /tmp/llama-strix-beta-bench +./build/bin/test-backend-ops test -b Vulkan0 -o LIGHTNING_INDEXER -p 'type_K=f16' > /tmp/dsv4-lightning-correctness.log 2>&1 +tail -n 30 /tmp/dsv4-lightning-correctness.log +``` + +Final 32k-equivalent microbench and pipeline selection: + +```sh +cd /tmp/llama-strix-beta-bench +GGML_VK_PIPELINE_STATS=lightning_indexer_cm_f16 ./build/bin/test-backend-ops perf -b Vulkan0 -o LIGHTNING_INDEXER -p 'kv=8704' > /tmp/dsv4-lightning-final-selected-32k.log 2>&1 +tail -n 16 /tmp/dsv4-lightning-final-selected-32k.log +``` + +Full Indexer depth matrix: + +```sh +cd /tmp/llama-strix-beta-bench +./build/bin/test-backend-ops perf -b Vulkan0 -o LIGHTNING_INDEXER -p 'nb=2048,nh=64,ns=1,nm=1,type_K=f16' > /tmp/dsv4-lightning-matrix.log 2>&1 +rg 'kv=(512|2560|4608|8704|16896|33280|66048|131584),nb=2048' /tmp/dsv4-lightning-matrix.log +``` + +Canonical 32k llama-bench. Run it only when needed, never while compiling, and inspect only the final block: + +```sh +cd /tmp/llama-strix-beta-bench +GGML_VK_PERF_LOGGER=1 ./build/bin/llama-bench -m /home/jaap/Projects/docker/localLLaMA/models/models--unsloth--DeepSeek-V4-Flash-0731-GGUF/snapshots/109848da2469efe1f1aab9e11acea08a065ccd4f/UD-IQ3_XXS/DeepSeek-V4-Flash-0731-UD-IQ3_XXS-00001-of-00004.gguf -r 1 -d 32768 -p 2048 -ub 2048 -fa 1 -n 0 > /tmp/dsv4-lightning-32k-llama-bench.log 2>&1 +last=$(grep -n 'Vulkan Timings:' /tmp/dsv4-lightning-32k-llama-bench.log | tail -n 1 | cut -d: -f1) +sed -n "${last},\$p" /tmp/dsv4-lightning-32k-llama-bench.log | tail -n 180 +``` + +Patch inspection: + +```sh +cd /tmp/llama-strix-beta-bench +git diff --check +git diff --stat +git diff +git status --short +``` + +## Next actions + +1. The user reviews and understands the four-file implementation and result summary. +2. Commit only after explicit user approval for that commit action. +3. Remove this scratch pad before a PR if it is not useful as permanent documentation. +4. Restore the separate decode microbench stash from the main worktree only if that work resumes. diff --git a/docs/development/DSV4-vulkan-sparse-prefill-progress.md b/docs/development/DSV4-vulkan-sparse-prefill-progress.md new file mode 100644 index 000000000000..e880ae63e99e --- /dev/null +++ b/docs/development/DSV4-vulkan-sparse-prefill-progress.md @@ -0,0 +1,587 @@ +# DeepSeek V4 Vulkan sparse prefill progress + +This file is a self-contained handoff for the DeepSeek V4 sparse-attention prompt-processing optimization on AMD Strix Halo. Read the repository `AGENTS.md` and `CONTRIBUTING.md` before continuing. + +## Repository state + +- Repository: `https://github.com/Nathanw1014/llama.cpp` +- Branch: `strix-halo-vulkan` +- Starting commit: `baf0025de861c6f6ea3720fa81c52ae1b2e6c078` +- Target GPU: AMD Radeon 8060S / gfx1151, RADV, Vulkan, wave64 +- The device reports `GL_KHR_cooperative_matrix`, f16 inputs with f32 accumulation, 64 KiB shared memory, and a maximum 512-thread workgroup used by this path. + +The implementation is not upstream `ggml-org/llama.cpp`. It builds on this branch's DeepSeek V4 graph, Lightning Indexer, sparse top-K hint, decode gather path, fused HC kernels, and Vulkan profiler changes. + +## Important execution constraints + +The system is an APU. CPU compilation and GPU benchmarking share power and memory bandwidth. Never build and benchmark at the same time. Serialize all builds, correctness tests, and performance tests. + +GPU commands must run with host GPU access. In an agent sandbox, request elevated/out-of-sandbox execution. A sandboxed benchmark showed only CPU activity and is invalid. + +Redirect the full model benchmark to a log. Inspect only the last profiler block with `tail`; do not load the full log into agent context. + +## Build and ccache + +The build directory is `build`, configured as Release with Vulkan enabled. The default ccache directory was read-only in the agent environment, so use a writable directory: + +```bash +cmake -S . -B build \ + -DGGML_VULKAN=ON \ + -DCMAKE_BUILD_TYPE=Release \ + -DCMAKE_C_COMPILER_LAUNCHER=ccache \ + -DCMAKE_CXX_COMPILER_LAUNCHER=ccache + +CCACHE_DIR=/tmp/llama-cpp-ccache cmake --build build --config Release \ + --target llama-bench test-backend-ops -j "$(nproc)" + +CCACHE_DIR=/tmp/llama-cpp-ccache ccache -s +``` + +ccache was verified active. The final build reported direct hits. Note that an incremental change to `ggml-vulkan.cpp` is one large C++ translation unit and therefore uses one compiler core even with `-j`. Shader object regeneration can run in parallel. + +## Canonical benchmark command + +Do not change or omit switches for the 32k acceptance run: + +```bash +GGML_VK_PERF_LOGGER=1 ./build/bin/llama-bench \ + -m ~/Projects/docker/localLLaMA/models/models--unsloth--DeepSeek-V4-Flash-0731-GGUF/snapshots/109848da2469efe1f1aab9e11acea08a065ccd4f/UD-IQ3_XXS/DeepSeek-V4-Flash-0731-UD-IQ3_XXS-00001-of-00004.gguf \ + -r 1 -d 32768 -p 2048 -ub 2048 -fa 1 -n 0 \ + > /tmp/dsv4-vulkan-32k.log 2>&1 + +tail -n 180 /tmp/dsv4-vulkan-32k.log +``` + +The known model path exists on the target system. + +## Graph and dispatch findings + +DeepSeek V4 builds the Lightning Indexer and sparse attention in `src/models/deepseek4.cpp`: + +1. `build_lid_top_k()` creates indexer Q/K/weights and calls `ggml_lightning_indexer()`. +2. `ggml_top_k()` selects up to `hparams.indexer_top_k` compressed-cache indices for every query token. +3. `build_csa_lid_attention()` concatenates the raw SWA K prefix with compressed CSA K, builds a dense mask carrying the same sparse selection, and calls `build_attn_mha(..., top_k, raw_k->ne[2])`. +4. `build_attn_mha()` attaches `top_k` and `n_kv_raw` to `GGML_OP_FLASH_ATTN_EXT` through `ggml_flash_attn_ext_add_top_k()`. + +At the final 32k PP2048 batch: + +- total K rows: 11,008 +- `n_kv_raw`: 2,304 raw SWA rows, always attended subject to the mask +- `n_top_k`: 512 selected compressed rows per query token +- active rows: 2,816 +- selectable compressed region: 8,704 + +The custom Vulkan path is selected by `ggml_vk_flash_attn_top_k()` before ordinary FA. Its gate requires the DeepSeek V4 shape and `total_k >= 3 * (n_kv_raw + n_top_k)`. The final shape satisfies `11008 >= 3 * 2816`. + +The old `flash_attn_top_k.comp` shader is scalar/subgroup code. One 512-thread workgroup covers eight heads for one query token. It stages 16 selected 512-wide K/V rows, computes QK with scalar FMAs and `subgroupAdd`, updates online softmax one key at a time, and accumulates PV manually. It does not use cooperative matrices. + +The top-K set differs by query token but is shared by all 64 query heads for that token. This makes the attention for one token a regular matrix problem across heads and selected keys despite sparse per-token indexing. + +## Root cause evidence + +The existing Vulkan timestamp infrastructure was extended with `ggml_vk_perf_mark_subop()` after the sparse dispatch. This reports the sparse kernel separately as `FA_TOP_K_SPARSE (sub-op)` or `FA_TOP_K_CM (sub-op)`. + +Focused test shape: + +```bash +GGML_VK_PERF_LOGGER=1 ./build/bin/test-backend-ops perf \ + -b Vulkan0 -o FLASH_ATTN_EXT \ + -p 'kv=32768,nb=512,n_kv_raw=1024,n_top_k=512,sinks=0' +``` + +Results: + +- old scalar sparse kernel: 61.42 ms, 1.68 TFLOPS of useful active-set work +- ordinary dense FA diagnostic (`GGML_VK_FA_TOPK=0`): 255.97 ms, about 8.7 TFLOPS over the full dense work +- final cooperative sparse kernel: 32.55 ms, 3.17 TFLOPS of useful active-set work + +The residual `FLASH_ATTN_EXT` interval after the sparse timestamp is only about 4-7 us. The cost is inside the shader, not dispatch or surrounding synchronization. + +A temporary uniform stage-profiling mode was used and removed. For the final 32-head tile: + +- selected K gather + cooperative QK: 14.30 ms +- gather + QK + serial softmax: 26.34 ms +- gather + QK + parallel softmax: 15.53 ms +- full kernel: 32.55 ms +- the remaining cooperative PV/output portion is about 17.0 ms + +The old scalar shader was compute/issue inefficient. Dense FA proved matrix hardware is much faster but was still too expensive because it processes all K rows. The final implementation preserves sparsity and uses the matrix hardware for both QK and PV. + +## Implementation + +Files changed: + +- `ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp` + - New cooperative-matrix sparse prefill shader. + - One 512-thread workgroup covers 32 query heads for one token. + - Eight wave64 subgroups cover two 16-head tiles by four 16-key or 16-output-dimension tiles. + - Processes 64 selected keys per online-softmax block. + - Stages only indexed selected K/V tiles, never the full K range. + - Uses f16 cooperative-matrix inputs and f32 accumulation for QK and PV. + - Uses a 16-lane segmented softmax per head. XOR subgroup shuffles reduce max and sum for four independent heads per wave without workgroup barriers. + - Keeps f32 output accumulators and normalizes after all active blocks. + - Preserves the raw prefix, top-K index validation, mask, sinks, stream strides, and K == V latent behavior. +- `ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp` + - Embeds the new shader when cooperative-matrix shader support is available. +- `ggml/src/ggml-vulkan/ggml-vulkan.cpp` + - Adds the cooperative sparse pipeline when the device supports the required 16x16x16 f16/f32 cooperative matrix shape. + - Selects it by capability and keeps the scalar shader as fallback. + - Adds `GGML_VK_FA_TOPK=0` to force ordinary dense FA for diagnostics. + - Adds `GGML_VK_FA_TOPK_CM=0` to force the old scalar sparse shader for A/B tests. + - Adds sparse sub-operation timestamps through the existing profiler. + +The path is capability-based, not hardcoded to Strix Halo. The current sparse shape gate remains DeepSeek V4-specific. Devices without the required cooperative matrix support keep the correct scalar sparse or dense fallback. + +Two discarded prototypes are useful context: + +- A 16-head, 64-key streaming cooperative tile was correct and reduced the focused test from 61.4 to 44.6 ms. +- Keeping 32 complete 512-wide K/V rows in LDS grew shared memory to about 41 KiB, reduced residency, doubled block/barrier count, and regressed to 73.8 ms. Do not retry full-row LDS staging without solving occupancy. +- A 32-head, 64-key streaming tile halved irregular row loads but initially stayed near 44.7 ms because serial softmax cost about 11.5 ms. Parallel segmented softmax produced the final 32.55 ms result. + +## Correctness validation + +Run: + +```bash +./build/bin/test-backend-ops test -b Vulkan0 -o FLASH_ATTN_EXT -p 'n_top_k=' +``` + +Final result: 8/8 sparse top-K FA cases passed against the CPU reference. Cases include: + +- decode and short batches that use dense/gather fallback +- sparse prefill batch sizes 64 and 128 +- `n_kv_raw` plus top-K selection +- invalid top-K index handling from the test fixture +- sinks enabled and disabled +- sparse threshold transitions +- an active-key count of 193, which exercises a partial final 64-key block + +The test uses the existing FA tolerance of NMSE <= `5e-4`. No NaN or Inf failure occurred. The implementation changes Q and probability inputs to f16 cooperative-matrix operands with f32 accumulation, matching the precision strategy of ordinary Vulkan cooperative FA. + +Still desirable before broader submission: + +- compare model logits on controlled prompts between `GGML_VK_FA_TOPK_CM=0` and the default cooperative path + +## Canonical 32k results + +Exact clean runs, same command and machine, no concurrent build: + +```text +32k context, PP 2048, ub 2048 + +Before (commit baf0025de): +112.29 tok/s +Total Vulkan: 18.1957 s +Sparse FA: 8.84496 s, 421.189 ms/layer +Lightning Indexer: 1.19438 s +TOP_K: 0.076948 s + +After: +152.32 tok/s +Total Vulkan: 13.4032 s +Sparse FA: 4.44701 s, 211.762 ms/layer +Lightning Indexer: 1.13156 s +TOP_K: 0.073945 s + +Change: +Throughput: +35.65% +Total Vulkan time: -26.34% +Sparse FA time: -49.72% +Sparse FA saved: 4.398 s +Total GPU time saved: 4.793 s +``` + +The profiler now lists the optimized dispatch as `FA_TOP_K_CM (sub-op)`. The following residual `FLASH_ATTN_EXT` line is only the post-mark interval and must not be interpreted as the kernel time. + +## Context-depth measurements + +All points use PP2048, ub2048, FA enabled, one repetition, and no token generation. They were run sequentially with no compiler active: + +```text +Existing depth tok/s Total Vulkan Final large FA Lightning Indexer TOP_K +0 253.44 8.041 s 0.294 s 0.095 s 0.001 s +8192 211.01 9.666 s 1.625 s 0.379 s 0.022 s +16384 177.19 11.518 s 3.031 s 0.678 s 0.044 s +32768 152.32 13.403 s 4.447 s 1.132 s 0.074 s +``` + +At 0, 8k, and 16k, total K is below the existing sparse-path gate `total_k >= 3 * (n_kv_raw + n_top_k)`. These points use the unchanged ordinary dense FA implementation, so the cooperative sparse change does not affect or regress them. At 32k, total K is 11,008 and the cooperative sparse path engages. The 32k `Final large FA` value is the `FA_TOP_K_CM (sub-op)` total; the lower-depth values are the large ordinary `FLASH_ATTN_EXT` totals. + +Logs: + +- `/tmp/dsv4-vulkan-cm-0k.log` +- `/tmp/dsv4-vulkan-cm-8k.log` +- `/tmp/dsv4-vulkan-cm-16k.log` +- `/tmp/dsv4-vulkan-cm-32k.log` +- `/tmp/dsv4-vulkan-baseline-clean.log` +- `/tmp/dsv4-fa-cm-correctness-final.log` + +## Next optimization target + +The cooperative sparse FA remains the largest context-dependent cost at about 4.45 s total. Stage profiling indicates approximately 14.3 ms of focused-test time in gather/QK, about 1.2 ms in parallel softmax, and about 17 ms in PV/output. + +The next useful work is PV and output accumulation, not TOP_K. Investigate: + +- reducing repeated selected V staging across the two 32-head workgroups per token without increasing LDS enough to lose occupancy +- reducing the eight output-dimension passes or retaining more PV state in cooperative fragments/registers +- checking register count and spills for the 32 f32 output accumulators per invocation using RADV shader statistics +- alternate 32-head layouts that keep the same eight-wave occupancy but improve PV scheduling +- query-tile overlap/union gathering only if measured top-K overlap is high enough; a whole-2048-query union is unlikely to help + +Do not optimize TOP_K first. At 32k it is only about 74 ms total. Lightning Indexer is about 1.13 s and is the next context-dependent target only after sparse FA improves further. + +## Useful diagnostics + +Force old scalar sparse path: + +```bash +GGML_VK_FA_TOPK_CM=0 GGML_VK_PERF_LOGGER=1 ./build/bin/test-backend-ops perf \ + -b Vulkan0 -o FLASH_ATTN_EXT \ + -p 'kv=32768,nb=512,n_kv_raw=1024,n_top_k=512,sinks=0' +``` + +Force ordinary dense FA: + +```bash +GGML_VK_FA_TOPK=0 GGML_VK_PERF_LOGGER=1 ./build/bin/test-backend-ops perf \ + -b Vulkan0 -o FLASH_ATTN_EXT \ + -p 'kv=32768,nb=512,n_kv_raw=1024,n_top_k=512,sinks=0' +``` + +Default cooperative sparse path: + +```bash +GGML_VK_PERF_LOGGER=1 ./build/bin/test-backend-ops perf \ + -b Vulkan0 -o FLASH_ATTN_EXT \ + -p 'kv=32768,nb=512,n_kv_raw=1024,n_top_k=512,sinks=0' +``` + +Always run these sequentially. Do not run a compiler concurrently on this APU. + +## Experimental raw-prefix split prototype + +An uncommitted follow-up prototype was tested after commit `24c7ead76cc1a9631fa1b42b0bfa53a15169e1ba`. Check `git status` before continuing. It was initially tested with `GGML_VK_FA_TOPK_SPLIT=1`. After the successful 32k run, the split path was changed to default-on for a llama-server coherence test. Set `GGML_VK_FA_TOPK_SPLIT=0` to restore the single cooperative sparse kernel. + +Stage profiling on the previously used 19,200-row sparse shape (`kv=19200,nb=2048,n_kv_raw=2304,n_top_k=512,sinks=0`) showed: + +```text +cooperative QK and selected-K gather: 88.885 ms +softmax increment: 4.909 ms +cooperative PV and output increment: 128.399 ms +full cooperative sparse kernel: 222.193 ms +``` + +PV and output were 57.8% of the kernel. More importantly, most active keys are not sparse: all 2,304 raw-prefix rows are contiguous, while only 512 compressed rows use top-K indices. The prototype therefore makes two attention partitions: + +1. Ordinary optimized Vulkan cooperative FA processes the contiguous raw prefix. +2. The cooperative top-K shader processes only the 512 selected compressed rows. +3. The existing split-K reduction combines both online-softmax partitions and applies sinks. + +This preserves the exact raw-prefix, sparse top-K, causal mask, and softmax semantics. It reuses the existing ordinary FA and split-K reduction rather than adding a new subsystem. + +The exact-shape microbenchmark improved from 222.19 ms to 111.33 ms. Its steady split stages were about 62-64 ms raw-prefix FA, 43-46 ms selected sparse FA, and 3.6-3.9 ms reduction. + +Correctness validation: + +```bash +GGML_VK_FA_TOPK_SPLIT=1 ./build/bin/test-backend-ops test \ + -b Vulkan0 -o FLASH_ATTN_EXT -p 'n_top_k=' + +./build/bin/test-backend-ops test -b Vulkan0 -o FLASH_ATTN_EXT +``` + +Results were 8/8 sparse top-K cases and 13,296/13,296 complete Vulkan FA cases against the CPU reference. No NaN or Inf failure occurred. + +Canonical 32k prototype command: + +```bash +GGML_VK_FA_TOPK_SPLIT=1 GGML_VK_PERF_LOGGER=1 ./build/bin/llama-bench \ + -m ~/Projects/docker/localLLaMA/models/models--unsloth--DeepSeek-V4-Flash-0731-GGUF/snapshots/109848da2469efe1f1aab9e11acea08a065ccd4f/UD-IQ3_XXS/DeepSeek-V4-Flash-0731-UD-IQ3_XXS-00001-of-00004.gguf \ + -r 1 -d 32768 -p 2048 -ub 2048 -fa 1 -n 0 \ + > /tmp/dsv4-vulkan-split-32k.log 2>&1 +``` + +Final 32k result: + +```text +32k context, PP 2048, ub 2048 + +Old scalar: 112.29 tok/s, 18.1957 s total, 8.84496 s sparse FA +Committed coopmat: 152.32 tok/s, 13.4032 s total, 4.44701 s sparse FA +Experimental split: 208.70 tok/s, 9.7704 s total, 1.20170 s split sparse FA + +Experimental split stages: +raw-prefix FA: 0.210775 s total, 10.037 ms/layer +selected sparse FA: 0.908190 s total, 43.247 ms/layer +split reduction: 0.082732 s total, 3.940 ms/layer +Lightning Indexer: 1.110410 s total, 52.877 ms/layer +TOP_K: 0.077394 s total, 3.685 ms/layer +``` + +Relative to the committed cooperative path, throughput improved 37.0%, total Vulkan time fell 27.1%, and sparse FA time fell 73.0%. Relative to the old scalar path, throughput improved 85.9%, total Vulkan time fell 46.3%, and sparse FA time fell 86.4%. + +The initial split implementation used 538,968,064 bytes (514 MiB) of scratch at PP2048 because each of two partitions stored an f32 partial output for 512 dimensions x 64 heads x 2048 queries, plus L/M data. This was subsequently reduced by query tiling as described below. + +The old sparse selection heuristic required `total_k >= 3 * active_k`. For the coherence test it now selects sparse attention whenever `total_k > active_k`; equality still uses dense FA because no keys are pruned. The shape, capability, and allocation gates remain unchanged. + +Prototype logs: + +- `/tmp/dsv4-fa-exact-qk.log` +- `/tmp/dsv4-fa-exact-softmax.log` +- `/tmp/dsv4-fa-exact-full.log` +- `/tmp/dsv4-fa-exact-split.log` +- `/tmp/dsv4-fa-split-correctness.log` +- `/tmp/dsv4-fa-all-correctness.log` +- `/tmp/dsv4-vulkan-split-32k.log` + +## Llama-server coherence check + +After making the split path default-on and changing the sparse crossover to `total_k > active_k`, `llama-server` produced a coherent response from a 2,044-token prompt at significant context depth. The final profiler block confirmed that all 21 sparse-attention layers used the split path: + +```text +FA_TOP_K_RAW: 0.179223 s total, 8.534 ms/layer +FA_TOP_K_SELECTED: 0.886629 s total, 42.220 ms/layer +FA_TOP_K_REDUCE: 0.083240 s total, 3.964 ms/layer +Split sparse FA: 1.149092 s total, 54.719 ms/layer +Lightning Indexer: 1.054610 s total, 50.220 ms/layer +TOP_K: 0.044540 s total, 2.121 ms/layer +Total Vulkan: 9.797560 s +``` + +This closely matches the canonical PP2048 llama-bench result of 9.770 s total and 1.202 s split sparse FA. The focused sparse CPU-reference test was rerun after the crossover change and passed 8/8 cases. + +## Tiled split scratch optimization + +Commit `4bbe53e4775f0707de8158e4977a16ab770829da` used two full PP2048 output partitions. The same two-partition algorithm now processes at most 256 query tokens per tile and reuses the split scratch between tiles. Q, mask, top-K, and destination descriptors are offset to the tile while K/V remain shared. A Vulkan pipeline barrier separates reuse of each scratch tile. + +Scratch at PP2048 changed from: + +```text +Before: 538,968,064 bytes (514 MiB) +After: 67,371,008 bytes (64.25 MiB) +Change: 8x reduction +``` + +The previously used 19,200-row microbenchmark (`kv=19200,nb=2048,n_kv_raw=2304,n_top_k=512,sinks=0`) measured: + +```text +Full-batch two partitions: 114.65 ms +256-query tiled path: 113.70 ms +``` + +Two one-partition alternatives were tested and discarded. Merging the raw partial directly inside the cooperative selected shader measured 119.17 ms. Writing selected output separately and using a lightweight merge kernel measured 120.45 ms. Both cut scratch in half but regressed because writing selected output outside the contiguous split layout increased the selected stage from about 44 ms to about 51 ms. Query tiling preserves the faster memory layout. + +Canonical 32k result after tiling: + +```text +32k context, PP 2048, ub 2048 + +Full-batch split: 208.70 tok/s, 9.77039 s total, 1.20170 s sparse FA +Tiled split: 209.62 tok/s, 9.72897 s total, 1.18721 s sparse FA + +Tiled split stages: +raw-prefix FA: 0.192879 s total, 168 tile dispatches +selected sparse FA: 0.910153 s total, 168 tile dispatches +split reduction: 0.084179 s total, 168 tile dispatches +Lightning Indexer: 1.100230 s total +TOP_K: 0.075158 s total +``` + +The 168 dispatch count is eight tiles x 21 sparse-attention layers. Relative to the full-batch split path, throughput improved 0.44%, total Vulkan time fell 0.42%, and sparse FA time fell 1.21%. The main result is the 8x scratch reduction without a performance regression. + +Final correctness after removing the discarded merge prototypes: + +- focused sparse top-K suite: 9/9 passed against the CPU reference, including a 257-query tile-boundary case +- complete Vulkan Flash Attention suite: 13,296/13,296 passed +- no NaN or Inf failure + +Logs: + +- `/tmp/dsv4-fa-exact-fused.log` +- `/tmp/dsv4-fa-exact-merge.log` +- `/tmp/dsv4-fa-exact-two-part-current.log` +- `/tmp/dsv4-fa-exact-tiled.log` +- `/tmp/dsv4-fa-tiled-final-correctness.log` +- `/tmp/dsv4-fa-tiled-boundary-correctness.log` +- `/tmp/dsv4-fa-tiled-all-correctness.log` +- `/tmp/dsv4-vulkan-tiled-32k.log` + +## Selected PV probability reuse + +The selected cooperative-matrix stage remained the largest split sparse-FA component. Profiling the tiled 19,200-row shape showed: + +```text +selected K gather and cooperative QK: about 20.0 ms +QK plus softmax: about 22.3 ms +full selected stage: about 45.5 ms +cooperative PV and output increment: about 23.2 ms +``` + +PV and output were about 51% of the selected stage. The shader previously loaded each of four 16x16 probability cooperative-matrix fragments again for every one of the eight 64-dimension PV passes. The optimized shader loads these four fragments once per 64-key block and retains them across all PV dimension passes. This follows the probability-fragment lifetime used by ordinary cooperative Vulkan FA. + +The shader also aliases the shared score and PV-output matrices because their lifetimes do not overlap. Shader resource statistics changed as follows: + +```text + Before After +VGPRs 192 192 +VGPR spills 0 0 +LDS 36,864 28,672 bytes +static instructions 11,524 11,358 +``` + +The 19,200-row microbenchmark changed from 113.70 ms for the committed tiled path to 110.11 ms. The selected stage fell from about 45.5 ms to about 43.7 ms. The raw-prefix and reduction implementations are unchanged. + +Two canonical 32k runs after this change measured: + +```text +32k context, PP 2048, ub 2048 + +Committed tiled reference: +209.62 tok/s +Total Vulkan: 9.72897 s +Split sparse FA: 1.18721 s + raw-prefix FA: 0.192879 s + selected sparse FA: 0.910153 s + split reduction: 0.084179 s +Lightning Indexer: 1.100230 s +TOP_K: 0.075158 s + +Probability reuse, run 1: +211.03 tok/s +Total Vulkan: 9.66363 s +Split sparse FA: 1.12776 s + raw-prefix FA: 0.192078 s + selected sparse FA: 0.849862 s + split reduction: 0.085821 s +Lightning Indexer: 1.098980 s +TOP_K: 0.071036 s + +Probability reuse, run 2: +209.57 tok/s +Total Vulkan: 9.72968 s +Split sparse FA: 1.13553 s + raw-prefix FA: 0.192987 s + selected sparse FA: 0.859890 s + split reduction: 0.082653 s +Lightning Indexer: 1.102550 s +TOP_K: 0.071474 s +``` + +The two-run selected-stage improvement is 5.5-6.6%, and the complete split sparse-FA improvement is 4.4-5.0%. The two-run throughput mean is 210.30 tok/s, 0.32% above the 209.62 tok/s reference. End-to-end noise in unrelated model kernels is larger than this small total-throughput change, but both profiler runs isolate a consistent gain in the modified selected stage. + +Correctness after the change: + +- focused sparse top-K suite: 9/9 passed against the CPU reference +- complete Vulkan Flash Attention suite: 13,297/13,297 passed +- pipeline statistics probe: 1/1 passed, with no SGPR or VGPR spills +- no NaN or Inf failure + +Logs: + +- `/tmp/dsv4-fa-selected-qk-tiled.log` +- `/tmp/dsv4-fa-selected-softmax-tiled.log` +- `/tmp/dsv4-fa-lds-alias-stats.log` +- `/tmp/dsv4-fa-pmat-stats.log` +- `/tmp/dsv4-fa-exact-lds-alias.log` +- `/tmp/dsv4-fa-exact-pmat.log` +- `/tmp/dsv4-fa-pmat-correctness.log` +- `/tmp/dsv4-fa-pmat-all-correctness.log` +- `/tmp/dsv4-vulkan-32k-pmat.log` +- `/tmp/dsv4-vulkan-32k-pmat-repeat.log` + +The next sparse-FA optimization should continue to target the selected PV/output half. The retained probability fragments remove redundant cooperative loads without increasing reported VGPR allocation. More invasive changes such as doubling the PV dimension tile can reduce barriers but must be designed around the eight available wave64 subgroups, 64 KiB LDS limit, and already high 192-VGPR allocation. Do not build a crossover matrix until the next kernel layout is settled because an optimization can shift the crossover. + +## Selected mask caching and deep-context micro matrix + +The sparse FA `kv` dimension is compressed K/V rows, not source-token context depth. For the tested DeepSeek V4 graph, a PP2048 batch uses 2,304 raw rows and approximately one compressed row per four source tokens. The useful synthetic mapping is: + +```text +Source context Sparse FA kv rows +32k 11,008 +64k 19,200 +128k 35,584 +256k 68,352 +512k 133,888 +``` + +The performance test registry now contains PP2048 cases at all five K extents with `n_kv_raw=2304` and `n_top_k=512`. Use this command template and replace `KV` with a value from the table: + +```bash +GGML_VK_PERF_LOGGER=1 ./build/bin/test-backend-ops perf \ + -b Vulkan0 -o FLASH_ATTN_EXT \ + -p 'kv=KV,nb=2048,n_kv_raw=2304,n_top_k=512,sinks=0' +``` + +These are synthetic sparse-FA depth tests. They do not include the Lightning Indexer or the rest of the model graph and do not replace the canonical 32k llama-bench. + +The selected shader previously loaded the same selected-key mask value independently for all 32 heads during both softmax passes. The shader now loads each mask value once per 64-key block and reuses it from LDS. Q and probability storage also share one LDS allocation, while K and V staging share another because both pairs have disjoint lifetimes. + +Shader resources changed from the probability-reuse commit: + +```text + Before After +VGPRs 192 192 +VGPR spills 0 0 +LDS 28,672 24,576 bytes +static instructions 11,358 11,233 +``` + +A 128-dimension V staging experiment was correct but rejected. With operand aliasing it used exactly 32 KiB LDS. It was 2.3% slower at simulated 32k, approximately equal at 64k, and 1.6% faster at 128k. The retained 64-dimension layout is better for the canonical 32k target and is simpler. + +Selected-mask caching improved every synthetic depth relative to the same 64-dimension operand-alias layout: + +```text +Depth Selected before Selected after Split total before Split total after +32k 43.16 ms 41.23 ms 109.45 ms 108.07 ms +64k 43.22 ms 41.42 ms 110.45 ms 109.11 ms +128k 52.26 ms 48.74 ms 119.15 ms 116.92 ms +256k 51.64 ms 49.02 ms 118.06 ms 116.93 ms +512k 51.33 ms 49.20 ms 117.87 ms 116.27 ms +``` + +The 256k case exposed a separate path-selection cutoff. The raw-prefix ordinary FA dispatch encoded its mask row stride in 16 bits and disabled split sparse FA above 65,535 K rows. It now uses a flagged convention that carries the full 32-bit mask stride in the existing `split_kv` push constant for this one-partition partial-output dispatch. No push-constant structure was enlarged. At simulated 256k, this changes the selected path from unsplit cooperative sparse FA at 206.22 ms to split sparse FA at 118.06 ms before mask caching, a 42.8% reduction. The 512k case also uses the split path successfully. + +Canonical 32k PP2048 after operand aliasing and selected-mask caching: + +```text +Previous commit, two runs: +209.57-211.03 tok/s +Total Vulkan: 9.664-9.730 s +Split sparse FA: 1.128-1.136 s + raw-prefix FA: 0.192-0.193 s + selected sparse FA: 0.850-0.860 s + split reduction: 0.083-0.086 s + +Current result: +210.47 tok/s +Total Vulkan: 9.68921 s +Split sparse FA: 1.10160 s + raw-prefix FA: 0.192860 s + selected sparse FA: 0.821381 s + split reduction: 0.087354 s +Lightning Indexer: 1.102230 s +TOP_K: 0.072488 s +``` + +The selected stage improves 3.4-4.5% and complete split sparse FA improves 2.3-3.0% relative to the previous commit's two canonical runs. End-to-end throughput remains within model-wide benchmark noise. + +Correctness: + +- focused sparse top-K suite: 9/9 passed +- complete Vulkan Flash Attention suite: 13,297/13,297 passed +- no NaN or Inf failure +- simulated 256k and 512k performance cases both selected the split path and completed successfully + +Benchmark policy for this laptop APU: do not repeat llama-bench when the exact-shape microbenchmarks and the first canonical profiler block agree. Sustained load can lower GPU clocks and bias a repeat. Repeat only when the first result is anomalous or contradicts the microbenchmarks. Never build and benchmark concurrently. + +Logs: + +- `/tmp/dsv4-fa-mask-cache-32k.log` +- `/tmp/dsv4-fa-mask-cache-64k.log` +- `/tmp/dsv4-fa-mask-cache-128k.log` +- `/tmp/dsv4-fa-mask-cache-256k.log` +- `/tmp/dsv4-fa-mask-cache-512k.log` +- `/tmp/dsv4-fa-final-micro-32k.log` +- `/tmp/dsv4-fa-final-focused-correctness.log` +- `/tmp/dsv4-fa-mask-cache-all-correctness.log` +- `/tmp/dsv4-vulkan-32k-mask-cache.log` diff --git a/ggml/include/ggml.h b/ggml/include/ggml.h index f0eda79a1ae2..fb83a874e877 100644 --- a/ggml/include/ggml.h +++ b/ggml/include/ggml.h @@ -2486,6 +2486,16 @@ extern "C" { struct ggml_tensor * a, struct ggml_tensor * sinks); + // sparse attention hint: attend only to the first n_kv_raw keys (dense prefix) plus the + // keys selected by top_k. top_k is I32 [n_top_k, n_tokens, 1, n_streams]; each index i + // selects absolute key n_kv_raw + i. Negative or out-of-range indices are ignored. + // Backends may ignore the hint: the kq_mask must still encode the same selection, so a + // dense fallback computes the identical result. + GGML_API void ggml_flash_attn_ext_add_top_k( + struct ggml_tensor * a, + struct ggml_tensor * top_k, + int64_t n_kv_raw); + // TODO: needs to be adapted to ggml_flash_attn_ext GGML_API struct ggml_tensor * ggml_flash_attn_back( struct ggml_context * ctx, diff --git a/ggml/src/ggml-backend.cpp b/ggml/src/ggml-backend.cpp index dc6b98946e20..0955a431b171 100644 --- a/ggml/src/ggml-backend.cpp +++ b/ggml/src/ggml-backend.cpp @@ -1446,7 +1446,10 @@ void ggml_backend_sched_split_graph(ggml_backend_sched_t sched, struct ggml_cgra } // check if the split has too many inputs // FIXME: count the number of inputs instead of only checking when full - if (split->n_inputs >= split->inputs_capacity) { + // cut on the constant, not on inputs_capacity: capacity doubles on demand and + // is never reset, so using it lets the cut point drift up and keeps every input + // copy of an ever-longer split live at once + if (split->n_inputs >= GGML_SCHED_MAX_SPLIT_INPUTS) { const size_t id = hash_id(src); int src_backend_id = sched->hv_tensor_backend_ids[id]; bool supported = ggml_backend_sched_buffer_supported(sched, src, cur_backend_id); diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 8147894b0e47..353edb87c466 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -64,6 +64,7 @@ typedef struct VkPhysicalDeviceCooperativeMatrixDecodeVectorFeaturesNV { #include #include #include +#include #include #include #include @@ -798,6 +799,22 @@ static bool ggml_vk_lightning_indexer_k_type_supported(ggml_type type) { return std::find(lightning_indexer_k_types.begin(), lightning_indexer_k_types.end(), type) != lightning_indexer_k_types.end(); } +// Indexer head counts we build kernels for: 4 = Qwen4-exp (qwen4_exp), 64 = DeepSeek-V4 and +// GLM-DSA, 32 kept because test-backend-ops exercises it. N_HEAD is a specialization constant, +// so each entry is a separately compiled kernel with the head loop fully unrolled - as a push +// constant the loop cannot unroll and decode costs ~55% more (measured, gfx1151, kv=131584). +static constexpr uint32_t LI_NH_VALUES[] = { 4, 32, 64 }; +#define LI_NH_COUNT (sizeof(LI_NH_VALUES) / sizeof(LI_NH_VALUES[0])) + +static int ggml_vk_li_nh_index(int64_t nh) { + for (size_t i = 0; i < LI_NH_COUNT; ++i) { + if ((int64_t) LI_NH_VALUES[i] == nh) { + return (int) i; + } + } + return -1; +} + struct vk_device_struct { std::recursive_mutex mutex; mutable std::shared_mutex pinned_memory_mutex; @@ -852,6 +869,7 @@ struct vk_device_struct { bool add_rms_fusion; uint32_t partials_binding_alignment; uint32_t max_nodes_per_submit; + uint64_t max_bytes_per_submit; bool shader_64b_indexing; @@ -938,6 +956,10 @@ struct vk_device_struct { vk_matmul_pipeline2 pipeline_dequant_mul_mat_mat_id[GGML_TYPE_COUNT]; vk_matmul_pipeline2 pipeline_dequant_mul_mat_mat_id_q8_1[GGML_TYPE_COUNT]; + // EXPERIMENT (GGML_VK_MMID_F16B=1): f16-B mul_mat_id pipelines on KHR_coopmat + // devices (upstream only builds f32-B there). Populated only when the env flag + // is set; empty otherwise. + vk_matmul_pipeline2 pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_COUNT]; vk_pipeline pipeline_matmul_split_k_reduce; vk_pipeline pipeline_quantize_q8_1_x4; @@ -978,6 +1000,7 @@ struct vk_device_struct { vk_pipeline pipeline_add_id_f32; vk_pipeline pipeline_concat_i8, pipeline_concat_i16, pipeline_concat_i32, pipeline_concat_i64; + vk_pipeline pipeline_concat_transpose_i32; vk_pipeline pipeline_upscale_nearest_f32, pipeline_upscale_bilinear_f32, pipeline_upscale_bicubic_f32, pipeline_upscale_bilinear_antialias_f32; vk_pipeline pipeline_scale_f32; vk_pipeline pipeline_log[2]; @@ -1104,6 +1127,21 @@ struct vk_device_struct { vk_pipeline pipeline_lightning_indexer_f32[GGML_TYPE_COUNT]; // [size_idx][kda] where size_idx: 0=d16, 1=d32, 2=d64, 3=d128 vk_pipeline pipeline_gated_delta_net[4][2]; + // One pipeline per supported indexer head count; see LI_NH_VALUES. + vk_pipeline pipeline_lightning_indexer_f16[LI_NH_COUNT]; + vk_pipeline pipeline_lightning_indexer_cm_f16[LI_NH_COUNT]; + vk_pipeline pipeline_lightning_indexer_cm_small_f16[LI_NH_COUNT]; + vk_pipeline pipeline_lightning_indexer_decode_cm_f16[LI_NH_COUNT]; + vk_pipeline pipeline_flash_attn_top_k_f16; + vk_pipeline pipeline_flash_attn_top_k_cm_f16; + vk_pipeline pipeline_flash_attn_gather_f16; + vk_pipeline pipeline_flash_attn_gather_dq[GGML_TYPE_COUNT]; + vk_pipeline pipeline_flash_attn_union_f16; + vk_pipeline pipeline_flash_attn_gather_union_f16; + vk_pipeline pipeline_flash_attn_gather_union_dq[GGML_TYPE_COUNT]; + vk_pipeline pipeline_dsv4_hc_pre_f32; + vk_pipeline pipeline_dsv4_hc_comb_f32; + vk_pipeline pipeline_dsv4_hc_post_f32; vk_pipeline pipeline_ssm_scan_f32_d128; vk_pipeline pipeline_ssm_scan_f32_d256; vk_pipeline pipeline_ssm_conv_f32; @@ -1367,6 +1405,7 @@ struct vk_mat_mat_id_push_constants { uint32_t nei0; uint32_t nei1; uint32_t nbi1; uint32_t ne11; uint32_t n_experts; uint32_t hoist_row_ids; + uint32_t fusion_flags; }; struct vk_mat_vec_id_push_constants { uint32_t ncols; @@ -1933,6 +1972,73 @@ struct vk_op_gated_delta_net_push_constants { uint32_t K; }; +// push constants for the fork's wave64 f16 lightning-indexer kernels (scalar-64 + CM family) +struct vk_op_lightning_indexer_cm_push_constants { + uint32_t n_kv, n_batch, n_stream, nem3; + uint32_t nb1, nb3; + uint32_t nbq1, nbq2, nbq3; + uint32_t nbk2, nbk3; + uint32_t nbw1, nbw3; + uint32_t nbm1, nbm3; +}; +static_assert(sizeof(vk_op_lightning_indexer_cm_push_constants) <= 128); + +struct vk_op_dsv4_hc_pre_push_constants { + uint32_t n_embd, hc, nr; + uint32_t sx0, sx1, sx2; + uint32_t sw0, sw1; + uint32_t sd0, sd1; +}; +static_assert(sizeof(vk_op_dsv4_hc_pre_push_constants) <= 128); + +struct vk_op_dsv4_hc_comb_push_constants { + uint32_t n_tokens; + uint32_t sm0, sm1; + uint32_t ss0; + uint32_t sb0; + uint32_t sd0, sd1, sd2; + float eps; + int32_t n_iter; +}; +static_assert(sizeof(vk_op_dsv4_hc_comb_push_constants) <= 128); + +struct vk_op_dsv4_hc_post_push_constants { + uint32_t n_embd, hc, nr; + uint32_t sx0, sx1; + uint32_t sr0, sr1, sr2; + uint32_t sp0, sp1; + uint32_t sc0, sc1, sc2; + uint32_t sd0, sd1, sd2; +}; +static_assert(sizeof(vk_op_dsv4_hc_post_push_constants) <= 128); + +struct vk_op_flash_attn_union_push_constants { + uint32_t n_kv, n_kv_raw, n_batch, n_top_k, max_union, nbt1, max_words, pad_to, count_only; +}; +// nbk1/nbk3 are in 4-byte WORDS, not elements: the gather relocates K rows verbatim and never +// interprets what is in them, so it works for any type whose row is a whole number of words. +struct vk_op_flash_attn_gather_union_push_constants { + uint32_t n_kv, n_kv_raw, kv_c_max, nbk1, nbm1, n_batch, row_words; +}; +struct vk_op_flash_attn_gather_push_constants { + uint32_t n_kv, n_kv_raw, n_top_k, kv_c; + uint32_t nbk1, nbk3, nbt1, nbt3, nbm1, nbm3, nem3, n_batch, row_words; +}; +static_assert(sizeof(vk_op_flash_attn_gather_push_constants) <= 128); + +struct vk_op_flash_attn_top_k_push_constants { + uint32_t n_batch, n_kv, n_kv_raw, n_top_k, n_head; + uint32_t nbq1, nbq2, nbq3; + uint32_t nbk1, nbk3; + uint32_t nbm1, nbm3; + uint32_t nbt1, nbt3; + uint32_t nb1, nb2, nb3; + float scale; + uint32_t has_sinks; + uint32_t split_mode; +}; +static_assert(sizeof(vk_op_flash_attn_top_k_push_constants) <= 128); + struct vk_op_ssm_scan_push_constants { uint32_t nb02, nb03, nb12, nb13; uint32_t nb21, nb22, nb31; @@ -2084,7 +2190,11 @@ struct vk_quantize_q8_1_push_constants { struct vk_op_flash_attn_split_k_reduce_push_constants { uint32_t D; uint32_t ne1; + // ne2 describes the SPLIT BUFFER (which may cover only a tile of queries); dst_ne2 is the + // destination's query count. They differ only when the caller tiles the split buffer, and + // the destination stride between streams must always use the full count. uint32_t ne2; + uint32_t dst_ne2; uint32_t ne3; uint32_t k_num; uint32_t sinks; @@ -2193,6 +2303,20 @@ static bool vk_enable_sync_logger = false; static uint32_t vk_perf_logger_frequency = 1; static std::string vk_pipeline_stats_filter; +// Total memory traffic of a node (dst + srcs). Used to bound command buffer +// execution time for bandwidth-bound ops with no flops estimate (large copies, +// set_rows, mask fills at long context) - packing too many of them into one +// submission can exceed the driver timeout. +static uint64_t ggml_vk_get_node_bytes(const ggml_tensor * node) { + uint64_t bytes = ggml_nbytes(node); + for (int i = 0; i < GGML_MAX_SRC; i++) { + if (node->src[i]) { + bytes += ggml_nbytes(node->src[i]); + } + } + return bytes; +} + static uint64_t ggml_vk_get_node_flops(const ggml_tensor * node) { if (node->op == GGML_OP_MUL_MAT || node->op == GGML_OP_MUL_MAT_ID) { const uint64_t m = node->ne[0]; @@ -2311,6 +2435,20 @@ class vk_perf_logger { if (node->op == GGML_OP_UNARY) { return fusion_str + ggml_unary_op_name(ggml_get_unary_op(node)); } + if (node->op == GGML_OP_MUL && getenv("GGML_VK_PERF_SHAPES")) { + std::string name = "MUL "; + name += "dst(" + std::to_string(node->ne[0]) + "," + std::to_string(node->ne[1]) + "," + + std::to_string(node->ne[2]) + ") b(" + std::to_string(node->src[1]->ne[0]) + "," + + std::to_string(node->src[1]->ne[1]) + "," + std::to_string(node->src[1]->ne[2]) + ")"; + name += std::string(" a=") + ggml_op_name(node->src[0]->op); + if (node->src[0]->op == GGML_OP_UNARY) { name += std::string(":") + ggml_unary_op_name(ggml_get_unary_op(node->src[0])); } + name += std::string(" b=") + ggml_op_name(node->src[1]->op); + if (node->src[1]->op == GGML_OP_UNARY) { name += std::string(":") + ggml_unary_op_name(ggml_get_unary_op(node->src[1])); } + if (node->src[1]->op == GGML_OP_RESHAPE && node->src[1]->src[0]) { + name += std::string("(") + ggml_op_name(node->src[1]->src[0]->op) + ")"; + } + return fusion_str + name; + } if (node->op == GGML_OP_MUL_MAT || node->op == GGML_OP_MUL_MAT_ID) { const uint64_t m = node->ne[0]; const uint64_t n = node->ne[1]; @@ -2383,6 +2521,13 @@ class vk_perf_logger { timings[name].push_back(time); } + // Log a sub-node interval under a caller-supplied name. Used for work a node's handler + // dispatches before the op itself (e.g. the FA K/V contiguize/dequant pass), which would + // otherwise be billed to the op and invisible. No flops: these move bytes, not math. + void log_timing_named(const char *name, uint64_t time) { + timings[std::string(name)].push_back(time); + } + void log_timing(const std::vector &nodes, const std::vector &names, uint64_t time) { uint64_t total_flops = 0; std::string name; @@ -2416,6 +2561,18 @@ struct ggml_backend_vk_context { ggml_vk_garbage_collector gc; size_t prealloc_size_x, prealloc_size_y, prealloc_size_split_k, prealloc_size_add_rms_partials, prealloc_size_add_rms_partials_offset; vk_buffer prealloc_x, prealloc_y, prealloc_split_k, prealloc_add_rms_partials, sync_staging; + // memoized capacity decision for the FA dequant-once scratch, see ggml_vk_fa_dequant_scratch_fits + uint64_t fa_dequant_gate_sz; + bool fa_dequant_gate_fits; + bool fa_dequant_gate_logged; + // DeepSeek V4 small-batch union: the compact row count is produced on the device, so the + // host prices compaction from the last count it wrote. See ggml_vk_fa_union_estimate. + // Per batch size, because the overlap depends on it (measured 0.64 at 2 tokens, 0.40 at 4, + // 0.24 at 8) and because a speculative decode varies the batch with the accept count, so a + // single slot would be invalidated on nearly every step. This path caps the batch at 64. + vk_buffer fa_union_stat; + float fa_union_est_ratio[64]; // union / candidates, decaying peak; 0 = unseeded + uint64_t fa_union_declines; vk::Fence fence, almost_ready_fence; bool submit_pending {}; bool almost_ready_fence_pending {}; @@ -2475,6 +2632,9 @@ struct ggml_backend_vk_context { std::vector query_fusion_node_count; std::vector query_nodes; std::vector query_node_idx; + // non-null => this query slot closes a sub-node interval logged under this literal name, + // not a graph node. See ggml_vk_perf_mark_subop. + std::vector query_sub_names; int32_t num_queries {}; int32_t query_idx {}; }; @@ -3880,8 +4040,46 @@ static vk_fa_tuning_params get_fa_tuning_params_coopmat1(const vk_device& device result.block_cols = coopmat_block_cols * num_subgroups; result.row_split = num_subgroups; result.subgroup_size = device->subgroup_size; + + // Pin a 32-wide subgroup only where narrowing is free. The shader derives cols_per_iter, + // threads_per_rowgroup and every strided load loop from gl_WorkGroupSize.x, and + // workgroup_size is num_subgroups * subgroup_size, so threads_per_rowgroup always equals the + // real subgroup size. Halving the subgroup halves the workgroup, and the per-lane O state + // grows as d_per_thread = ceil((HSV/4) / threads_per_rowgroup). Pin only when that count is + // unchanged. Above that point the narrow subgroup issues roughly 1.5x to 1.8x the + // instructions for the same number of SIMD passes, which loses on an issue-bound kernel: + // hd256 measures 6 to 18 percent slower. The test depends on HSV only; HSK does not enter + // d_per_thread. On a 64-wide device it reduces exactly to hsv <= 128. + // On by default (=1, applies the rule); =0 disables. =2 forces the pin regardless of + // head size, for measuring the configurations the rule rejects. Diagnostic only. + static const int fa_wave32 = [] { + const char * e = getenv("GGML_VK_FA_WAVE32"); + return e ? atoi(e) : 1; + }(); + // Multi-row dispatches only. Decode dispatches carry N = gqa_ratio rows (8 on + // Qwen3-Coder-30B) and there the narrow subgroup loses: tg64 at d8192/d32768 measured + // 3.5 to 4 percent slower with the pin (q8_0 KV), while prefill gains up to 10 percent + // at d8192. Keep the pin to the prefill shapes it was measured on. + if (fa_wave32 != 0 && + n_rows >= 32 && + device->subgroup_size_control && + 32 < device->subgroup_size && // narrow only, never widen + device->subgroup_min_size <= 32 && 32 <= device->subgroup_max_size && + (result.block_cols % 32) == 0 && // cols_per_thread stays >= 1 + (result.block_cols * result.block_rows / 4) >= num_subgroups * 32 && // mask_cache != 0 + (fa_wave32 == 2 || + CEIL_DIV(hsv / 4, 32u) == CEIL_DIV(hsv / 4, device->subgroup_size))) { + result.subgroup_size = 32; + } + result.workgroup_size = num_subgroups * result.subgroup_size; + // threads_per_rowgroup == the real subgroup size is load-bearing in three places: + // the subgroupMax row reduction, the subgroupAdd of Lf, and tmpsh[gl_SubgroupID], which is + // sized by row_split and would be written out of bounds if gl_NumSubgroups exceeded it. + GGML_ASSERT(result.workgroup_size == result.row_split * result.subgroup_size); + GGML_ASSERT(result.block_cols % result.subgroup_size == 0); + const uint32_t D_lsb = D ^ (D & (D-1)); // extract lowest set bit result.d_split = std::min(std::min(result.subgroup_size, 8u), D_lsb / 4); @@ -3963,14 +4161,16 @@ static vk_fa_tuning_params get_fa_tuning_params(const vk_device& device, uint32_ } static vk_fa_pipeline_state get_fa_pipeline_state(const vk_device& device, const vk_fa_tuning_params& params, uint32_t hsk, uint32_t hsv, bool aligned, bool f32acc, - bool use_mask, bool use_mask_opt, bool use_logit_softcap, ggml_type k_type, ggml_type v_type) { + bool use_mask, bool use_mask_opt, bool use_logit_softcap, ggml_type k_type, ggml_type v_type, + bool use_dynamic_kv = false) { const bool old_amd_windows = device->vendor_id == VK_VENDOR_ID_AMD && device->driver_id == vk::DriverId::eAmdProprietary && (device->architecture == AMD_GCN || device->architecture == AMD_RDNA1 || device->architecture == AMD_RDNA2); uint32_t flags = (use_mask_opt ? 1 : 0) | (use_mask ? 2 : 0) | (use_logit_softcap ? 4 : 0) | - (old_amd_windows ? 8 : 0); + (old_amd_windows ? 8 : 0) | + (use_dynamic_kv ? 16 : 0); const uint32_t subgroup_size = params.disable_subgroups ? 0 : params.subgroup_size; @@ -4294,6 +4494,34 @@ struct CompileTask { uint32_t required_subgroup_size; }; +// EXPERIMENT (GGML_VK_MMID_F16B=1): convert the contiguous f32 activations (B) of +// quantized MUL_MAT_ID to f16 and run the f16-B matmul_id kernels instead of the +// f32-B ones. Halves B bytes and buf_b shared memory (better occupancy) at the +// f32->f16 rounding cost upstream already accepts on the coopmat2 path. +static bool ggml_vk_mmid_f16b_enabled() { + static const bool enabled = [] { + const char * env = getenv("GGML_VK_MMID_F16B"); + return env == nullptr || atoi(env) != 0; // on by default; =0 disables + }(); + return enabled; +} + +// GGML_VK_DENSE_F16B: same idea as GGML_VK_MMID_F16B but for plain MUL_MAT. Halves the B bytes +// moved. Numerically identical: mul_mm stages B into shared FLOAT_TYPE either way, so the f32-B +// kernel already rounds B to f16. Helps wide dense models, costs ~1% on narrow ones. +// 0 = off, 1 = all quantized dense matmuls, 2 = auto (only the K we have positive data for) +static int ggml_vk_dense_f16b_mode() { + static const int mode = [] { + const char * e = getenv("GGML_VK_DENSE_F16B"); + if (e == nullptr) return 0; + if (e[0] == 'a') return 2; + return atoi(e) != 0 ? 1 : 0; + }(); + return mode; +} + +static bool ggml_vk_dense_f16b_enabled() { return ggml_vk_dense_f16b_mode() != 0; } + static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { VK_LOG_DEBUG("ggml_vk_load_shaders(" << device->name << ")"); @@ -4435,6 +4663,23 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { l_warptile = { 256, 128, 128, 16, mm_warp_8, 64, 2, tm_m, tn_m, tk_m, mm_warp_8 }; l_warptile_mmq = l_warptile_mmq_int = { 256, 128, 128, 32, mm_warp_8, 64, 2, tm_m, tn_m, tk_m, mm_warp_8 }; l_warptile_mmq_int_k = { 256, 128, 128, 32, mm_warp_16, 64, 1, 4, 2, 1, mm_warp_16 }; + + // EXPERIMENT (GGML_VK_MMID_WG256=1): the dense large tile above runs 256 threads on a + // 128x128 tile, but the mul_mat_id variants still run 128. Give MoE the same thread + // count per tile: same BM/BN/BK (so same shared memory), twice the threads sharing each + // A/B tile load, half the accumulators per thread. Warp split stays legal: + // (BM/WM)*(BN/WN) == wg/subgroup == 4, WNITER == (WM*WN)/(WARP*TM*TN*WMITER) == 2. + // A 256-expert MoE at ub=2048 sees only ~64 rows per expert, so the tile that actually + // runs is the medium one, not the large one. Override both. + static const char * mmid_wg256_env = getenv("GGML_VK_MMID_WG256"); + if (mmid_wg256_env && atoi(mmid_wg256_env) != 0) { + l_warptile_mmqid = { 256, 128, 128, 32, mul_mat_subgroup_size_8, 64, 2, tm_m, tn_m, tk_m, mul_mat_subgroup_size_8 }; + l_warptile_mmqid_int = { 256, 128, 128, 32, mul_mat_subgroup_size_8, 64, 2, 4, 4, 1, mul_mat_subgroup_size_8 }; + // BM=BN=64 at 4 warps needs WM=WN=32: (BM/WM)*(BN/WN) == 4, cms_per_row/col == 2. + m_warptile_mmqid = { 256, 64, 64, 32, 32, 32, 2, tm_m, tn_m, tk_m, mul_mat_subgroup_size_8 }; + m_warptile_mmqid_int = { 256, 64, 64, 32, 32, 32, 2, 2, 2, 1, mul_mat_subgroup_size_8 }; + fprintf(stderr, "ggml_vulkan: MUL_MAT_ID medium+large tiles at 256 threads (GGML_VK_MMID_WG256)\n"); + } } else if (device->vendor_id == VK_VENDOR_ID_INTEL && device->coopmat_support) { // Xe2/Xe3 with coopmat enabled - warptile performance tuning l_warptile = { 512, 128, 128, 16, mm_warp_8, 32, 2, tm_m, tn_m, tk_m, mm_warp_8 }; @@ -4645,7 +4890,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { } name = aligned ? "flash_attn_f32_f16_aligned" : "flash_attn_f32_f16"; } - ggml_vk_create_pipeline(device, fa.second, name, spv_size, spv_data, "main", 7, + ggml_vk_create_pipeline(device, fa.second, name, spv_size, spv_data, "main", 8, sizeof(vk_flash_attn_push_constants), {Br, 1, 1}, get_fa_spec_constants(fa.first), aligned ? Bc : 1, true, !fa_ds, !fa_ds ? fa_sgs : 0); @@ -4681,7 +4926,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { else { spv_data = flash_attn_f32_f16_f16acc_cm1_data; spv_size = flash_attn_f32_f16_f16acc_cm1_len; } name = aligned ? "flash_attn_f32_f16_aligned_cm1" : "flash_attn_f32_f16_cm1"; } - ggml_vk_create_pipeline(device, fa.second, name, spv_size, spv_data, "main", 7, + ggml_vk_create_pipeline(device, fa.second, name, spv_size, spv_data, "main", 8, sizeof(vk_flash_attn_push_constants), {Br, 1, 1}, get_fa_spec_constants(fa.first), aligned ? Bc : 1, true, !fa_ds, !fa_ds ? fa_sgs : 0); @@ -4718,7 +4963,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { if (f32acc) { spv_data = flash_attn_f32_f16_cm2_data; spv_size = flash_attn_f32_f16_cm2_len; name = "flash_attn_f32_f16_f32acc_cm2"; } else { spv_data = flash_attn_f32_f16_f16acc_cm2_data; spv_size = flash_attn_f32_f16_f16acc_cm2_len; name = "flash_attn_f32_f16_f16acc_cm2"; } } - ggml_vk_create_pipeline(device, fa.second, name, spv_size, spv_data, "main", 7, + ggml_vk_create_pipeline(device, fa.second, name, spv_size, spv_data, "main", 8, sizeof(vk_flash_attn_push_constants), {Br, 1, 1}, get_fa_spec_constants(fa.first), aligned ? Bc : 1, true, false, 0); } @@ -4737,7 +4982,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { return spec; }; - const int mul_mat_id_param_count = 5; + const int mul_mat_id_param_count = 6; // a, b, d, ids, expert_counts, fused scale #if defined(VK_NV_cooperative_matrix2) && defined(GGML_VULKAN_COOPMAT2_GLSLC_SUPPORT) if (device->coopmat2) { @@ -4813,48 +5058,48 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { GGML_ASSERT(device->subgroup_ballot); - CREATE_MM2(pipeline_matmul_id_f16, matmul_id_subgroup_f16, wg_denoms, warptile, vk_mat_mat_id_push_constants, 5) + CREATE_MM2(pipeline_matmul_id_f16, matmul_id_subgroup_f16, wg_denoms, warptile, vk_mat_mat_id_push_constants, mul_mat_id_param_count) #if defined(GGML_VULKAN_BFLOAT16_GLSLC_SUPPORT) if (device->coopmat_bf16_support) { - CREATE_MM(pipeline_matmul_id_bf16, matmul_id_subgroup_bf16, , wg_denoms, warptile, vk_mat_mat_id_push_constants, 5) + CREATE_MM(pipeline_matmul_id_bf16, matmul_id_subgroup_bf16, , wg_denoms, warptile, vk_mat_mat_id_push_constants, mul_mat_id_param_count) } #endif - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q1_0], matmul_id_subgroup_q1_0_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q2_0], matmul_id_subgroup_q2_0_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_0], matmul_id_subgroup_q4_0_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_1], matmul_id_subgroup_q4_1_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q5_0], matmul_id_subgroup_q5_0_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q5_1], matmul_id_subgroup_q5_1_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q8_0], matmul_id_subgroup_q8_0_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q2_K], matmul_id_subgroup_q2_k_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_TQ2_0], matmul_id_subgroup_tq2_0_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q3_K], matmul_id_subgroup_q3_k_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_K], matmul_id_subgroup_q4_k_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q5_K], matmul_id_subgroup_q5_k_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q6_K], matmul_id_subgroup_q6_k_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ1_S], matmul_id_subgroup_iq1_s_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ1_M], matmul_id_subgroup_iq1_m_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ2_XXS], matmul_id_subgroup_iq2_xxs_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ2_XS], matmul_id_subgroup_iq2_xs_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ2_S], matmul_id_subgroup_iq2_s_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ3_XXS], matmul_id_subgroup_iq3_xxs_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ3_S], matmul_id_subgroup_iq3_s_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_XS], matmul_id_subgroup_iq4_xs_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_NL], matmul_id_subgroup_iq4_nl_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_0_ROCMFP4], matmul_id_subgroup_rocmfp4_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_0_ROCMFP4_FAST], matmul_id_subgroup_rocmfp4_fast_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q3_0_ROCMFPX], matmul_id_subgroup_rocmfpx_fp3_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q6_0_ROCMFPX], matmul_id_subgroup_rocmfpx_fp6_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q8_0_ROCMFPX], matmul_id_subgroup_rocmfpx_fp8_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q1_0], matmul_id_subgroup_q1_0_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q2_0], matmul_id_subgroup_q2_0_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_0], matmul_id_subgroup_q4_0_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_1], matmul_id_subgroup_q4_1_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q5_0], matmul_id_subgroup_q5_0_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q5_1], matmul_id_subgroup_q5_1_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q8_0], matmul_id_subgroup_q8_0_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q2_K], matmul_id_subgroup_q2_k_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_TQ2_0], matmul_id_subgroup_tq2_0_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q3_K], matmul_id_subgroup_q3_k_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_K], matmul_id_subgroup_q4_k_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q5_K], matmul_id_subgroup_q5_k_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q6_K], matmul_id_subgroup_q6_k_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ1_S], matmul_id_subgroup_iq1_s_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ1_M], matmul_id_subgroup_iq1_m_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ2_XXS], matmul_id_subgroup_iq2_xxs_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ2_XS], matmul_id_subgroup_iq2_xs_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ2_S], matmul_id_subgroup_iq2_s_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ3_XXS], matmul_id_subgroup_iq3_xxs_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ3_S], matmul_id_subgroup_iq3_s_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_XS], matmul_id_subgroup_iq4_xs_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_NL], matmul_id_subgroup_iq4_nl_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_0_ROCMFP4], matmul_id_subgroup_rocmfp4_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_0_ROCMFP4_FAST], matmul_id_subgroup_rocmfp4_fast_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q3_0_ROCMFPX], matmul_id_subgroup_rocmfpx_fp3_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q6_0_ROCMFPX], matmul_id_subgroup_rocmfpx_fp6_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q8_0_ROCMFPX], matmul_id_subgroup_rocmfpx_fp8_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) #if defined(GGML_VULKAN_FLOAT_E2M1_GLSLC_SUPPORT) && defined(GGML_VULKAN_FLOAT_E4M3_GLSLC_SUPPORT) if (device->ocp_fp4) { - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_MXFP4], matmul_id_subgroup_mxfp4_f16_ocp, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_NVFP4], matmul_id_subgroup_nvfp4_f16_ocp, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_MXFP4], matmul_id_subgroup_mxfp4_f16_ocp, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_NVFP4], matmul_id_subgroup_nvfp4_f16_ocp, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) } else #endif { - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_MXFP4], matmul_id_subgroup_mxfp4_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) - CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_NVFP4], matmul_id_subgroup_nvfp4_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, 5) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_MXFP4], matmul_id_subgroup_mxfp4_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) + CREATE_MM2(pipeline_dequant_mul_mat_mat_id[GGML_TYPE_NVFP4], matmul_id_subgroup_nvfp4_f16, mmqid_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count) } #undef CREATE_MM #undef CREATE_MM2 @@ -4862,20 +5107,97 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { #endif // defined(VK_NV_cooperative_matrix2) && defined(GGML_VULKAN_COOPMAT2_GLSLC_SUPPORT) #if defined(VK_KHR_cooperative_matrix) && defined(GGML_VULKAN_COOPMAT_GLSLC_SUPPORT) if (device->coopmat_support) { + // Deterministic subgroup sizing for the dense coopmat pipelines. Two parts: + // + // (1) The required subgroup size is the tile's own WARP, not the driver's choice. The cm1 + // shaders derive their warp grid from the real subgroup (warp_i = gl_SubgroupID) and + // size shared arrays with NUM_WARPS = BLOCK_SIZE / WARP, so WARP and the actual + // subgroup must agree - leaving that to the driver makes the agreement incidental. + // + // (2) The QUANTISED dense tiles run at wave32. RDNA3.x WMMA is wave32-native, so a wave64 + // subgroup issues each coopmat op as two halves. BLOCK_SIZE is kept, so the subgroup + // count doubles and WM (or WN) halves until the warp grid tiles BM x BN again: + // NUM_WARPS == (BM/WM)*(BN/WN). A tile that cannot be retiled is left at wave64. + // + // The FLOAT tiles are deliberately excluded. Measured on gfx1151 with a standalone + // MUL_MAT microbench at both dense FFN shapes (m=25600 k=5120, m=5120 k=25600): + // q6_K +5.2..+10.8%, q8_0 +5.4..+8.4%, q4_K +0.7..+9.1%, q4_0 -1.5..+1.8%, but + // f16 -6.7..+6.4% and bf16 ~0. The win tracks inline dequant instruction count + // (q6_K 3907 -> 3433 instructions at wave32, identical 192 VGPRs and 8 subgroups/SIMD), + // so it lands on the issue-bound quantised kernels and not on the float ones, which are + // bandwidth-bound on the weight stream. + // + // Scoped to AMD coopmat1 on a wave64 default; other vendors keep the driver default, + // since none of the above is validated there. + // GGML_VK_DENSE_WAVE32=0 disables, =2 additionally retiles the float tiles (probe). + const bool dense_sgs_scope = + device->vendor_id == VK_VENDOR_ID_AMD && + device->driver_id != vk::DriverId::eAmdProprietary && + device->subgroup_size_control; + const bool dense_wave32_possible = + dense_sgs_scope && + device->subgroup_min_size <= 32 && 32 <= device->subgroup_max_size && + device->subgroup_size > 32; + const char * dense_wave32_env = getenv("GGML_VK_DENSE_WAVE32"); + const int dense_wave32 = dense_wave32_env ? atoi(dense_wave32_env) : 1; + + // The mul_mat_id tiles below start from these. Snapshot them before the dense + // wave32 rewrite so an mmid pipeline created without a required subgroup size + // (GGML_VK_MMID_WAVE32=0) keeps a WARP that matches the real wave64 subgroup. + const auto l_warptile_mmq_w64 = l_warptile_mmq; + const auto m_warptile_mmq_w64 = m_warptile_mmq; + const auto s_warptile_mmq_w64 = s_warptile_mmq; + if (dense_wave32_possible && dense_wave32 != 0) { + auto wave32_tile = [](std::vector & w) -> bool { + // {BLOCK_SIZE, BM, BN, BK, WM, WN, WMITER, TM, TN, TK, WARP} + std::vector t = w; + t[10] = 32; + for (int guard = 0; guard < 4 && t[0] / t[10] != (t[1] / t[4]) * (t[2] / t[5]); ++guard) { + if (t[4] >= t[5] && t[4] > t[7]) { + t[4] /= 2; // halve WM, keeping WM >= TM + } else { + t[5] /= 2; // halve WN + } + } + if (t[0] / t[10] != (t[1] / t[4]) * (t[2] / t[5]) || t[4] < t[7] || t[5] < t[8]) { + return false; + } + w = t; + return true; + }; + wave32_tile(l_warptile_mmq); + wave32_tile(m_warptile_mmq); + wave32_tile(s_warptile_mmq); + if (dense_wave32 >= 2) { + wave32_tile(l_warptile); + wave32_tile(m_warptile); + wave32_tile(s_warptile); + } + } + + // WARP -> required subgroup size, or 0 where the device cannot honor one. + auto dense_req_sgs = [dense_sgs_scope, &device](const std::vector & w) -> uint32_t { + const uint32_t warp = w[10]; + if (!dense_sgs_scope || warp < device->subgroup_min_size || warp > device->subgroup_max_size) { + return 0; + } + return warp; + }; + // Create 6 variants, {s,m,l}x{unaligned,aligned} #define CREATE_MM(TYPE, PIPELINE_NAME, NAMELC, F16ACC, WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \ if (device->mul_mat ## ID ## _l[TYPE]) \ - ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->l, #NAMELC #F16ACC "_l", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), l_ ## WG_DENOMS, ggml_vk_mul_mm_spec(l_ ## WARPTILE, false), 1, false, true); \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->l, #NAMELC #F16ACC "_l", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), l_ ## WG_DENOMS, ggml_vk_mul_mm_spec(l_ ## WARPTILE, false), 1, false, true, dense_req_sgs(l_ ## WARPTILE)); \ if (device->mul_mat ## ID ## _m[TYPE]) \ - ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->m, #NAMELC #F16ACC "_m", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), m_ ## WG_DENOMS, ggml_vk_mul_mm_spec(m_ ## WARPTILE, false), 1, false, true); \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->m, #NAMELC #F16ACC "_m", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), m_ ## WG_DENOMS, ggml_vk_mul_mm_spec(m_ ## WARPTILE, false), 1, false, true, dense_req_sgs(m_ ## WARPTILE)); \ if (device->mul_mat ## ID ## _s[TYPE]) \ - ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->s, #NAMELC #F16ACC "_s", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, ggml_vk_mul_mm_spec(s_ ## WARPTILE, false), 1, false, true); \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->s, #NAMELC #F16ACC "_s", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, ggml_vk_mul_mm_spec(s_ ## WARPTILE, false), 1, false, true, dense_req_sgs(s_ ## WARPTILE)); \ if (device->mul_mat ## ID ## _l[TYPE]) \ - ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->a_l, #NAMELC #F16ACC "_aligned_l", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), l_ ## WG_DENOMS, ggml_vk_mul_mm_spec(l_ ## WARPTILE, true), l_align, false, true); \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->a_l, #NAMELC #F16ACC "_aligned_l", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), l_ ## WG_DENOMS, ggml_vk_mul_mm_spec(l_ ## WARPTILE, true), l_align, false, true, dense_req_sgs(l_ ## WARPTILE)); \ if (device->mul_mat ## ID ## _m[TYPE]) \ - ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->a_m, #NAMELC #F16ACC "_aligned_m", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), m_ ## WG_DENOMS, ggml_vk_mul_mm_spec(m_ ## WARPTILE, true), m_align, false, true); \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->a_m, #NAMELC #F16ACC "_aligned_m", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), m_ ## WG_DENOMS, ggml_vk_mul_mm_spec(m_ ## WARPTILE, true), m_align, false, true, dense_req_sgs(m_ ## WARPTILE)); \ if (device->mul_mat ## ID ## _s[TYPE]) \ - ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->a_s, #NAMELC #F16ACC "_aligned_s", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, ggml_vk_mul_mm_spec(s_ ## WARPTILE, true), s_align, false, true); \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->a_s, #NAMELC #F16ACC "_aligned_s", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, ggml_vk_mul_mm_spec(s_ ## WARPTILE, true), s_align, false, true, dense_req_sgs(s_ ## WARPTILE)); \ // Create 2 variants, {f16,f32} accumulator #define CREATE_MM2(TYPE, PIPELINE_NAME, NAMELC, WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \ @@ -4910,6 +5232,21 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MM2(GGML_TYPE_Q4_K, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q4_K], matmul_q4_k_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); CREATE_MM2(GGML_TYPE_Q5_K, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q5_K], matmul_q5_k_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); CREATE_MM2(GGML_TYPE_Q6_K, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q6_K], matmul_q6_k_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + + // f16-B variants of the same quant pipelines. The _f16 SPIR-V is already built for + // every type; upstream only instantiates it in the coopmat2 branch. + if (ggml_vk_dense_f16b_enabled()) { + CREATE_MM2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_f16[GGML_TYPE_Q4_0], matmul_q4_0_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_f16[GGML_TYPE_Q4_1], matmul_q4_1_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM2(GGML_TYPE_Q5_0, pipeline_dequant_mul_mat_mat_f16[GGML_TYPE_Q5_0], matmul_q5_0_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM2(GGML_TYPE_Q5_1, pipeline_dequant_mul_mat_mat_f16[GGML_TYPE_Q5_1], matmul_q5_1_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_f16[GGML_TYPE_Q8_0], matmul_q8_0_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM2(GGML_TYPE_Q2_K, pipeline_dequant_mul_mat_mat_f16[GGML_TYPE_Q2_K], matmul_q2_k_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM2(GGML_TYPE_Q3_K, pipeline_dequant_mul_mat_mat_f16[GGML_TYPE_Q3_K], matmul_q3_k_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM2(GGML_TYPE_Q4_K, pipeline_dequant_mul_mat_mat_f16[GGML_TYPE_Q4_K], matmul_q4_k_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM2(GGML_TYPE_Q5_K, pipeline_dequant_mul_mat_mat_f16[GGML_TYPE_Q5_K], matmul_q5_k_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM2(GGML_TYPE_Q6_K, pipeline_dequant_mul_mat_mat_f16[GGML_TYPE_Q6_K], matmul_q6_k_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + } CREATE_MM2(GGML_TYPE_IQ1_S, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ1_S], matmul_iq1_s_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); CREATE_MM2(GGML_TYPE_IQ1_M, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ1_M], matmul_iq1_m_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); CREATE_MM2(GGML_TYPE_IQ2_XXS, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ2_XXS], matmul_iq2_xxs_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); @@ -4947,6 +5284,117 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { } #endif + // EXPERIMENT (GGML_VK_MMID_TILE16=1): narrow the matmul_id small tile to BN=16. + // MoE prefill leaves ~nei1*nei0/n_expert rows per expert (~16 at ub512/256E/8a), + // so even the 32-wide small tile runs half empty. Only meaningful stacked on + // GGML_VK_MMID_SMALLN=1 (routes mmid to the small tile) + the row-list prepass. + // Shadows the s-tile config for the mmid quant pipelines only; dense unaffected. + auto s_warptile_mmq_id16 = s_warptile_mmq_w64; + auto s_mmq_wg_denoms_id16 = s_mmq_wg_denoms; + auto m_warptile_mmq_id128 = m_warptile_mmq_w64; + auto m_mmq_wg_denoms_id128 = m_mmq_wg_denoms; + auto l_warptile_mmq_idw = l_warptile_mmq_w64; + uint32_t mmid_req_sgs = 0; + { + const char * tile16_env = getenv("GGML_VK_MMID_TILE16"); + if (tile16_env && atoi(tile16_env) != 0) { + s_warptile_mmq_id16[2] = 16; // BN + s_warptile_mmq_id16[5] = 16; // WN + s_mmq_wg_denoms_id16[1] = 16; + } + // GGML_VK_MMID_BM64=1: taller small tile (BM 32->64, two warps). Same + // A-traffic (ic-tile count unchanged), halves per-expert B re-reads + // (ir-tile count), larger WGs to hide latency. + const char * bm64_env = getenv("GGML_VK_MMID_BM64"); + if (!bm64_env || atoi(bm64_env) != 0) { // on by default; =0 disables + s_warptile_mmq_id16[0] = 2 * mul_mat_subgroup_size; // BLOCK_SIZE + s_warptile_mmq_id16[1] = 64; // BM + s_mmq_wg_denoms_id16[0] = 64; + } + // GGML_VK_MMID_M128=1: same idea for the medium tile (BM 64->128, four + // warps) — the tile the per-expert-n heuristic picks at n~64 (e.g. ub2048). + const char * m128_env = getenv("GGML_VK_MMID_M128"); + if (!m128_env || atoi(m128_env) != 0) { // on by default; =0 disables + m_warptile_mmq_id128[0] = 4 * mul_mat_subgroup_size; // BLOCK_SIZE + m_warptile_mmq_id128[1] = 128; // BM + m_mmq_wg_denoms_id128[0] = 128; + } + // GGML_VK_MMID_WAVE32=1: force required subgroup size 32 on the mmid + // quant coopmat pipelines (RDNA3.x WMMA is wave32-native; probe whether + // RADV lowers KHR_coopmat better at wave32). The cm1 path derives its + // warp grid from the real subgroup (warp_i = gl_SubgroupID, tiw = + // gl_SubgroupInvocationID) and sizes shared arrays (coopmat_stage, + // ballots_sh) with NUM_WARPS = BLOCK_SIZE / WARP, so the WARP spec + // constant must equal the forced size, and the coverage invariant + // NUM_WARPS == (BM/WM)*(BN/WN) must be restored. BLOCK_SIZE is kept + // (same workgroup shape, load loops and shmem as the wave64 stack), so + // the subgroup count doubles and WM (or WN) is halved until the warp + // grid exactly tiles BM x BN again. Per-lane coopmat accumulator + // footprint is unchanged: half the lanes per subgroup, half the + // (WM/TM)*(WN/TN) fragments per subgroup. Runs after the BM64/M128 + // gates so it composes with the probe stack. Applies only when the + // driver honors a required size (subgroup_size_control covering 32); + // otherwise WARP=32 with a real subgroup of 64 would corrupt tiling. + const char * wave32_env = getenv("GGML_VK_MMID_WAVE32"); + if ((!wave32_env || atoi(wave32_env) != 0) && device->subgroup_size_control && // on by default; =0 disables + device->subgroup_min_size <= 32 && 32 <= device->subgroup_max_size) { + mmid_req_sgs = 32; + auto wave32_tile = [](std::vector &w) { + // {BLOCK_SIZE, BM, BN, BK, WM, WN, WMITER, TM, TN, TK, WARP} + w[10] = 32; // WARP: must match the forced subgroup size + for (int guard = 0; guard < 4 && w[0] / w[10] != (w[1] / w[4]) * (w[2] / w[5]); ++guard) { + if (w[4] >= w[5] && w[4] > w[7]) { + w[4] /= 2; // halve WM, keeping WM >= TM + } else { + w[5] /= 2; // halve WN + } + } + GGML_ASSERT(w[0] / w[10] == (w[1] / w[4]) * (w[2] / w[5])); // NUM_WARPS == (BM/WM)*(BN/WN) + GGML_ASSERT(w[4] >= w[7] && w[5] >= w[8]); // WM >= TM, WN >= TN + }; + wave32_tile(s_warptile_mmq_id16); + wave32_tile(m_warptile_mmq_id128); + wave32_tile(l_warptile_mmq_idw); + } else { + // Wave64 stack: the tiles are the dense wave64 tiles plus the BM64/M128 reshapes, + // which keep NUM_WARPS == (BM/WM)*(BN/WN). Pin the subgroup to the tile's WARP + // where the driver honours it, so the spec constant can never disagree with + // the real subgroup size. + for (const auto * w : { &s_warptile_mmq_id16, &m_warptile_mmq_id128, &l_warptile_mmq_idw }) { + GGML_ASSERT((*w)[0] / (*w)[10] == ((*w)[1] / (*w)[4]) * ((*w)[2] / (*w)[5])); + } + const uint32_t warp = s_warptile_mmq_id16[10]; + if (device->subgroup_size_control && warp == m_warptile_mmq_id128[10] && warp == l_warptile_mmq_idw[10] && + device->subgroup_min_size <= warp && warp <= device->subgroup_max_size) { + mmid_req_sgs = warp; + } + } + } + { + const auto &s_warptile_mmq = s_warptile_mmq_id16; + const auto &s_mmq_wg_denoms = s_mmq_wg_denoms_id16; + const auto &m_warptile_mmq = m_warptile_mmq_id128; + const auto &m_mmq_wg_denoms = m_mmq_wg_denoms_id128; + const auto &l_warptile_mmq = l_warptile_mmq_idw; + + // Same expansion as CREATE_MM above, plus a trailing required subgroup + // size (0 = driver default) for the GGML_VK_MMID_WAVE32 probe. Scoped to + // the mmid quant pipelines below; dense pipelines are untouched. +#undef CREATE_MM +#define CREATE_MM(TYPE, PIPELINE_NAME, NAMELC, F16ACC, WG_DENOMS, WARPTILE, PUSHCONST, PARAMCOUNT, ID) \ + if (device->mul_mat ## ID ## _l[TYPE]) \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->l, #NAMELC #F16ACC "_l", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), l_ ## WG_DENOMS, ggml_vk_mul_mm_spec(l_ ## WARPTILE, false), 1, false, true, mmid_req_sgs); \ + if (device->mul_mat ## ID ## _m[TYPE]) \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->m, #NAMELC #F16ACC "_m", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), m_ ## WG_DENOMS, ggml_vk_mul_mm_spec(m_ ## WARPTILE, false), 1, false, true, mmid_req_sgs); \ + if (device->mul_mat ## ID ## _s[TYPE]) \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->s, #NAMELC #F16ACC "_s", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, ggml_vk_mul_mm_spec(s_ ## WARPTILE, false), 1, false, true, mmid_req_sgs); \ + if (device->mul_mat ## ID ## _l[TYPE]) \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->a_l, #NAMELC #F16ACC "_aligned_l", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), l_ ## WG_DENOMS, ggml_vk_mul_mm_spec(l_ ## WARPTILE, true), l_align, false, true, mmid_req_sgs); \ + if (device->mul_mat ## ID ## _m[TYPE]) \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->a_m, #NAMELC #F16ACC "_aligned_m", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), m_ ## WG_DENOMS, ggml_vk_mul_mm_spec(m_ ## WARPTILE, true), m_align, false, true, mmid_req_sgs); \ + if (device->mul_mat ## ID ## _s[TYPE]) \ + ggml_vk_create_pipeline(device, device-> PIPELINE_NAME ->a_s, #NAMELC #F16ACC "_aligned_s", NAMELC ## F16ACC ## _cm1_len, NAMELC ## F16ACC ## _cm1_data, "main", PARAMCOUNT, sizeof(PUSHCONST), s_ ## WG_DENOMS, ggml_vk_mul_mm_spec(s_ ## WARPTILE, true), s_align, false, true, mmid_req_sgs); \ + CREATE_MM2(GGML_TYPE_Q1_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q1_0], matmul_id_subgroup_q1_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MM2(GGML_TYPE_Q2_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q2_0], matmul_id_subgroup_q2_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MM2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_0], matmul_id_subgroup_q4_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); @@ -4984,6 +5432,37 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MM2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_MXFP4], matmul_id_subgroup_mxfp4_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MM2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_NVFP4], matmul_id_subgroup_nvfp4_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); } + + // EXPERIMENT (GGML_VK_MMID_F16B=1): f16-B variants of the quant matmul_id + // pipelines. The _f16 SPIR-V exists for every quant type; upstream just never + // instantiates it in the KHR_coopmat branch. Same warptiles as the f32-B lines + // (including the SMALLN/BM64/M128 tile experiments shadowed above). + if (ggml_vk_mmid_f16b_enabled()) { + CREATE_MM2(GGML_TYPE_Q1_0, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_Q1_0], matmul_id_subgroup_q1_0_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_Q4_0], matmul_id_subgroup_q4_0_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_Q4_1], matmul_id_subgroup_q4_1_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_Q5_0, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_Q5_0], matmul_id_subgroup_q5_0_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_Q5_1, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_Q5_1], matmul_id_subgroup_q5_1_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_Q8_0, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_Q8_0], matmul_id_subgroup_q8_0_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_Q2_K, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_Q2_K], matmul_id_subgroup_q2_k_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_Q3_K, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_Q3_K], matmul_id_subgroup_q3_k_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_Q4_K, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_Q4_K], matmul_id_subgroup_q4_k_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_Q5_K, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_Q5_K], matmul_id_subgroup_q5_k_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_Q6_K, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_Q6_K], matmul_id_subgroup_q6_k_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_IQ1_S, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_IQ1_S], matmul_id_subgroup_iq1_s_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_IQ1_M, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_IQ1_M], matmul_id_subgroup_iq1_m_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_IQ2_XXS, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_IQ2_XXS], matmul_id_subgroup_iq2_xxs_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_IQ2_XS, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_IQ2_XS], matmul_id_subgroup_iq2_xs_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_IQ2_S, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_IQ2_S], matmul_id_subgroup_iq2_s_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_IQ3_XXS, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_IQ3_XXS], matmul_id_subgroup_iq3_xxs_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_IQ3_S, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_IQ3_S], matmul_id_subgroup_iq3_s_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_IQ4_XS, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_IQ4_XS], matmul_id_subgroup_iq4_xs_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_IQ4_NL], matmul_id_subgroup_iq4_nl_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_MXFP4, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_MXFP4], matmul_id_subgroup_mxfp4_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_NVFP4, pipeline_dequant_mul_mat_mat_id_f16b[GGML_TYPE_NVFP4], matmul_id_subgroup_nvfp4_f16, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + } + } + #undef CREATE_MM2 #undef CREATE_MM } else @@ -5618,8 +6097,13 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q4_1], "dequant_q4_1", dequant_q4_1_len, dequant_q4_1_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q5_0], "dequant_q5_0", dequant_q5_0_len, dequant_q5_0_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q5_1], "dequant_q5_1", dequant_q5_1_len, dequant_q5_1_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_dequant_transpose[GGML_TYPE_Q4_0], "dequant_q4_0_transpose", dequant_q4_0_transpose_len, dequant_q4_0_transpose_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_dequant_transpose[GGML_TYPE_Q4_1], "dequant_q4_1_transpose", dequant_q4_1_transpose_len, dequant_q4_1_transpose_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_dequant_transpose[GGML_TYPE_Q5_0], "dequant_q5_0_transpose", dequant_q5_0_transpose_len, dequant_q5_0_transpose_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_dequant_transpose[GGML_TYPE_Q5_1], "dequant_q5_1_transpose", dequant_q5_1_transpose_len, dequant_q5_1_transpose_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q8_0], "dequant_q8_0", dequant_q8_0_len, dequant_q8_0_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant_transpose[GGML_TYPE_Q8_0], "dequant_q8_0_transpose", dequant_q8_0_transpose_len, dequant_q8_0_transpose_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_dequant_transpose[GGML_TYPE_F16], "dequant_f16_transpose", dequant_f16_transpose_len, dequant_f16_transpose_data, "main", 2, 5 * sizeof(uint32_t), {256 * 8, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q2_K], "dequant_q2_k", dequant_q2_k_len, dequant_q2_k_data, "main", 2, 5 * sizeof(uint32_t), {256 * 64, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_TQ2_0], "dequant_tq2_0", dequant_tq2_0_len, dequant_tq2_0_data, "main", 2, 5 * sizeof(uint32_t), {256 * 64, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q3_K], "dequant_q3_k", dequant_q3_k_len, dequant_q3_k_data, "main", 2, 5 * sizeof(uint32_t), {256 * 64, 1, 1}, {}, 1); @@ -5640,6 +6124,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q3_0_ROCMFPX], "dequant_rocmfpx_fp3", dequant_rocmfpx_fp3_len, dequant_rocmfpx_fp3_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q6_0_ROCMFPX], "dequant_rocmfpx_fp6", dequant_rocmfpx_fp6_len, dequant_rocmfpx_fp6_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q8_0_ROCMFPX], "dequant_rocmfpx_fp8", dequant_rocmfpx_fp8_len, dequant_rocmfpx_fp8_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_dequant_transpose[GGML_TYPE_IQ4_NL], "dequant_iq4_nl_transpose", dequant_iq4_nl_transpose_len, dequant_iq4_nl_transpose_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_MXFP4], "dequant_mxfp4", dequant_mxfp4_len, dequant_mxfp4_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_NVFP4], "dequant_nvfp4", dequant_nvfp4_len, dequant_nvfp4_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); @@ -5868,6 +6353,8 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { ggml_vk_create_pipeline(device, device->pipeline_concat_i8, "concat_i8", concat_i8_len, concat_i8_data, "main", 3, sizeof(vk_op_binary_push_constants), {512, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_concat_i16, "concat_i16", concat_i16_len, concat_i16_data, "main", 3, sizeof(vk_op_binary_push_constants), {512, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_concat_i32, "concat_i32", concat_i32_len, concat_i32_data, "main", 3, sizeof(vk_op_binary_push_constants), {512, 1, 1}, {}, 1); + // One workgroup per 32x32 tile: elements are passed as (rows, cols, 1). + ggml_vk_create_pipeline(device, device->pipeline_concat_transpose_i32, "concat_transpose_i32", concat_transpose_i32_len, concat_transpose_i32_data, "main", 3, sizeof(vk_op_binary_push_constants), {32, 32, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_concat_i64, "concat_i64", concat_i64_len, concat_i64_data, "main", 3, sizeof(vk_op_binary_push_constants), {512, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_upscale_nearest_f32, "upscale_f32", upscale_f32_len, upscale_f32_data, "main", 2, sizeof(vk_op_upscale_push_constants), {512, 1, 1}, {GGML_SCALE_MODE_NEAREST}, 1); @@ -6185,6 +6672,92 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { } } + if (device->subgroup_arithmetic && device->subgroup_size == 64) { + for (size_t nhi = 0; nhi < LI_NH_COUNT; ++nhi) { + ggml_vk_create_pipeline(device, device->pipeline_lightning_indexer_f16[nhi], + "lightning_indexer_f16", lightning_indexer_f16_len, lightning_indexer_f16_data, "main", 5, + sizeof(vk_op_lightning_indexer_cm_push_constants), {8, 1, 1}, {device->subgroup_size, LI_NH_VALUES[nhi]}, 1, true, true, + device->subgroup_size); + } +#if defined(VK_KHR_cooperative_matrix) && defined(GGML_VULKAN_COOPMAT_GLSLC_SUPPORT) + if (device->coopmat_support && device->coopmat_support_16x16x16_f32acc && device->subgroup_size_control) { + for (size_t nhi = 0; nhi < LI_NH_COUNT; ++nhi) { + ggml_vk_create_pipeline(device, device->pipeline_lightning_indexer_cm_small_f16[nhi], + "lightning_indexer_cm_small_f16", lightning_indexer_cm_small_f16_len, lightning_indexer_cm_small_f16_data, "main", 5, + sizeof(vk_op_lightning_indexer_cm_push_constants), {16, 16, 1}, {device->subgroup_size, LI_NH_VALUES[nhi]}, 1, true, true, + device->subgroup_size); + } + if (device->properties.limits.maxComputeWorkGroupInvocations >= 512 && + device->properties.limits.maxComputeWorkGroupSize[0] >= 512 && + device->properties.limits.maxComputeSharedMemorySize >= 64 * 1024) { + for (size_t nhi = 0; nhi < LI_NH_COUNT; ++nhi) { + ggml_vk_create_pipeline(device, device->pipeline_lightning_indexer_cm_f16[nhi], + "lightning_indexer_cm_f16", lightning_indexer_cm_f16_len, lightning_indexer_cm_f16_data, "main", 5, + sizeof(vk_op_lightning_indexer_cm_push_constants), {128, 16, 1}, {device->subgroup_size, LI_NH_VALUES[nhi]}, 1, true, true, + device->subgroup_size); + } + } + for (size_t nhi = 0; nhi < LI_NH_COUNT; ++nhi) { + ggml_vk_create_pipeline(device, device->pipeline_lightning_indexer_decode_cm_f16[nhi], + "lightning_indexer_decode_cm_f16", lightning_indexer_decode_cm_f16_len, lightning_indexer_decode_cm_f16_data, "main", 5, + sizeof(vk_op_lightning_indexer_cm_push_constants), {16, 1, 1}, {device->subgroup_size, LI_NH_VALUES[nhi]}, 1, true, true, + device->subgroup_size); + } + ggml_vk_create_pipeline(device, device->pipeline_flash_attn_top_k_cm_f16, + "flash_attn_top_k_cm_f16", flash_attn_top_k_cm_f16_len, flash_attn_top_k_cm_f16_data, "main", 6, + sizeof(vk_op_flash_attn_top_k_push_constants), {1, 1, 1}, {512, device->subgroup_size}, 1, true, true, + device->subgroup_size); + } +#endif + ggml_vk_create_pipeline(device, device->pipeline_flash_attn_top_k_f16, + "flash_attn_top_k_f16", flash_attn_top_k_f16_len, flash_attn_top_k_f16_data, "main", 6, + sizeof(vk_op_flash_attn_top_k_push_constants), {1, 1, 1}, {512, device->subgroup_size}, 1, true, true, + device->subgroup_size); + ggml_vk_create_pipeline(device, device->pipeline_flash_attn_gather_f16, + "flash_attn_gather_f16", flash_attn_gather_f16_len, flash_attn_gather_f16_data, "main", 5, + sizeof(vk_op_flash_attn_gather_push_constants), {1, 1, 1}, {}, 1, true, true, + device->subgroup_size); +#define CREATE_FA_GATHER_TOK_DQ(TYPE, NAMED) \ + ggml_vk_create_pipeline(device, device->pipeline_flash_attn_gather_dq[TYPE], \ + "flash_attn_gather_dq_" #NAMED, flash_attn_gather_dq_ ## NAMED ## _len, \ + flash_attn_gather_dq_ ## NAMED ## _data, "main", 5, \ + sizeof(vk_op_flash_attn_gather_push_constants), {1, 1, 1}, {}, 1, true, true, \ + device->subgroup_size); + CREATE_FA_GATHER_TOK_DQ(GGML_TYPE_Q4_0, q4_0) + CREATE_FA_GATHER_TOK_DQ(GGML_TYPE_Q8_0, q8_0) +#undef CREATE_FA_GATHER_TOK_DQ + if (device->subgroup_arithmetic) { + ggml_vk_create_pipeline(device, device->pipeline_flash_attn_union_f16, + "flash_attn_union_f16", flash_attn_union_f16_len, flash_attn_union_f16_data, "main", 3, + sizeof(vk_op_flash_attn_union_push_constants), {1, 1, 1}, {}, 1, true, true, + device->subgroup_size); + ggml_vk_create_pipeline(device, device->pipeline_flash_attn_gather_union_f16, + "flash_attn_gather_union_f16", flash_attn_gather_union_f16_len, flash_attn_gather_union_f16_data, "main", 6, + sizeof(vk_op_flash_attn_gather_union_push_constants), {1, 1, 1}, {}, 1, true, true, + device->subgroup_size); +#define CREATE_FA_GATHER_DQ(TYPE, NAMED) \ + ggml_vk_create_pipeline(device, device->pipeline_flash_attn_gather_union_dq[TYPE], \ + "flash_attn_gather_union_dq_" #NAMED, flash_attn_gather_union_dq_ ## NAMED ## _len, \ + flash_attn_gather_union_dq_ ## NAMED ## _data, "main", 6, \ + sizeof(vk_op_flash_attn_gather_union_push_constants), {1, 1, 1}, {}, 1, true, true, \ + device->subgroup_size); + CREATE_FA_GATHER_DQ(GGML_TYPE_Q4_0, q4_0) + CREATE_FA_GATHER_DQ(GGML_TYPE_Q8_0, q8_0) +#undef CREATE_FA_GATHER_DQ + } + } + + // DSv4 fused hyper-connection ops: plain f32 compute, no subgroup/coopmat requirements + ggml_vk_create_pipeline(device, device->pipeline_dsv4_hc_pre_f32, + "dsv4_hc_pre_f32", dsv4_hc_pre_f32_len, dsv4_hc_pre_f32_data, "main", 3, + sizeof(vk_op_dsv4_hc_pre_push_constants), {256, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_dsv4_hc_comb_f32, + "dsv4_hc_comb_f32", dsv4_hc_comb_f32_len, dsv4_hc_comb_f32_data, "main", 4, + sizeof(vk_op_dsv4_hc_comb_push_constants), {256, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_dsv4_hc_post_f32, + "dsv4_hc_post_f32", dsv4_hc_post_f32_len, dsv4_hc_post_f32_data, "main", 5, + sizeof(vk_op_dsv4_hc_post_push_constants), {256, 1, 1}, {}, 1); + if (device->subgroup_arithmetic && device->subgroup_require_full_support) { ggml_vk_create_pipeline(device, device->pipeline_ssm_scan_f32_d128, "ssm_scan_128_f32", ssm_scan_subgroup_f32_len, ssm_scan_subgroup_f32_data, "main", 8, sizeof(vk_op_ssm_scan_push_constants), {1, 1, 1}, {128, device->subgroup_size}, 1, true, true); ggml_vk_create_pipeline(device, device->pipeline_ssm_scan_f32_d256, "ssm_scan_256_f32", ssm_scan_subgroup_f32_len, ssm_scan_subgroup_f32_data, "main", 8, sizeof(vk_op_ssm_scan_push_constants), {1, 1, 1}, {256, device->subgroup_size}, 1, true, true); @@ -6722,6 +7295,15 @@ static vk_device ggml_vk_get_device(size_t idx) { device->max_nodes_per_submit = std::max(max_nodes_per_submit, 1u); } + // Also submit once a batch has accumulated enough memory traffic, so that + // bandwidth-bound nodes with no flops estimate cannot grow a command buffer + // past the driver timeout. 0 disables the limit. + device->max_bytes_per_submit = 8ull * 1024 * 1024 * 1024; + const char* GGML_VK_MAX_MB_PER_SUBMIT = getenv("GGML_VK_MAX_MB_PER_SUBMIT"); + if (GGML_VK_MAX_MB_PER_SUBMIT != nullptr) { + device->max_bytes_per_submit = std::stoull(GGML_VK_MAX_MB_PER_SUBMIT) * 1024 * 1024; + } + const bool force_disable_f16 = getenv("GGML_VK_DISABLE_F16") != nullptr; device->fp16 = !force_disable_f16 && fp16_storage && fp16_compute; @@ -7892,6 +8474,11 @@ static void ggml_vk_init(ggml_backend_vk_context * ctx, size_t idx) { ctx->prealloc_size_x = 0; ctx->prealloc_size_y = 0; ctx->prealloc_size_split_k = 0; + ctx->fa_dequant_gate_sz = 0; + ctx->fa_dequant_gate_fits = false; + ctx->fa_dequant_gate_logged = false; + memset(ctx->fa_union_est_ratio, 0, sizeof(ctx->fa_union_est_ratio)); + ctx->fa_union_declines = 0; // Fixed size of 1KB, for deterministic behavior ctx->prealloc_size_add_rms_partials = 1024; @@ -8000,7 +8587,8 @@ static vk_matmul_pipeline ggml_vk_get_mul_mat_mat_pipeline(ggml_backend_vk_conte return pipelines; } - if (src1_type != GGML_TYPE_F32 && !ctx->device->coopmat2) { + if (src1_type != GGML_TYPE_F32 && !ctx->device->coopmat2 && + !(src1_type == GGML_TYPE_F16 && ggml_vk_dense_f16b_enabled())) { return nullptr; } @@ -8044,7 +8632,9 @@ static vk_matmul_pipeline ggml_vk_get_mul_mat_mat_pipeline(ggml_backend_vk_conte return prec == GGML_PREC_DEFAULT ? ctx->device->pipeline_dequant_mul_mat_mat_f16[src0_type].f16acc : ctx->device->pipeline_dequant_mul_mat_mat_f16[src0_type].f32acc; } if (ctx->device->coopmat_support) { - return (ctx->device->fp16 && ctx->device->coopmat_acc_f16_support && prec == GGML_PREC_DEFAULT) ? ctx->device->pipeline_dequant_mul_mat_mat[src0_type].f16acc : ctx->device->pipeline_dequant_mul_mat_mat[src0_type].f32acc; + vk_matmul_pipeline2 & p = (src1_type == GGML_TYPE_F16) ? ctx->device->pipeline_dequant_mul_mat_mat_f16[src0_type] + : ctx->device->pipeline_dequant_mul_mat_mat[src0_type]; + return (ctx->device->fp16 && ctx->device->coopmat_acc_f16_support && prec == GGML_PREC_DEFAULT) ? p.f16acc : p.f32acc; } return (ctx->device->fp16 && prec == GGML_PREC_DEFAULT) ? ctx->device->pipeline_dequant_mul_mat_mat[src0_type].f16acc : ctx->device->pipeline_dequant_mul_mat_mat[src0_type].f32acc; } @@ -8179,7 +8769,8 @@ static vk_matmul_pipeline ggml_vk_get_mul_mat_mat_id_pipeline(ggml_backend_vk_co return pipelines; } - GGML_ASSERT(src1_type == GGML_TYPE_F32 || (ctx->device->coopmat2 && src1_type == GGML_TYPE_F16)); + GGML_ASSERT(src1_type == GGML_TYPE_F32 || + ((ctx->device->coopmat2 || ggml_vk_mmid_f16b_enabled()) && src1_type == GGML_TYPE_F16)); switch (src0_type) { case GGML_TYPE_Q1_0: @@ -8216,7 +8807,11 @@ static vk_matmul_pipeline ggml_vk_get_mul_mat_mat_id_pipeline(ggml_backend_vk_co return nullptr; } - vk_matmul_pipeline2& mmp = ctx->device->pipeline_dequant_mul_mat_mat_id[src0_type]; + // GGML_VK_MMID_F16B: on KHR_coopmat devices the f16-B mmid pipelines live in a + // parallel array (coopmat2's main array already holds f16-B pipelines). + vk_matmul_pipeline2& mmp = (src1_type == GGML_TYPE_F16 && !ctx->device->coopmat2) + ? ctx->device->pipeline_dequant_mul_mat_mat_id_f16b[src0_type] + : ctx->device->pipeline_dequant_mul_mat_mat_id[src0_type]; // XXX TODO 'prec' is not actually allowed in mul_mat_id. bool prefer_fp16acc = ctx->device->fp16 /*&& prec == GGML_PREC_DEFAULT*/; bool support_fp16acc = !mmp.f16acc->is_empty(); @@ -9221,14 +9816,14 @@ static void ggml_vk_matmul_id( uint32_t m, uint32_t n, uint32_t k, uint32_t stride_a, uint32_t stride_b, uint32_t stride_d, uint32_t batch_stride_a, uint32_t batch_stride_b, uint32_t batch_stride_d, uint32_t n_as, uint32_t nei0, uint32_t nei1, uint32_t nbi1, uint32_t ne11, - bool hoist_row_ids) { + bool hoist_row_ids, const vk_subbuffer & fused_scale, uint32_t fusion_flags) { VK_LOG_DEBUG("ggml_vk_matmul_id(a: (" << a.buffer->buffer << ", " << a.offset << ", " << a.size << "), b: (" << b.buffer->buffer << ", " << b.offset << ", " << b.size << "), d: (" << d.buffer->buffer << ", " << d.offset << ", " << d.size << "), ids: (" << ids.buffer->buffer << ", " << ids.offset << ", " << ids.size << "), expert_count: (" << expert_count_buf.buffer->buffer << ", " << expert_count_buf.offset << ", " << expert_count_buf.size << "), " << "m: " << m << ", n: " << n << ", k: " << k << ", stride_a: " << stride_a << ", stride_b: " << stride_b << ", stride_d: " << stride_d << ", " << "batch_stride_a: " << batch_stride_a << ", batch_stride_b: " << batch_stride_b << ", batch_stride_d: " << batch_stride_d << ", " << "n_as: " << n_as << ", nei0: " << nei0 << ", nei1: " << nei1 << ", nbi1: " << nbi1 << ", ne11: " << ne11 << ")"); const vk_mat_mat_id_push_constants pc = { m, n, k, stride_a, stride_b, stride_d, batch_stride_a, batch_stride_b, batch_stride_d, - nei0, nei1, nbi1, ne11, n_as, uint32_t(hoist_row_ids) }; - ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { a, b, d, ids, expert_count_buf }, pc, { m, nei1, n_as }); + nei0, nei1, nbi1, ne11, n_as, uint32_t(hoist_row_ids), fusion_flags }; + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { a, b, d, ids, expert_count_buf, fused_scale }, pc, { m, nei1, n_as }); } static bool ggml_vk_dim01_contiguous(const ggml_tensor * tensor) { @@ -9534,7 +10129,26 @@ static void ggml_vk_mul_mat_q_f16(ggml_backend_vk_context * ctx, vk_context& sub // Reformat and convert to fp16 if non-contiguous, or for coopmat2 for better perf const bool x_non_contig = (ctx->device->coopmat2 && src0->type == GGML_TYPE_F32) || !ggml_vk_dim01_contiguous(src0); - const bool y_non_contig = (ctx->device->coopmat2 && src1->type == GGML_TYPE_F32) || + // Route quantized MUL_MAT through the f16-B kernels. Treating contiguous f32 B as + // y_non_contig reuses the convert-to-prealloc_y plumbing, like coopmat2 does. + // auto mode restricts to ne10 == 5120: the only width measured to gain. Narrow on purpose, + // so widths measured as losses (3584 dense, 2048 MoE) cannot trigger it. + const bool dense_f16b = ggml_vk_dense_f16b_enabled() && + (ggml_vk_dense_f16b_mode() == 1 || ne10 == 5120) && + ctx->device->coopmat_support && !ctx->device->coopmat2 && + ggml_is_quantized(src0->type) && src1->type == GGML_TYPE_F32 && + !(ctx->device->pipeline_dequant_mul_mat_mat_f16[src0->type].f16acc->is_empty() && + ctx->device->pipeline_dequant_mul_mat_mat_f16[src0->type].f32acc->is_empty()); + if (dense_f16b) { + static bool dense_f16b_logged = false; + if (!dense_f16b_logged) { + dense_f16b_logged = true; + fprintf(stderr, "ggml_vulkan: MUL_MAT f16-B path engaged (GGML_VK_DENSE_F16B)\n"); + } + } + + const bool y_non_contig = dense_f16b || + (ctx->device->coopmat2 && src1->type == GGML_TYPE_F32) || (src0->type == GGML_TYPE_BF16 && src1->type != GGML_TYPE_BF16) || !ggml_vk_dim01_contiguous(src1); @@ -10475,7 +11089,7 @@ static void ggml_vk_mul_mat(ggml_backend_vk_context * ctx, vk_context& subctx, c } } -static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, const ggml_tensor * ids, ggml_tensor * dst) { +static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, const ggml_tensor * ids, ggml_tensor * dst, const ggml_tensor * fused_scale = nullptr, ggml_tensor * fused_dst = nullptr) { VK_LOG_DEBUG("ggml_vk_mul_mat_id_q_f16((" << src0 << ", name=" << src0->name << ", type=" << src0->type << ", ne0=" << src0->ne[0] << ", ne1=" << src0->ne[1] << ", ne2=" << src0->ne[2] << ", ne3=" << src0->ne[3] << ", nb0=" << src0->nb[0] << ", nb1=" << src0->nb[1] << ", nb2=" << src0->nb[2] << ", nb3=" << src0->nb[3]; std::cerr << "), (" << src1 << ", name=" << src1->name << ", type=" << src1->type << ", ne0=" << src1->ne[0] << ", ne1=" << src1->ne[1] << ", ne2=" << src1->ne[2] << ", ne3=" << src1->ne[3] << ", nb0=" << src1->nb[0] << ", nb1=" << src1->nb[1] << ", nb2=" << src1->nb[2] << ", nb3=" << src1->nb[3]; std::cerr << "), (" << ids << ", name=" << ids->name << ", type=" << ids->type << ", ne0=" << ids->ne[0] << ", ne1=" << ids->ne[1] << ", ne2=" << ids->ne[2] << ", ne3=" << ids->ne[3] << ", nb0=" << ids->nb[0] << ", nb1=" << ids->nb[1] << ", nb2=" << ids->nb[2] << ", nb3=" << ids->nb[3]; @@ -10513,7 +11127,9 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& hoisted_row_id_words * sizeof(uint32_t) <= ctx->device->properties.limits.maxStorageBufferRange; - ggml_backend_vk_buffer_context * dst_buf_ctx = (ggml_backend_vk_buffer_context *)dst->buffer->context; + // When the following MUL is fused in, write the scaled result straight to its destination. + const ggml_tensor * out_dst = fused_dst ? fused_dst : dst; + ggml_backend_vk_buffer_context * dst_buf_ctx = (ggml_backend_vk_buffer_context *)out_dst->buffer->context; ggml_backend_vk_buffer_context * src0_buf_ctx = (ggml_backend_vk_buffer_context *)src0->buffer->context; ggml_backend_vk_buffer_context * src1_buf_ctx = (ggml_backend_vk_buffer_context *)src1->buffer->context; ggml_backend_vk_buffer_context * ids_buf_ctx = (ggml_backend_vk_buffer_context *)ids->buffer->context; @@ -10559,7 +11175,27 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& #else const bool y_decode_vector_staging = false; #endif + // EXPERIMENT (GGML_VK_MMID_F16B=1): route quantized MUL_MAT_ID through the f16-B + // kernels. Treating contiguous f32 B as y_non_contig reuses the existing + // convert-to-prealloc_y plumbing (to_fp16_vk_1), exactly like coopmat2 does. + // Gated on coopmat_support because the f16b pipelines are only created there. + const bool mmid_f16b = ggml_vk_mmid_f16b_enabled() && + ctx->device->coopmat_support && !ctx->device->coopmat2 && + ggml_is_quantized(src0->type) && src1->type == GGML_TYPE_F32 && + // only take the f16-B path if a pipeline exists for this src0 type + // (e.g. Q2_0 has none); otherwise fall through to the normal f32-B path. + !(ctx->device->pipeline_dequant_mul_mat_mat_id_f16b[src0->type].f16acc->is_empty() && + ctx->device->pipeline_dequant_mul_mat_mat_id_f16b[src0->type].f32acc->is_empty()); + if (mmid_f16b) { + static bool mmid_f16b_logged = false; + if (!mmid_f16b_logged) { + mmid_f16b_logged = true; + fprintf(stderr, "ggml_vulkan: MUL_MAT_ID f16-B path engaged (GGML_VK_MMID_F16B)\n"); + } + } + const bool y_non_contig = y_decode_vector_staging || + mmid_f16b || (ctx->device->coopmat2 && src1->type == GGML_TYPE_F32) || (src0->type == GGML_TYPE_BF16 && src1->type != GGML_TYPE_BF16) || !ggml_vk_dim01_contiguous(src1); @@ -10587,7 +11223,18 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& const ggml_type effective_src1_type = quantize_y ? GGML_TYPE_Q8_1 : (y_f32_kernel ? GGML_TYPE_F32 : src1->type); - const uint32_t kpad = quantize_y ? 0 : ggml_vk_align_size(ne10, ggml_vk_guess_matmul_id_pipeline_align(ctx, mmp, ne01, nei1, qx_needs_dequant ? f16_type : src0->type, effective_src1_type)); + // EXPERIMENT (GGML_VK_MMID_SMALLN=1): select the matmul tile by the EXPECTED PER-EXPERT + // token count rather than the whole batch. With E experts and nei0 active per token, each + // expert sees ~nei1*nei0/E rows; selecting by aggregate nei1 always picks the widest tile + // and leaves most N-lanes empty at MoE prefill (measured: MUL_MAT_ID at ~20% of dense + // matmul efficiency, ~78% of MoE prefill time on Qwen3.6-35B-A3B). + uint32_t n_for_tile = (uint32_t)nei1; + static const char * mmid_smalln_env = getenv("GGML_VK_MMID_SMALLN"); + if (!(mmid_smalln_env && atoi(mmid_smalln_env) == 0) && ne02 > 1) { // on by default; =0 disables + n_for_tile = std::max(1u, (uint32_t)((nei1 * nei0 + ne02 - 1) / ne02)); + } + + const uint32_t kpad = quantize_y ? 0 : ggml_vk_align_size(ne10, ggml_vk_guess_matmul_id_pipeline_align(ctx, mmp, ne01, n_for_tile, qx_needs_dequant ? f16_type : src0->type, effective_src1_type)); // Coopmat2 MUL_MAT_ID BK specialization constants in ggml_vk_load_shaders are at most 64. const uint32_t y_staged_row_stride = ctx->device->coopmat2 && !quantize_y ? ggml_vk_align_size(ne10, 64) : ne10; const bool y_needs_k_padding = ne10 != y_staged_row_stride; @@ -10596,10 +11243,21 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& // Not implemented GGML_ASSERT(y_needs_reformat || !qy_needs_dequant); // NOLINT - const bool aligned = !quantize_y && ne10 == kpad && ne01 > 8 && nei1 > 8; - vk_pipeline pipeline = ggml_vk_guess_matmul_id_pipeline(ctx, mmp, ne01, nei1, aligned, qx_needs_dequant ? f16_type : src0->type, effective_src1_type); + vk_pipeline pipeline = ggml_vk_guess_matmul_id_pipeline(ctx, mmp, ne01, n_for_tile, aligned, qx_needs_dequant ? f16_type : src0->type, effective_src1_type); + + // PROBE (GGML_VK_MMID_PROBE=1): which mmid tile actually runs, and with how many threads. + static const char * mmid_probe_env = getenv("GGML_VK_MMID_PROBE"); + if (mmid_probe_env && atoi(mmid_probe_env) != 0) { + static std::set seen; + std::string key = pipeline->name + ":" + std::to_string(n_for_tile); + if (seen.insert(key).second) { + fprintf(stderr, "ggml_vulkan: mmid pipeline=%s n_for_tile=%u m=%u wg=(%u,%u,%u)\n", + pipeline->name.c_str(), n_for_tile, (uint32_t)ne01, + pipeline->wg_denoms[0], pipeline->wg_denoms[1], pipeline->wg_denoms[2]); + } + } if (ggml_nbytes(src0) > ctx->device->properties.limits.maxStorageBufferRange) { pipeline = ggml_vk_get_64b_indexing_pipeline(ctx, pipeline); @@ -10691,7 +11349,7 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& } vk_buffer d_D = dst_buf_ctx->dev_buffer; - const uint64_t d_buf_offset = vk_tensor_offset(dst) + dst->view_offs; + const uint64_t d_buf_offset = vk_tensor_offset(out_dst) + out_dst->view_offs; GGML_ASSERT(d_D != nullptr); vk_buffer d_X; uint64_t x_buf_offset = 0; @@ -10827,7 +11485,9 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& { d_D, d_buf_offset, d_sz }, { d_ids, ids_buf_offset, ids_sz }, expert_count_buf, ne01, ne21, ne10, ne10, stride_b_y, ne01, stride_batch_x, stride_batch_y, ne20*ne21, - n_as, nei0, nei1, nbi1 / ggml_type_size(ids->type), ne11, hoist_row_ids + n_as, nei0, nei1, nbi1 / ggml_type_size(ids->type), ne11, hoist_row_ids, + fused_scale ? ggml_vk_tensor_subbuffer(ctx, fused_scale) : vk_subbuffer{ d_D, d_buf_offset, d_sz }, + fused_scale ? 1u : 0u ); // NOLINT if (x_non_contig || qx_needs_dequant) { @@ -11091,7 +11751,16 @@ static void ggml_vk_mul_mat_id(ggml_backend_vk_context * ctx, vk_context& subctx if (ggml_vk_use_mul_mat_vec_id(cgraph, node_idx)) { ggml_vk_mul_mat_vec_id_q_f16(ctx, subctx, cgraph, node_idx); } else { - ggml_vk_mul_mat_id_q_f16(ctx, subctx, src0, src1, src2, dst); + // Fused scale epilogue: the MUL's other operand is applied as the matmul writes out, + // and the result goes straight to the MUL's destination. + const ggml_tensor * fused_scale = nullptr; + ggml_tensor * fused_dst = nullptr; + if (ctx->num_additional_fused_ops == 1) { + ggml_tensor * mul = cgraph->nodes[node_idx + 1]; + fused_scale = (mul->src[0] == dst) ? mul->src[1] : mul->src[0]; + fused_dst = mul; + } + ggml_vk_mul_mat_id_q_f16(ctx, subctx, src0, src1, src2, dst, fused_scale, fused_dst); } } @@ -11169,8 +11838,8 @@ static bool ggml_vk_flash_attn_coopmat_shmem_support(const vk_device& device, co const uint32_t qstride = hsk_pad / 4 + 2; const uint32_t Qf = Br * qstride * f16vec4; - const uint32_t psh_stride = Br / 4 + 2; - const uint32_t Psh = Bc * psh_stride * f16vec4; + const uint32_t psh_stride = Bc / 4 + 2; + const uint32_t Psh = Br * psh_stride * f16vec4; const uint32_t sfshstride = (hsk <= 128) ? (Br + 8) : Br; const uint32_t sfsh = Bc * sfshstride * acctype; @@ -11194,6 +11863,800 @@ static bool ggml_vk_flash_attn_coopmat_shmem_support(const vk_device& device, co return supported; } +// K/V types the FA shaders can read directly (scalar/coopmat1 select the dequant code via the +// FaTypeK/FaTypeV spec constants; a type outside this list silently reads garbage). Types that +// are FA-supported but not listed here (iq4_nl) are only correct through the dequant-once +// scratch path, so supports_op and the dispatch-time gate must agree on when that path runs. +static bool ggml_vk_fa_kv_native(ggml_type t, bool coopmat2) { + switch (t) { + case GGML_TYPE_F32: + case GGML_TYPE_F16: + case GGML_TYPE_BF16: + case GGML_TYPE_Q4_0: + case GGML_TYPE_Q4_1: + case GGML_TYPE_Q5_0: + case GGML_TYPE_Q5_1: + case GGML_TYPE_Q8_0: + case GGML_TYPE_IQ4_NL: // native FA support since upstream 8161641 + return true; + default: + return false; + } +} + +// Capacity gate for the dequant-once f16 K/V scratch. On discrete devices the scratch can push the +// working set past free VRAM, and the driver then silently pages device-local memory (~15x prefill +// regression measured on an 8 GB card at long context). UMA has no separate pool to overflow. +// +// Gates our own device-local usage against the physical heap size, less a reserve for memory not +// visible in-process. heapBudget is deliberately not the signal: ggml_backend_vk_get_device_memory +// computes heapBudget - heapUsage unsigned, which wraps when the device is oversubscribed. Nor can +// the allocation gate itself: on WDDM vkAllocateMemory only fails near physical heap size, above +// the free-VRAM level where paging starts. The reserve is necessarily conservative because other +// processes' VRAM use is invisible to us. GGML_VK_FA_DEQUANT=0/1 forces the path off/on; +// GGML_VK_FA_DEQUANT_RESERVE_MB overrides the reserve. +static bool ggml_vk_fa_dequant_scratch_fits(ggml_backend_vk_context * ctx, uint64_t scratch_sz) { + const vk_device& device = ctx->device; + + if (device->uma) { + return true; + } + + // Decided once per scratch size, so the decision cannot flip between layers as usage grows. + if (ctx->fa_dequant_gate_sz == scratch_sz) { + return ctx->fa_dequant_gate_fits; + } + + static const uint64_t reserve = [] { + const char * env = getenv("GGML_VK_FA_DEQUANT_RESERVE_MB"); + return (uint64_t)(env ? atoi(env) : 1024) * 1024 * 1024; + }(); + + bool fits = false; + + // Without VK_EXT_memory_budget our usage is unknowable, so leave the path disabled. + if (vk_instance.device_supports_membudget[device->idx]) { + vk::PhysicalDeviceMemoryBudgetPropertiesEXT budgetprops; + vk::PhysicalDeviceMemoryProperties2 memprops = {}; + memprops.pNext = &budgetprops; + device->physical_device.getMemoryProperties2(&memprops); + + uint64_t heap_size = 0; + uint64_t heap_used = 0; + for (uint32_t i = 0; i < memprops.memoryProperties.memoryHeapCount; ++i) { + const vk::MemoryHeap & heap = memprops.memoryProperties.memoryHeaps[i]; + if (heap.flags & vk::MemoryHeapFlagBits::eDeviceLocal) { + heap_size += heap.size; + heap_used += budgetprops.heapUsage[i]; + } + } + // heap_used already covers scratch allocated on a previous ubatch, so counting scratch_sz + // in full is conservative by up to the current scratch size. + fits = heap_size > reserve && heap_used + scratch_sz + reserve <= heap_size; + } + + if (!fits && !ctx->fa_dequant_gate_logged) { + ctx->fa_dequant_gate_logged = true; + GGML_LOG_INFO("ggml_vulkan: flash attention dequant-once disabled: %llu MiB K/V scratch does not fit " + "device-local memory with a %llu MiB reserve. Set GGML_VK_FA_DEQUANT=1 to force it on.\n", + (unsigned long long)(scratch_sz >> 20), (unsigned long long)(reserve >> 20)); + } + + ctx->fa_dequant_gate_sz = scratch_sz; + ctx->fa_dequant_gate_fits = fits; + return fits; +} + +// Close a timestamp interval mid-node so work a handler dispatches before its op is billed +// separately instead of being folded into the op's own time. `name` must be a string literal +// (stored by pointer). No-op unless the perf logger is on in per-op mode. +static void ggml_vk_perf_mark_subop(ggml_backend_vk_context * ctx, vk_context& subctx, const char * name) { + if (!vk_perf_logger_enabled || vk_perf_logger_concurrent || ctx->query_pool == VK_NULL_HANDLE) { + return; + } + if (ctx->query_idx >= (int)ctx->num_queries) { + return; // pool headroom exhausted; drop the mark rather than overflow + } + ctx->query_nodes[ctx->query_idx] = nullptr; + ctx->query_fusion_names[ctx->query_idx] = nullptr; + ctx->query_sub_names[ctx->query_idx] = name; + subctx->s->buffer->buf.writeTimestamp(vk::PipelineStageFlagBits::eAllCommands, ctx->query_pool, ctx->query_idx++); +} + +static bool ggml_vk_flash_attn_top_k(ggml_backend_vk_context * ctx, vk_context & subctx, + const ggml_tensor * q, const ggml_tensor * k, const ggml_tensor * v, + const ggml_tensor * mask, const ggml_tensor * sinks, ggml_tensor * dst) { + const ggml_tensor * top_k = dst->src[5]; + static const char * top_k_env = getenv("GGML_VK_FA_TOPK"); + // The sparse shaders read f16 only. Serve a quantised cache by dequantising it once into + // the f16 scratch, the same pass the dense path uses. K and V are the same tensor here, so + // one pass covers both. Without this a quantised cache declines to dense FA, which costs + // O(kv) where the sparse path costs O(n_kv_raw + n_top_k). + static const char * fa_dequant_env = getenv("GGML_VK_FA_DEQUANT"); + const uint64_t kv_f16_sz = (uint64_t) ggml_nelements(k) * sizeof(ggml_fp16_t); + const bool dequant_kv = top_k && k->type != GGML_TYPE_F16 && + !(fa_dequant_env && fa_dequant_env[0] == '0') && + ctx->device->pipeline_dequant_transpose[k->type] != nullptr && + k->nb[0] == ggml_type_size(k->type) && + ggml_is_contiguously_allocated(k) && + kv_f16_sz <= ctx->device->properties.limits.maxStorageBufferRange && + ggml_vk_fa_dequant_scratch_fits(ctx, kv_f16_sz); + if ((top_k_env && top_k_env[0] == '0') || + !top_k || (!ctx->device->pipeline_flash_attn_top_k_f16 && !ctx->device->pipeline_flash_attn_top_k_cm_f16) || + q->type != GGML_TYPE_F32 || (k->type != GGML_TYPE_F16 && !dequant_kv) || v->type != k->type || + !mask || mask->type != GGML_TYPE_F16 || top_k->type != GGML_TYPE_I32 || + q->ne[0] != 512 || q->ne[1] < 64 || k->ne[0] != 512 || v->ne[0] != 512 || + q->ne[2] != 64 || k->ne[2] != 1 || v->ne[2] != 1 || + q->ne[1] != top_k->ne[1] || q->ne[3] != top_k->ne[3] || + k->ne[1] != v->ne[1] || k->buffer != v->buffer || k->data != v->data || + !ggml_is_contiguous(mask) || !ggml_is_contiguous(top_k)) { + return false; + } + + float scale = 0.0f; + float max_bias = 0.0f; + float logit_softcap = 0.0f; + memcpy(&scale, (const float *) dst->op_params + 0, sizeof(float)); + memcpy(&max_bias, (const float *) dst->op_params + 1, sizeof(float)); + memcpy(&logit_softcap, (const float *) dst->op_params + 2, sizeof(float)); + if (max_bias != 0.0f || logit_softcap != 0.0f) { + return false; + } + + const int32_t n_kv_raw = ggml_get_op_params_i32(dst, 4); + if (n_kv_raw < 0 || n_kv_raw > k->ne[1] || top_k->ne[0] > k->ne[1] - n_kv_raw) { + return false; + } + const int64_t n_kv_active = n_kv_raw + top_k->ne[0]; + // Crossover against ordinary dense FA. The coopmat sparse path (and especially the + // raw/selected split below) is cheap enough to win as soon as any key is pruned; the + // scalar fallback shader is not, and regresses against dense FA until the pruned + // fraction is large, so it keeps its original 3x margin. Equality always uses dense + // FA because nothing is pruned. + const bool have_cm_sparse = ctx->device->pipeline_flash_attn_top_k_cm_f16 != nullptr; + const int64_t min_total_k = have_cm_sparse ? n_kv_active + 1 : 3 * n_kv_active; + if (k->ne[1] < min_total_k) { + return false; + } + + // ---- diagnostic: adjacent-token top-k overlap (GGML_VK_TOPK_OVERLAP=1) --------------- + // Decides whether a deduplicated union is worth building for the small-batch path. A + // speculative draft attends at adjacent POSITIONS, so overlap between adjacent prefill + // tokens is the same quantity and a plain deep prefill samples it thousands of times + // without needing a draft model. Reports |union| / (W * n_top_k) for window sizes W. + // Sampled, host-side, off by default; the readback would be far too costly otherwise. + { + static const char * ov_env = getenv("GGML_VK_TOPK_OVERLAP"); + if (ov_env && ov_env[0] == '1') { + static std::mutex ov_mu; + static uint64_t ov_seen = 0; + static std::map> ov_acc; // W -> {sum ratio, n} + std::lock_guard lock(ov_mu); + if ((ov_seen++ % 32) == 0) { // sample: this is a multi-MB readback + const uint32_t nb = (uint32_t) q->ne[1]; + const uint32_t tk = (uint32_t) top_k->ne[0]; + const int32_t rng = (int32_t) (k->ne[1] - n_kv_raw); + std::vector idx((size_t) nb * tk); + vk_subbuffer sb = ggml_vk_tensor_subbuffer(ctx, top_k); + ggml_vk_buffer_read(sb.buffer, sb.offset, idx.data(), idx.size() * sizeof(int32_t)); + for (uint32_t W : {2u, 4u, 8u}) { + if (nb < W) continue; + double sum = 0.0; uint64_t n = 0; + for (uint32_t t0 = 0; t0 + W <= nb; t0 += W) { // disjoint windows + std::unordered_set u; + uint64_t valid = 0; + for (uint32_t t = t0; t < t0 + W; ++t) { + for (uint32_t j = 0; j < tk; ++j) { + const int32_t v = idx[(size_t) t * tk + j]; + if (v >= 0 && v < rng) { u.insert(v); ++valid; } + } + } + if (valid) { sum += (double) u.size() / (double) valid; ++n; } + } + if (n) { auto & a = ov_acc[W]; a.first += sum; a.second += n; } + } + fprintf(stderr, "[topk-overlap] sample %llu (n_batch=%u, kv=%lld)\n", + (unsigned long long) ov_seen, nb, (long long) k->ne[1]); + for (const auto & e : ov_acc) { + const double r = e.second.first / (double) e.second.second; + const double now = (double) n_kv_raw + (double) e.first * tk; + const double dedup = (double) n_kv_raw + r * (double) e.first * tk; + fprintf(stderr, "[topk-overlap] W=%u union/selected=%.3f (overlap %.1f%%) " + "kv_c %.0f -> %.0f projected small-batch gain %.1f%%\n", + e.first, r, 100.0 * (1.0 - r), now, dedup, 100.0 * (1.0 - dedup / now)); + } + } + } + } + + vk_subbuffer k_buf = ggml_vk_tensor_subbuffer(ctx, k); + if (dequant_kv) { + if (ctx->prealloc_size_x < kv_f16_sz) { + ctx->prealloc_size_x = kv_f16_sz; + ggml_vk_preallocate_buffers(ctx, subctx); + } + vk_pipeline tr_k = ctx->device->pipeline_dequant_transpose[k->type]; + ggml_pipeline_request_descriptor_sets(ctx, tr_k, 1); + if (ctx->prealloc_x_need_sync) { + ggml_vk_sync_buffers(ctx, subctx); + } + const vk_subbuffer k_dst = vk_subbuffer{ ctx->prealloc_x, 0, kv_f16_sz }; + const uint32_t k_nel = (uint32_t) ggml_nelements(k); + const std::vector tr_pc = { (uint32_t) k->ne[0], (uint32_t) k->ne[2], (uint32_t) k->ne[1], 0, k_nel }; + ggml_vk_dispatch_pipeline(ctx, subctx, tr_k, { k_buf, k_dst }, tr_pc, { k_nel, 1, 1 }); + ggml_vk_sync_buffers(ctx, subctx); + ctx->prealloc_x_need_sync = true; + k_buf = k_dst; + ggml_vk_perf_mark_subop(ctx, subctx, "FA_KV_DEQUANT (sub-op)"); + } + // strides of what the shaders actually read: the source cache, or the contiguous + // [HS, KV, n_head_kv, ns] f16 scratch the dequant just wrote + const uint32_t k_stride = dequant_kv ? (uint32_t) k->ne[0] : (uint32_t) (k->nb[1] / sizeof(ggml_fp16_t)); + const uint32_t k_nb3_el = dequant_kv ? (uint32_t) ((uint64_t) k->ne[0] * k->ne[1] * k->ne[2]) + : (uint32_t) (k->nb[3] / sizeof(ggml_fp16_t)); + const uint32_t k_nb2_byte = dequant_kv ? (uint32_t) ((uint64_t) k->ne[0] * k->ne[1] * sizeof(ggml_fp16_t)) + : (uint32_t) k->nb[2]; + const uint32_t k_nb3_byte = dequant_kv ? (uint32_t) (k_nb3_el * sizeof(ggml_fp16_t)) : (uint32_t) k->nb[3]; + + vk_op_flash_attn_top_k_push_constants pc = { + (uint32_t) q->ne[1], (uint32_t) k->ne[1], (uint32_t) n_kv_raw, + (uint32_t) top_k->ne[0], (uint32_t) q->ne[2], + (uint32_t) (q->nb[1] / sizeof(float)), + (uint32_t) (q->nb[2] / sizeof(float)), + (uint32_t) (q->nb[3] / sizeof(float)), + k_stride, + k_nb3_el, + (uint32_t) (mask->nb[1] / sizeof(ggml_fp16_t)), + (uint32_t) (mask->nb[3] / sizeof(ggml_fp16_t)), + (uint32_t) (top_k->nb[1] / sizeof(int32_t)), + (uint32_t) (top_k->nb[3] / sizeof(int32_t)), + (uint32_t) (dst->nb[1] / sizeof(float)), + (uint32_t) (dst->nb[2] / sizeof(float)), + (uint32_t) (dst->nb[3] / sizeof(float)), + scale, sinks != nullptr, 0, + }; + + const vk_subbuffer q_buf = ggml_vk_tensor_subbuffer(ctx, q); + const vk_subbuffer sinks_buf = sinks ? ggml_vk_tensor_subbuffer(ctx, sinks) : q_buf; + static const char * top_k_cm_env = getenv("GGML_VK_FA_TOPK_CM"); + const bool use_cm = (!top_k_cm_env || top_k_cm_env[0] != '0') && ctx->device->pipeline_flash_attn_top_k_cm_f16; + vk_pipeline pipeline = use_cm ? ctx->device->pipeline_flash_attn_top_k_cm_f16 : ctx->device->pipeline_flash_attn_top_k_f16; + + static const char * top_k_split_env = getenv("GGML_VK_FA_TOPK_SPLIT"); + const uint32_t mask_stride = (uint32_t) (mask->nb[1] / sizeof(ggml_fp16_t)); + const bool try_split = use_cm && (!top_k_split_env || top_k_split_env[0] != '0') && n_kv_raw > 0 && top_k->ne[0] > 0; + if (try_split) { + const uint32_t N = (uint32_t) q->ne[1]; + const uint32_t D = 512; + const uint32_t NH = 64; + const uint32_t NS = (uint32_t) q->ne[3]; + const uint32_t raw_kv = (uint32_t) n_kv_raw; + const uint32_t partitions = 2; + // Query tiling caps the split scratch at 256 queries regardless of batch or stream + // count. The reduce takes the tile height as ne2 and the full query count as dst_ne2, + // so the destination stream stride stays correct across tiles. + const uint32_t tile_size = std::min(N, 256u); + const uint32_t n_tiles = CEIL_DIV(N, tile_size); + const bool f32acc = true; + vk_fa_tuning_params tuning = get_fa_tuning_params(ctx->device, D, D, N, raw_kv, GGML_TYPE_F16, GGML_TYPE_F16, f32acc); + + const uint32_t q_stride = (uint32_t) (q->nb[1] / sizeof(float)); + const bool aligned = raw_kv % tuning.block_cols == 0 && (q_stride & 7) == 0 && (k_stride & 7) == 0; + const vk_fa_pipeline_state raw_state = get_fa_pipeline_state(ctx->device, tuning, D, D, aligned, f32acc, + true, false, false, GGML_TYPE_F16, GGML_TYPE_F16); + if (raw_state.path == FA_COOPMAT1 && ctx->device->pipeline_flash_attn_split_k_reduce) { + vk_pipeline raw_pipeline; + { + std::lock_guard guard(ctx->device->compile_mutex); + auto & pipelines = ctx->device->pipeline_flash_attn_f32_f16; + auto it = pipelines.find(raw_state); + if (it != pipelines.end()) { + raw_pipeline = it->second; + } else { + pipelines[raw_state] = raw_pipeline = std::make_shared(); + } + } + const uint64_t partition_size = ((uint64_t) D * NH + NH * 2) * sizeof(float) * tile_size * NS; + const uint64_t split_size = partition_size * partitions; + if (split_size <= ctx->device->properties.limits.maxStorageBufferRange) { + ggml_pipeline_request_descriptor_sets(ctx, raw_pipeline, n_tiles); + ggml_pipeline_request_descriptor_sets(ctx, pipeline, n_tiles); + ggml_pipeline_request_descriptor_sets(ctx, ctx->device->pipeline_flash_attn_split_k_reduce, n_tiles); + if (ctx->prealloc_size_split_k < split_size) { + ctx->prealloc_size_split_k = split_size; + ggml_vk_preallocate_buffers(ctx, subctx); + } + if (ctx->prealloc_split_k_need_sync) { + ggml_vk_sync_buffers(ctx, subctx); + } + + // ALiBi is disabled (max_bias == 0), so n_head_log2 is never read; keep it 0 + // rather than a value that looks computed. + const uint32_t n_head_log2 = 0; + // The raw dispatch smuggles the mask row stride through split_kv, which the FA + // shader also uses to derive its KV range as min(KV, (split_k_index+1)*split_kv). + // That is only safe while split_k_index == 0 and the stride covers the whole raw + // prefix -- both hold here, but assert rather than rely on it silently. + GGML_ASSERT(mask_stride >= raw_kv && "split_kv carries the mask stride; it must not clip the raw KV range"); + const uint32_t mask_stride_in_split_kv = 1u << 31; + const uint32_t packed_gqa = mask_stride_in_split_kv | 1u; + const uint32_t packed_partitions = (partitions << 16) | 1; + const vk_subbuffer split_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_split_k, 0); + const vk_subbuffer mask_buf = ggml_vk_tensor_subbuffer(ctx, mask); + const vk_subbuffer top_buf = ggml_vk_tensor_subbuffer(ctx, top_k); + const vk_subbuffer dst_buf = ggml_vk_tensor_subbuffer(ctx, dst); + const auto sliced = [](const vk_subbuffer & buf, uint64_t offset) { + return vk_subbuffer{buf.buffer, buf.offset + offset, buf.size - offset}; + }; + + for (uint32_t tile = 0; tile < n_tiles; ++tile) { + if (tile != 0) { + ggml_vk_sync_buffers(ctx, subctx); + } + const uint32_t token_offset = tile * tile_size; + const uint32_t tile_n = std::min(tile_size, N - token_offset); + const vk_subbuffer tile_q = sliced(q_buf, (uint64_t) token_offset * q->nb[1]); + const vk_subbuffer tile_mask = sliced(mask_buf, (uint64_t) token_offset * mask->nb[1]); + const vk_subbuffer tile_top = sliced(top_buf, (uint64_t) token_offset * top_k->nb[1]); + const vk_subbuffer tile_dst = sliced(dst_buf, (uint64_t) token_offset * dst->nb[2]); + const vk_flash_attn_push_constants raw_pc = { + tile_n, raw_kv, + NH, tile_n, NS, + NH, NS, + 1, NS, + 1, NS, + (uint32_t) mask->ne[1], (uint32_t) mask->ne[2], (uint32_t) mask->ne[3], + q_stride, (uint32_t) q->nb[2], (uint32_t) q->nb[3], + k_stride, k_nb2_byte, k_nb3_byte, + k_stride, k_nb2_byte, k_nb3_byte, + scale, 0.0f, 0.0f, + n_head_log2, 1.0f, 1.0f, + packed_gqa, mask_stride, packed_partitions, + }; + + ggml_vk_dispatch_pipeline(ctx, subctx, raw_pipeline, + {tile_q, k_buf, k_buf, tile_mask, tile_q, split_buf, tile_q, tile_q /* dyn-KV: unused */}, + raw_pc, {tile_n, NH, NS}); + ggml_vk_perf_mark_subop(ctx, subctx, "FA_TOP_K_RAW (sub-op)"); + + pc.n_batch = tile_n; + pc.split_mode = 1; + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, + {tile_q, k_buf, tile_mask, sinks_buf, tile_top, split_buf}, + pc, {tile_n, (uint32_t) CEIL_DIV(q->ne[2], 32), NS}); + ggml_vk_perf_mark_subop(ctx, subctx, "FA_TOP_K_SELECTED (sub-op)"); + + ctx->prealloc_split_k_need_sync = true; + ggml_vk_sync_buffers(ctx, subctx); + const vk_op_flash_attn_split_k_reduce_push_constants reduce_pc = {D, NH, tile_n, N, NS, partitions, sinks != nullptr}; + ggml_vk_dispatch_pipeline(ctx, subctx, ctx->device->pipeline_flash_attn_split_k_reduce, + {split_buf, sinks_buf, tile_dst}, reduce_pc, {NH, D, tile_n * NS}); + ctx->prealloc_split_k_need_sync = true; + ggml_vk_perf_mark_subop(ctx, subctx, "FA_TOP_K_REDUCE (sub-op)"); + } + return true; + } + } + } + + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, + {q_buf, k_buf, ggml_vk_tensor_subbuffer(ctx, mask), sinks_buf, + ggml_vk_tensor_subbuffer(ctx, top_k), ggml_vk_tensor_subbuffer(ctx, dst)}, + pc, {(uint32_t) q->ne[1], (uint32_t) CEIL_DIV(q->ne[2], use_cm ? 32 : 8), (uint32_t) q->ne[3]}); + ggml_vk_perf_mark_subop(ctx, subctx, use_cm ? "FA_TOP_K_CM (sub-op)" : "FA_TOP_K_SPARSE (sub-op)"); + return true; +} + +struct vk_fa_compact_state { + bool active = false; + bool dynamic_kv = false; // KV row count lives in kv_buf, not the push constant + uint32_t kv_c = 0; // upper bound; the real count is runtime when dynamic_kv + uint32_t n_batch = 1; + uint32_t row_bytes = 0; // bytes per compact K row; K may be quantised + uint32_t row_elems = 0; // K row stride in ELEMENTS/blocks, for the FA push constant + bool dequantized = false; // scratch holds f16 because the gather decoded on the way in + vk_subbuffer kv_buf; + vk_subbuffer kc_buf, mc_buf; +}; + +// Small host-visible buffer holding the last union count the device produced: +// [0] padded compact rows (also read by the gather and by the FA), [1] raw union size, +// [2] the candidate count it came from, [3] the batch it came from. Host-visible is a +// requirement rather than a preference here - the whole point is that the host can read it +// without submitting. +static bool ggml_vk_fa_union_stat_init(ggml_backend_vk_context * ctx) { + if (ctx->fa_union_stat) { + return ctx->fa_union_stat->ptr != nullptr; + } + try { + ctx->fa_union_stat = ggml_vk_create_buffer(ctx->device, 64, + {vk::MemoryPropertyFlagBits::eDeviceLocal | vk::MemoryPropertyFlagBits::eHostVisible | vk::MemoryPropertyFlagBits::eHostCoherent, + vk::MemoryPropertyFlagBits::eHostVisible | vk::MemoryPropertyFlagBits::eHostCoherent}); + } catch (const vk::SystemError &) { + return false; + } + if (ctx->fa_union_stat->ptr == nullptr) { + return false; + } + memset(ctx->fa_union_stat->ptr, 0, 64); + return true; +} + +// Host-coherent memory still needs the write made available to the host domain. This is the +// only barrier in the path that names eHost, and it is cheap: nothing waits on it, it just +// lets the next graph's host-side read see the count this one produced. +static void ggml_vk_fa_union_stat_host_barrier(vk_context & subctx) { + subctx->s->buffer->buf.pipelineBarrier( + vk::PipelineStageFlagBits::eComputeShader, + vk::PipelineStageFlagBits::eHost, + {}, + { { vk::AccessFlagBits::eShaderWrite, vk::AccessFlagBits::eHostRead } }, + {}, {}); +} + +// Price the union without synchronising for it. The count read here belongs to whichever call +// last wrote the slot, which under a decode is the final flash-attention op of the previous +// graph - one token stale, and that is fine: selection overlap is a property of the model and +// the draft, not of an individual op, and a wrong estimate costs part of one step rather than +// correctness (the compact buffers are still sized for the worst case). +// +// Tracks the latest measurement, lightly smoothed. The asymmetry runs the other way from what +// a conservative estimator would assume: an estimate that is too HIGH declines compaction and +// forgoes 2-3x for as long as it stays high, while one that is too low costs a single step at +// roughly dense cost and is corrected by the count that step produces. An earlier version held +// a decaying peak instead and measured 1992 us where the union delivers 900, because a spell of +// genuinely low overlap pinned the estimate and 0.999 per read took hundreds of steps to relax. +// +// The words are read without ordering against the device write, so they can come from +// different calls. The sample is filed under the batch the device reported rather than the +// batch being priced, so a torn read costs one mispriced step for that batch and then +// corrects, which is the same failure the estimate already tolerates. +// +// Returns the padded compact row count to gate on, or 0 when this batch is unseeded. +static uint32_t ggml_vk_fa_union_estimate(ggml_backend_vk_context * ctx, uint32_t n_kv_raw, + uint32_t n_batch, uint32_t n_cand) { + const volatile uint32_t * stat = (const volatile uint32_t *) ctx->fa_union_stat->ptr; + const uint32_t u = stat[1]; + const uint32_t cand = stat[2]; + const uint32_t nb_obs = stat[3]; + + if (u > 0 && cand > 0 && nb_obs > 0 && nb_obs < 64) { + const float r = std::min(1.0f, (float) u / (float) cand); + float & e = ctx->fa_union_est_ratio[nb_obs]; + e = e > 0.0f ? 0.5f * r + 0.5f * e : r; + } + + const float ratio = ctx->fa_union_est_ratio[n_batch]; + if (ratio <= 0.0f) { + return 0; + } + return GGML_PAD(n_kv_raw + (uint32_t) ceilf(ratio * (float) n_cand), 256u); +} + +// V4 sparse decode (gather-to-compact): the sparse prefill shader above gates on +// q->ne[1] >= 64, so single-token decode otherwise attends densely over the whole +// compressed KV, at a cost that grows with context. Instead, gather the active rows +// (dense prefix + top-k selection; MQA, so all query heads share one set) into a +// compact contiguous scratch in prealloc_y, and let the ordinary dense FA below run +// over the compacted K/V/mask. Correct by the same contract as the sparse shader: +// the source mask carries the selection, and the gathered mask preserves it. +static bool ggml_vk_flash_attn_gather_compact(ggml_backend_vk_context * ctx, vk_context & subctx, + const ggml_tensor * q, const ggml_tensor * k, const ggml_tensor * v, + const ggml_tensor * mask, ggml_tensor * dst, vk_fa_compact_state & st) { + const ggml_tensor * top_k = dst->src[5]; + static const char * gather_env = getenv("GGML_VK_FA_TOPK_GATHER"); + // K/V type: the gather relocates rows verbatim and never reads a value out of them, so it + // does not have to be f16 - it has to be a type whose row is a whole number of 4-byte words, + // and one flash-attention has a NATIVE shader for. The dequant/contiguize pass is disabled + // while this scratch is active and asserts if it turns out to be needed, so admitting a + // non-native type here would abort rather than fall back. + const bool kv_word_addressable = + k->type == v->type && + ggml_vk_fa_kv_native(k->type, ctx->device->coopmat2) && + k->ne[0] % ggml_blck_size(k->type) == 0 && + ggml_row_size(k->type, k->ne[0]) % 4 == 0 && + k->nb[1] % 4 == 0 && k->nb[3] % 4 == 0; + + if ((gather_env && gather_env[0] == '0') || + !top_k || !ctx->device->pipeline_flash_attn_gather_f16 || + q->ne[1] < 1 || q->ne[1] >= 64 || // 1..63: >=64 goes to the sparse prefill path + q->type != GGML_TYPE_F32 || !kv_word_addressable || + !mask || mask->type != GGML_TYPE_F16 || top_k->type != GGML_TYPE_I32 || + q->ne[0] != 512 || k->ne[0] != 512 || v->ne[0] != 512 || q->ne[2] != 64 || + k->ne[2] != 1 || v->ne[2] != 1 || + q->ne[1] != top_k->ne[1] || q->ne[3] != top_k->ne[3] || + k->ne[1] != v->ne[1] || k->buffer != v->buffer || k->data != v->data || + !ggml_is_contiguous(mask) || !ggml_is_contiguous(top_k)) { + return false; + } + + const uint32_t k_row_bytes = (uint32_t) ggml_row_size(k->type, k->ne[0]); + const uint32_t k_row_words = k_row_bytes / 4; + + float max_bias = 0.0f; + float logit_softcap = 0.0f; + memcpy(&max_bias, (const float *) dst->op_params + 1, sizeof(float)); + memcpy(&logit_softcap, (const float *) dst->op_params + 2, sizeof(float)); + if (max_bias != 0.0f || logit_softcap != 0.0f) { + return false; + } + + const int32_t n_kv_raw = ggml_get_op_params_i32(dst, 4); + if (n_kv_raw < 0 || n_kv_raw > k->ne[1] || top_k->ne[0] > k->ne[1] - n_kv_raw) { + return false; + } + + const uint32_t n_batch = (uint32_t) q->ne[1]; + const uint32_t n_cand = (uint32_t) ((int64_t) n_batch * top_k->ne[0]); + // Worst case: every token gets its own top-k block, so the compact set is + // n_kv_raw + n_batch*n_top_k -- independent of context depth, which is the point, but + // GROWING WITH BATCH. Attention work is then n_batch * (n_kv_raw + n_batch*n_top_k), i.e. + // quadratic in batch, against dense's n_batch * n_kv. Break-even is + // n_batch = (n_kv - n_kv_raw) / n_top_k, and a kv >= 2*kv_c gate on this worst case caps + // the useful batch at (n_kv/2 - n_kv_raw) / n_top_k -- about 6 at 32k depth, ~30 at 128k. + // Beyond that the gate declines and dense runs, so this can never be slower; it just stops + // helping. The union below lifts that ceiling by gating on what the selections actually + // deduplicate to; kv_c remains the bound for every allocation and dispatch count. + const uint32_t kv_c = GGML_PAD((uint32_t) (n_kv_raw + (int64_t) n_cand), 256u); + + // ---- deduplicated union (default on, GGML_VK_FA_TOPK_UNION=0 disables) -------------- + // Same compact layout, but one row per DISTINCT selected key instead of one block per + // token. Measured on real draft tokens the union is 0.56 of the selections at batch 4, + // so it is materially smaller. Its size is only known on the GPU, so the FA + // reads its KV bound from a buffer (the DYNAMIC_KV pipeline flag) rather than a push + // constant; padding the count to 256 keeps KV % Bc == 0 so the aligned variant still + // applies. Single stream only: the FA takes one KV for all streams. + // + // Gated on the ESTIMATED union rather than on kv_c, which is why this sits above the + // worst-case gate: at batch 8 and 32k depth the worst case is 6400 rows against 11008 + // source rows and would decline, while the union measures around 3300 and is well worth + // compacting. kv_c stays the worst case for every allocation and dispatch bound, so a + // wrong estimate is a slow step, never a wrong answer. + const uint32_t max_words = 12288; // shared bitmap capacity in flash_attn_union.comp + const bool bitmap_fits = (uint64_t) ((k->ne[1] - n_kv_raw) + 31) / 32 <= max_words; + + static const char * union_env = getenv("GGML_VK_FA_TOPK_UNION"); + if ((!union_env || union_env[0] != '0') && q->ne[3] == 1 && n_batch > 1 && bitmap_fits && + ctx->device->pipeline_flash_attn_union_f16 && ctx->device->pipeline_flash_attn_gather_union_f16 && + ggml_vk_fa_union_stat_init(ctx)) { + const uint32_t max_union = n_cand; + const uint32_t kv_c_est = ggml_vk_fa_union_estimate(ctx, (uint32_t) n_kv_raw, n_batch, n_cand); + // Two separate questions. Does the compact set fit under the gate at all, and does + // deduplicating actually shrink it: with no overlap to exploit the union is the same + // size as the per-token blocks and the scan is pure cost, measured at 1.2% of the op + // at 512k depth. The worst-case bound on the source keeps a collapse in overlap to + // roughly dense cost for the one step it takes the estimate to catch up. + const bool worth_it = kv_c_est != 0 && kv_c_est < kv_c && + (uint64_t) k->ne[1] >= 2ull * kv_c_est && + (uint64_t) k->ne[1] >= (uint64_t) kv_c; + + // GGML_VK_FA_UNION_STATS=N: report every Nth call (N=1 means every call) what the gate + // decided and on what measurement. The alternative is inferring engagement from a + // timing, which is how a sparse path gets credited for a run it never took. + // + // The period is a parameter because a fixed one is a way to miss the answer: a 64-token + // speculative decode over 25k of context produced ONE line at a period of 256, which + // said only that the first call was unseeded. Running totals rather than instants, so a + // single late line still reports whether the path engaged. + static const char * stats_env = getenv("GGML_VK_FA_UNION_STATS"); + if (stats_env && stats_env[0] != '\0' && stats_env[0] != '0') { + static uint64_t calls = 0, taken = 0; + const uint64_t period = std::max(1ull, (unsigned long long) atoll(stats_env)); + taken += worth_it ? 1 : 0; + if ((calls++ % period) == 0) { + fprintf(stderr, "[fa-union] n_kv=%lld n_kv_raw=%d n_batch=%u cand=%u " + "union/cand=%.3f kv_c %u -> est %u %s (%llu/%llu taken)\n", + (long long) k->ne[1], n_kv_raw, n_batch, n_cand, (double) ctx->fa_union_est_ratio[n_batch], + kv_c, kv_c_est, worth_it ? "UNION" : "declined", + (unsigned long long) taken, (unsigned long long) calls); + } + } + + if (!worth_it) { + // Nothing is known about this batch shape, or the last count says compaction does + // not pay. Either way the count is what settles it, so produce one: the scan is a + // single workgroup and depth-independent, and count_only needs no index list. Dense + // runs this step; the next one decides on a measurement instead of a bound. + // + // Every declined step, not a sample: a decline is exactly the state in which the + // estimate stops being refreshed by the compact path, so sampling it leaves a stale + // estimate latched for as many steps as the sampling period. The dispatch is one + // workgroup against the ~2.2 ms dense op it is riding along with. + ctx->fa_union_declines++; + { + const vk_op_flash_attn_union_push_constants ppc = { + (uint32_t) k->ne[1], (uint32_t) n_kv_raw, n_batch, (uint32_t) top_k->ne[0], max_union, + (uint32_t) (top_k->nb[1] / sizeof(int32_t)), max_words, 256u, 1u, + }; + const vk_subbuffer stat_buf = ggml_vk_subbuffer(ctx, ctx->fa_union_stat); + ggml_pipeline_request_descriptor_sets(ctx, ctx->device->pipeline_flash_attn_union_f16, 1); + // The probe writes the same slot a taken union path reads back as its row + // count, and a graph can contain both: the estimate is refreshed from the + // device as the graph is recorded, so an op late in the graph can be admitted + // after an earlier one was declined. Order it explicitly rather than rely on + // the compact path's own sync, which this path does not go through. + ggml_vk_sync_buffers(ctx, subctx); + ggml_vk_dispatch_pipeline(ctx, subctx, ctx->device->pipeline_flash_attn_union_f16, + { ggml_vk_tensor_subbuffer(ctx, top_k), stat_buf, stat_buf }, ppc, { 1, 1, 1 }); + ggml_vk_fa_union_stat_host_barrier(subctx); + } + goto union_unavailable; + } + + // Decoding on the way in makes the scratch f16 and hands flash attention its f16 path, + // which is worth far more than the extra scratch bytes: the inline decode it replaces + // costs a measured 0.15 us per KV row attended, every step. + const bool dq = ggml_is_quantized(k->type) && ctx->device->pipeline_flash_attn_gather_union_dq[k->type]; + const uint32_t u_row_by = dq ? (uint32_t) (k->ne[0] * sizeof(ggml_fp16_t)) : k_row_bytes; + const size_t ukc_sz = (size_t) kv_c * u_row_by; + const size_t umc_sz = (size_t) n_batch * kv_c * sizeof(ggml_fp16_t); + const size_t ul_sz = (size_t) max_union * sizeof(uint32_t); + const size_t need = ukc_sz + umc_sz + ul_sz; + if (ctx->prealloc_size_y < need) { + ctx->prealloc_size_y = need; + ggml_vk_preallocate_buffers(ctx, subctx); + } + // Unconditional, not gated on prealloc_y_need_sync: this also orders the count slot + // against a probe dispatched earlier in the same graph, which leaves that flag clear. + ggml_vk_sync_buffers(ctx, subctx); + const vk_subbuffer kc_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_y, 0); + const vk_subbuffer mc_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_y, ukc_sz); + const vk_subbuffer ul_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_y, ukc_sz + umc_sz); + // The count lives in the stat buffer rather than in prealloc_y so that the same write + // that the gather and the FA consume is also the one the host prices the next step from. + // Both are single-slot and rewritten by every layer, so the WAR hazard is unchanged: the + // prealloc_y sync above is a global barrier and orders the previous layer's read. + const vk_subbuffer uc_buf = ggml_vk_subbuffer(ctx, ctx->fa_union_stat); + + vk_pipeline gather_pipe = dq ? ctx->device->pipeline_flash_attn_gather_union_dq[k->type] + : ctx->device->pipeline_flash_attn_gather_union_f16; + ggml_pipeline_request_descriptor_sets(ctx, ctx->device->pipeline_flash_attn_union_f16, 1); + ggml_pipeline_request_descriptor_sets(ctx, gather_pipe, 1); + + const vk_op_flash_attn_union_push_constants upc = { + (uint32_t) k->ne[1], (uint32_t) n_kv_raw, n_batch, (uint32_t) top_k->ne[0], max_union, + (uint32_t) (top_k->nb[1] / sizeof(int32_t)), max_words, 256u, 0u, + }; + ggml_vk_dispatch_pipeline(ctx, subctx, ctx->device->pipeline_flash_attn_union_f16, + { ggml_vk_tensor_subbuffer(ctx, top_k), ul_buf, uc_buf }, upc, { 1, 1, 1 }); + ggml_vk_sync_buffers(ctx, subctx); + + // the fused decoder steps K in BLOCKS and writes elements; the verbatim one does words + const vk_op_flash_attn_gather_union_push_constants gpc = { + (uint32_t) k->ne[1], (uint32_t) n_kv_raw, kv_c, + dq ? (uint32_t) (k->nb[1] / ggml_type_size(k->type)) : (uint32_t) (k->nb[1] / 4), + (uint32_t) (mask->nb[1] / sizeof(ggml_fp16_t)), + n_batch, dq ? (uint32_t) k->ne[0] : k_row_words, + }; + ggml_vk_dispatch_pipeline(ctx, subctx, gather_pipe, + { ggml_vk_tensor_subbuffer(ctx, k), ul_buf, ggml_vk_tensor_subbuffer(ctx, mask), + kc_buf, mc_buf, uc_buf }, gpc, { kv_c, 1, 1 }); + ggml_vk_sync_buffers(ctx, subctx); + ggml_vk_fa_union_stat_host_barrier(subctx); + ctx->prealloc_y_need_sync = true; + + st.active = true; + st.dynamic_kv = true; + st.kv_c = kv_c; + st.n_batch = n_batch; + st.row_bytes = u_row_by; + st.row_elems = dq ? (uint32_t) k->ne[0] : (uint32_t) (k->ne[0] / ggml_blck_size(k->type)); + st.dequantized = dq; + st.kc_buf = kc_buf; + st.mc_buf = mc_buf; + st.kv_buf = uc_buf; + return true; + } +union_unavailable:; + + // Per-token blocks have no dedup, so this form really does cost kv_c: the gather writes + // then re-reads ~the active bytes while dense reads the source KV once, so compaction only + // pays when the source is comfortably larger than the active set. + if ((uint64_t) k->ne[1] < 2ull * kv_c) { + return false; + } + + // Decode on the way in, as the union path does: the gather touches each row once, while + // flash attention decodes once per query block that reads it. The union only covers + // n_batch > 1, so single-token decode lands here. + const bool tok_dq = ggml_is_quantized(k->type) && ctx->device->pipeline_flash_attn_gather_dq[k->type]; + const uint32_t row_by = tok_dq ? (uint32_t) (k->ne[0] * sizeof(ggml_fp16_t)) : k_row_bytes; + + const uint32_t ns = (uint32_t) q->ne[3]; + const size_t kc_sz = (size_t) ns * kv_c * row_by; + const size_t mc_sz = (size_t) ns * n_batch * kv_c * sizeof(ggml_fp16_t); + + if (ctx->prealloc_size_y < kc_sz + mc_sz) { + ctx->prealloc_size_y = kc_sz + mc_sz; + ggml_vk_preallocate_buffers(ctx, subctx); + } + if (ctx->prealloc_y_need_sync) { + ggml_vk_sync_buffers(ctx, subctx); + } + + vk_pipeline pipeline = tok_dq ? ctx->device->pipeline_flash_attn_gather_dq[k->type] + : ctx->device->pipeline_flash_attn_gather_f16; + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); + + // the fused decoder steps K in BLOCKS and writes elements; the verbatim one does words + const vk_op_flash_attn_gather_push_constants pc = { + (uint32_t) k->ne[1], (uint32_t) n_kv_raw, (uint32_t) top_k->ne[0], kv_c, + tok_dq ? (uint32_t) (k->nb[1] / ggml_type_size(k->type)) : (uint32_t) (k->nb[1] / 4), + tok_dq ? (uint32_t) (k->nb[3] / ggml_type_size(k->type)) : (uint32_t) (k->nb[3] / 4), + (uint32_t) (top_k->nb[1] / sizeof(int32_t)), + (uint32_t) (top_k->nb[3] / sizeof(int32_t)), + (uint32_t) (mask->nb[1] / sizeof(ggml_fp16_t)), + (uint32_t) (mask->nb[3] / sizeof(ggml_fp16_t)), + (uint32_t) mask->ne[3], + n_batch, tok_dq ? (uint32_t) k->ne[0] : k_row_words, + }; + + st.kc_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_y, 0); + st.mc_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_y, kc_sz); + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, + { ggml_vk_tensor_subbuffer(ctx, k), ggml_vk_tensor_subbuffer(ctx, top_k), + ggml_vk_tensor_subbuffer(ctx, mask), st.kc_buf, st.mc_buf }, + pc, { kv_c, 1, ns }); + ggml_vk_sync_buffers(ctx, subctx); + ctx->prealloc_y_need_sync = true; + + // ---- diagnostic: measure real top-k overlap between the query tokens ---------------- + // GGML_VK_TOPK_OVERLAP=1. Off by default and never touched on the hot path. Answers the + // one question that decides whether a deduplicated union is worth building: this path + // gathers n_batch*n_top_k rows, a union would gather |union| rows, and the op's cost is + // linear in that count (measured: 204.4 / 205.7 / 205.5 us per row at batch 2 / 4 / 8). + // Needs a real model - synthetic top-k selections say nothing about real overlap, which + // is exactly how the mul_mat_id large-tile probe misled once before. + static const char * overlap_env = getenv("GGML_VK_TOPK_OVERLAP"); + if (overlap_env && overlap_env[0] == '1' && n_batch > 1) { + static std::mutex ov_mutex; + static uint64_t ov_calls = 0, ov_selected = 0, ov_union = 0; + static std::map> ov_by_batch; // n_batch -> {selected, union} + const size_t n_idx = (size_t) n_batch * top_k->ne[0]; + std::vector idx(n_idx); + vk_subbuffer top_sb = ggml_vk_tensor_subbuffer(ctx, top_k); + ggml_vk_buffer_read(top_sb.buffer, top_sb.offset, idx.data(), n_idx * sizeof(int32_t)); + + const int32_t range = (int32_t) (k->ne[1] - n_kv_raw); + std::unordered_set uni; + uint64_t valid = 0; + for (size_t i = 0; i < n_idx; ++i) { + const int32_t v = idx[i]; + if (v >= 0 && v < range) { uni.insert(v); ++valid; } + } + std::lock_guard lock(ov_mutex); + ov_calls++; ov_selected += valid; ov_union += uni.size(); + auto & e = ov_by_batch[n_batch]; + e.first += valid; e.second += uni.size(); + if ((ov_calls % 256) == 0) { + fprintf(stderr, "[topk-overlap] calls=%llu selected=%llu union=%llu " + "union/selected=%.3f => a dedup union would gather %.1f%% fewer compressed rows\n", + (unsigned long long) ov_calls, (unsigned long long) ov_selected, + (unsigned long long) ov_union, + ov_selected ? (double) ov_union / (double) ov_selected : 0.0, + ov_selected ? 100.0 * (1.0 - (double) ov_union / (double) ov_selected) : 0.0); + for (const auto & kv : ov_by_batch) { + const double ratio = kv.second.first ? (double) kv.second.second / (double) kv.second.first : 0.0; + // projected op speedup uses the measured linear cost model on kv_c + const double now = (double) n_kv_raw + (double) kv.first * (double) top_k->ne[0]; + const double dedup = (double) n_kv_raw + ratio * (double) kv.first * (double) top_k->ne[0]; + fprintf(stderr, "[topk-overlap] n_batch=%u union/selected=%.3f " + "kv_c %.0f -> %.0f projected op gain %.1f%%\n", + kv.first, ratio, now, dedup, 100.0 * (1.0 - dedup / now)); + } + } + } + + st.active = true; + st.kv_c = kv_c; + st.n_batch = n_batch; + st.row_bytes = row_by; + st.row_elems = tok_dq ? (uint32_t) k->ne[0] : (uint32_t) (k->ne[0] / ggml_blck_size(k->type)); + st.dequantized = tok_dq; + return true; +} + static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * q, const ggml_tensor * k, const ggml_tensor * v, const ggml_tensor * mask, const ggml_tensor * sinks, ggml_tensor * dst) { VK_LOG_DEBUG("ggml_vk_flash_attn((" << q << ", name=" << q->name << ", type=" << q->type << ", ne0=" << q->ne[0] << ", ne1=" << q->ne[1] << ", ne2=" << q->ne[2] << ", ne3=" << q->ne[3] << ", nb0=" << q->nb[0] << ", nb1=" << q->nb[1] << ", nb2=" << q->nb[2] << ", nb3=" << q->nb[3]; std::cerr << "), (" << k << ", name=" << k->name << ", type=" << k->type << ", ne0=" << k->ne[0] << ", ne1=" << k->ne[1] << ", ne2=" << k->ne[2] << ", ne3=" << k->ne[3] << ", nb0=" << k->nb[0] << ", nb1=" << k->nb[1] << ", nb2=" << k->nb[2] << ", nb3=" << k->nb[3]; @@ -11213,15 +12676,15 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx GGML_TENSOR_LOCALS(int64_t, ne, dst, ne) GGML_TENSOR_LOCALS(size_t, nb, dst, nb) - const uint32_t nem0 = mask ? mask->ne[0] : 0; - const uint32_t nem1 = mask ? mask->ne[1] : 0; - const uint32_t nem2 = mask ? mask->ne[2] : 0; - const uint32_t nem3 = mask ? mask->ne[3] : 0; + uint32_t nem0 = mask ? mask->ne[0] : 0; + uint32_t nem1 = mask ? mask->ne[1] : 0; + uint32_t nem2 = mask ? mask->ne[2] : 0; + uint32_t nem3 = mask ? mask->ne[3] : 0; const uint32_t HSK = nek0; const uint32_t HSV = nev0; uint32_t N = neq1; - const uint32_t KV = nek1; + uint32_t KV = nek1; GGML_ASSERT(ne0 == HSV); GGML_ASSERT(ne2 == N); @@ -11245,6 +12708,23 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx assert(dst->type == GGML_TYPE_F32); assert(q->type == GGML_TYPE_F32); + if (ggml_vk_flash_attn_top_k(ctx, subctx, q, k, v, mask, sinks, dst)) { + return; + } + // V4 sparse decode: gather the active set into a compact scratch and run the dense + // FA below on it. Overrides KV, the mask geometry, and (further down) the K/V/mask + // bindings and strides; every other decision then sizes itself to the compact KV. + vk_fa_compact_state fa_compact; + if (ggml_vk_flash_attn_gather_compact(ctx, subctx, q, k, v, mask, dst, fa_compact)) { + KV = fa_compact.kv_c; + nem0 = fa_compact.kv_c; + nem1 = fa_compact.n_batch; + nem2 = 1; + nem3 = (uint32_t) q->ne[3]; + nem1 = N; + nem2 = 1; + nem3 = (uint32_t) q->ne[3]; + } uint32_t gqa_ratio = 1; uint32_t qk_ratio = neq2 / nek2; uint32_t workgroups_x = (uint32_t)neq1; @@ -11253,7 +12733,14 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx const bool f32acc = !ctx->device->fp16 || dst->op_params[3] == GGML_PREC_F32 || k->type == GGML_TYPE_BF16; - // dequant K/V once into an f16 scratch, reordered KV layout so FA can read without a stride + // For prefill with quantized K/V, dequantize+transpose K/V once into a per-head-contiguous + // f16 scratch and run the f16 FA path, instead of the coopmat1 shader re-dequantizing the + // whole KV inside every Q-workgroup. The KV-cache view reaching FA is [0,2,1,3]-permuted but + // dense, so we require dense allocation (not ggml_is_contiguous) and block-contiguous dim0, + // and only engage where a fused dequant-transpose shader exists. Prefill only + // (n_rows >= 64); measured neutral at shallow depth and up to ~2x at long context. + // The scratch is bound as a single storage buffer holding K and V back to back, so the SUM of + // the two must fit maxStorageBufferRange, not each half independently. auto is_dense_kv_cache = [](const ggml_tensor * t) { return t->nb[0] == ggml_type_size(t->type) && t->nb[2] == ggml_row_size(t->type, t->ne[0]) && @@ -11262,19 +12749,48 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx }; const bool k_quant = k->type != GGML_TYPE_F16 && k->type != GGML_TYPE_BF16 && k->type != GGML_TYPE_F32; const bool v_quant = v->type != GGML_TYPE_F16 && v->type != GGML_TYPE_BF16 && v->type != GGML_TYPE_F32; - const bool use_dequant_kv = k_quant && v_quant && neq1 >= 64 && + const uint64_t kv_f16_sz = ((uint64_t)ggml_nelements(k) + (uint64_t)ggml_nelements(v)) * sizeof(ggml_fp16_t); + static const char * fa_dequant_env = getenv("GGML_VK_FA_DEQUANT"); + const bool fa_dequant_off = fa_dequant_env && fa_dequant_env[0] == '0'; + const bool fa_dequant_on = fa_dequant_env && fa_dequant_env[0] == '1'; + // Contiguize pass for f16 K/V (GGML_VK_FA_KV_CONTIG=0 opts out). The KV-cache view + // reaching FA is head-interleaved ([HS, NH, KV] physically), and the cm1 + // direct-from-global coopMatLoads run ~2-5x slower on those strided rows than on + // per-head-contiguous K/V. Copy K/V once into the scratch instead (dequant_f16_transpose + // is a pure strided copy). Only engages when the rows are actually strided, and shares + // the quant path's prefill/allocation/scratch-capacity gates below. + static const char * fa_kv_contig_env = getenv("GGML_VK_FA_KV_CONTIG"); + const bool fa_kv_contig = !(fa_kv_contig_env && fa_kv_contig_env[0] == '0'); + const bool kv_f16_strided = k->type == GGML_TYPE_F16 && v->type == GGML_TYPE_F16 && + (k->nb[1] != (uint64_t)HSK * sizeof(ggml_fp16_t) || + v->nb[1] != (uint64_t)HSV * sizeof(ggml_fp16_t)) && + (HSK % 8) == 0 && (HSV % 8) == 0; + // A K/V type the FA shaders cannot read directly (iq4_nl) is only correct through the + // dequant path. supports_op only admits such types when the hard conditions below hold, + // and the VRAM heuristic must not veto them (there is no fallback), so force the path on. + const bool kv_needs_dequant = !ggml_vk_fa_kv_native(k->type, ctx->device->coopmat2) || + !ggml_vk_fa_kv_native(v->type, ctx->device->coopmat2); + const bool use_dequant_kv = !fa_dequant_off && + // the gather-to-compact scratch is already contiguous f16; the + // dequant/contiguize pass must not run on top of it + !fa_compact.active && + ((k_quant && v_quant) || kv_needs_dequant || (fa_kv_contig && kv_f16_strided)) && neq1 >= 64 && is_dense_kv_cache(k) && is_dense_kv_cache(v) && - (uint64_t)ggml_nelements(k) * sizeof(ggml_fp16_t) <= ctx->device->properties.limits.maxStorageBufferRange && - (uint64_t)ggml_nelements(v) * sizeof(ggml_fp16_t) <= ctx->device->properties.limits.maxStorageBufferRange && + kv_f16_sz <= ctx->device->properties.limits.maxStorageBufferRange && ctx->device->pipeline_dequant_transpose[k->type] != nullptr && ctx->device->pipeline_dequant_transpose[v->type] != nullptr && - // coopmat2 path does not benefit from the f16 scratch - !ctx->device->coopmat2 && + // coopmat2 reads its native types directly; non-native still needs the scratch + (kv_needs_dequant || !ctx->device->coopmat2) && // Intel Xe1 regresses, see PR 25494 - (ctx->device->vendor_id != VK_VENDOR_ID_INTEL || - (ctx->device->coopmat_support && ctx->device->architecture != vk_device_architecture::INTEL_XE1)); - const ggml_type k_type_eff = use_dequant_kv ? GGML_TYPE_F16 : k->type; - const ggml_type v_type_eff = use_dequant_kv ? GGML_TYPE_F16 : v->type; + (kv_needs_dequant || + ctx->device->vendor_id != VK_VENDOR_ID_INTEL || + (ctx->device->coopmat_support && ctx->device->architecture != vk_device_architecture::INTEL_XE1)) && + (fa_dequant_on || kv_needs_dequant || ggml_vk_fa_dequant_scratch_fits(ctx, kv_f16_sz)); + // If this fires, supports_op admitted a non-native K/V type the gate then rejected; the + // native shader would return garbage rather than fail, so abort instead. + GGML_ASSERT(use_dequant_kv || !kv_needs_dequant); + const ggml_type k_type_eff = (use_dequant_kv || fa_compact.dequantized) ? GGML_TYPE_F16 : k->type; + const ggml_type v_type_eff = (use_dequant_kv || fa_compact.dequantized) ? GGML_TYPE_F16 : v->type; // For scalar/coopmat1 FA, we can use the "large" size to accommodate qga. // For coopmat2 FA, we always use the small size (which is still pretty large for gqa). @@ -11296,6 +12812,12 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx const uint32_t q_stride = (uint32_t)(nbq1 / ggml_type_size(q->type)); uint32_t k_stride = (uint32_t)(nbk1 / ggml_type_size(k->type)); uint32_t v_stride = (uint32_t)(nbv1 / ggml_type_size(v->type)); + if (fa_compact.active) { + // rows are tightly packed in the compact scratch; for a quantised K this is the block + // count per row, which is what nbk1 / ggml_type_size would have given for the source + k_stride = fa_compact.row_elems; + v_stride = fa_compact.row_elems; + } // For F32, the shader treats it as a block of size 4 (for vec4 loads) if (k->type == GGML_TYPE_F32) { @@ -11342,7 +12864,8 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx bool use_mask_opt = mask && nem1 >= 32 && nem0 * nem1 > 32768 && nem0 >= tuning_params.block_cols * 16 && (ctx->device->architecture != vk_device_architecture::AMD_GCN || HSK > 256 || HSV > 256); vk_fa_pipeline_state fa_pipeline_state = get_fa_pipeline_state(ctx->device, tuning_params, HSK, HSV, aligned, f32acc, - mask != nullptr, use_mask_opt, logit_softcap != 0, k_type_eff, v_type_eff); + mask != nullptr, use_mask_opt, logit_softcap != 0, k_type_eff, v_type_eff, + fa_compact.dynamic_kv); vk_pipeline pipeline = nullptr; @@ -11443,6 +12966,11 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx vk_subbuffer v_buf = ggml_vk_tensor_subbuffer(ctx, v); vk_subbuffer dst_buf = ggml_vk_tensor_subbuffer(ctx, dst); vk_subbuffer mask_buf = mask ? ggml_vk_tensor_subbuffer(ctx, mask) : q_buf; + if (fa_compact.active) { + k_buf = fa_compact.kc_buf; + v_buf = fa_compact.kc_buf; // V is the K latent; one gather serves both + mask_buf = fa_compact.mc_buf; + } vk_subbuffer sinks_buf = sinks ? ggml_vk_tensor_subbuffer(ctx, sinks) : q_buf; vk_subbuffer mask_opt_buf = use_mask_opt ? ggml_vk_subbuffer(ctx, ctx->prealloc_y, 0) : q_buf; @@ -11472,6 +13000,11 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx ggml_vk_sync_buffers(ctx, subctx); k_buf = k_dst; v_buf = v_dst; + // Bill the K/V contiguize/dequant pass on its own line. Without this it is charged to + // FLASH_ATTN_EXT, which makes the copy invisible and the kernel look slower than it is. + ggml_vk_perf_mark_subop(ctx, subctx, kv_needs_dequant || (k_quant && v_quant) + ? "FA_KV_DEQUANT (sub-op)" + : "FA_KV_CONTIGUIZE (sub-op)"); } uint32_t mask_n_head_log2 = ((sinks != nullptr) << 24) | n_head_log2; @@ -11496,6 +13029,12 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx ggml_vk_sync_buffers(ctx, subctx); } + // compact scratch layout: [512, kv_c, 1, ns] tightly packed, in K's own type + const uint32_t eff_nbk2 = fa_compact.active ? fa_compact.kv_c * fa_compact.row_bytes : nbk2_eff; + const uint32_t eff_nbk3 = fa_compact.active ? fa_compact.kv_c * fa_compact.row_bytes : nbk3_eff; + const uint32_t eff_nbv2 = fa_compact.active ? eff_nbk2 : nbv2_eff; + const uint32_t eff_nbv3 = fa_compact.active ? eff_nbk3 : nbv3_eff; + const vk_flash_attn_push_constants pc = { N, KV, (uint32_t)ne1, (uint32_t)ne2, (uint32_t)ne3, (uint32_t)neq2, (uint32_t)neq3, @@ -11503,8 +13042,8 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx (uint32_t)nev2, (uint32_t)nev3, nem1, nem2, nem3, q_stride, (uint32_t)nbq2, (uint32_t)nbq3, - k_stride, nbk2_eff, nbk3_eff, - v_stride, nbv2_eff, nbv3_eff, + k_stride, eff_nbk2, eff_nbk3, + v_stride, eff_nbv2, eff_nbv3, scale, max_bias, logit_softcap, mask_n_head_log2, m0, m1, gqa_ratio, split_kv, split_k }; @@ -11528,11 +13067,11 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx vk_subbuffer split_k_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_split_k, 0); ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, - {q_buf, k_buf, v_buf, mask_buf, sinks_buf, split_k_buf, mask_opt_buf}, + {q_buf, k_buf, v_buf, mask_buf, sinks_buf, split_k_buf, mask_opt_buf, fa_compact.dynamic_kv ? fa_compact.kv_buf : q_buf}, pc, { dispatch_x, workgroups_y, workgroups_z }); ggml_vk_sync_buffers(ctx, subctx); - const vk_op_flash_attn_split_k_reduce_push_constants pc2 = { HSV, (uint32_t)ne1, (uint32_t)ne2, (uint32_t)ne3, split_k, (sinks != nullptr) }; + const vk_op_flash_attn_split_k_reduce_push_constants pc2 = { HSV, (uint32_t)ne1, (uint32_t)ne2, (uint32_t)ne2, (uint32_t)ne3, split_k, (sinks != nullptr) }; ggml_vk_dispatch_pipeline(ctx, subctx, ctx->device->pipeline_flash_attn_split_k_reduce, {split_k_buf, sinks_buf, dst_buf}, pc2, { (uint32_t)ne1, HSV, (uint32_t)(ne2 * ne3) }); @@ -11543,7 +13082,7 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx workgroups_x *= pipeline->wg_denoms[0]; } ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, - {q_buf, k_buf, v_buf, mask_buf, sinks_buf, dst_buf, mask_opt_buf}, + {q_buf, k_buf, v_buf, mask_buf, sinks_buf, dst_buf, mask_opt_buf, fa_compact.dynamic_kv ? fa_compact.kv_buf : q_buf}, pc, { workgroups_x, workgroups_y, workgroups_z }); } @@ -11584,6 +13123,33 @@ static vk_conv_shapes ggml_vk_conv_select_shape(ggml_backend_vk_context * ctx, u } } +// The delta-net conv-state path does ggml_transpose() straight into a dim-0 ggml_concat(), so +// the generic concat kernel walks src1 with a conv_channels * 4 byte stride. On qwen35 that is +// 40960 B = 160 * 256 B and 160 % 16 == 0, so every read lands on one of the 16 memory channels: +// 13.7 GB/s against 138.9 GB/s for the tiled path. Route that exact shape to a tiled-transpose +// kernel. On by default; GGML_VK_CONCAT_TRANSPOSE=0 opts out. +// Qwen3.8-27B pp2048: +0.4% at ub 256, +4.7% at ub 1024, +7.2% at ub 2048. +static bool ggml_vk_concat_is_transposed(const ggml_tensor * src0, const ggml_tensor * src1, const ggml_tensor * dst) { + static const char * env = getenv("GGML_VK_CONCAT_TRANSPOSE"); + if (env && env[0] == '0') { + return false; + } + if (ggml_get_op_params_i32(dst, 0) != 0) { // dim 0 only + return false; + } + if (src0->ne[2] != 1 || src0->ne[3] != 1 || src1->ne[2] != 1 || src1->ne[3] != 1) { + return false; + } + const size_t ts = ggml_type_size(src0->type); + if (src0->nb[0] != ts || dst->nb[0] != ts) { // src0 and dst rows must be contiguous + return false; + } + if (src1->nb[0] <= src1->nb[1]) { // src1 must actually be transposed + return false; + } + return src0->ne[1] == src1->ne[1] && dst->ne[1] == src1->ne[1]; +} + static vk_pipeline ggml_vk_op_get_pipeline(ggml_backend_vk_context * ctx, const ggml_tensor * src0, const ggml_tensor * src1, const ggml_tensor * src2, const ggml_tensor * dst, ggml_op op) { switch (op) { case GGML_OP_GET_ROWS: @@ -11676,6 +13242,11 @@ static vk_pipeline ggml_vk_op_get_pipeline(ggml_backend_vk_context * ctx, const if (!ggml_vk_concat_supported(src0, src1, dst)) { return nullptr; } + // Tiled-transpose path handles unquantized 4-byte elements only. + if (!ggml_is_quantized(src0->type) && ggml_vk_concat_unit_size(src0->type) == 4 && + ggml_vk_concat_is_transposed(src0, src1, dst)) { + return ctx->device->pipeline_concat_transpose_i32; + } switch (ggml_vk_concat_unit_size(src0->type)) { case 1: return ctx->device->pipeline_concat_i8; @@ -12114,6 +13685,52 @@ static vk_pipeline ggml_vk_op_get_pipeline(ggml_backend_vk_context * ctx, const } return nullptr; case GGML_OP_LIGHTNING_INDEXER: + // fork fast path: f16 K on wave64 subgroup-arithmetic devices routes to the tuned + // scalar-64/CM kernels (head counts in LI_NH_VALUES); anything else falls through to + // the generic pipeline table + if (src0->type == GGML_TYPE_F32 && src1->type == GGML_TYPE_F16 && dst->type == GGML_TYPE_F32 && + src0->ne[0] == 128 && src1->ne[1] == 1 && + ggml_vk_li_nh_index(src0->ne[1]) >= 0 && ctx->device->pipeline_lightning_indexer_f16[0]) { + const int nhi = ggml_vk_li_nh_index(src0->ne[1]); + // Small-batch routing. Three arms, one env var each, most specific first: + // + // default decode CM for the whole 2-15 window + // ..._DECODE_CM_BATCH=0 small CM for 4-15, scalar for 2-3 (PR #6) + // ..._DECODE_CM_BATCH=0 SMALL_CM=0 scalar for 2-15 (pre-PR #6 baseline) + // + // The decode CM shader puts 16 HEADS in the coopmat N dimension and dispatches one + // workgroup per token, so it issues 4*n_batch tiles per 16 KV rows - the arithmetic + // minimum, since 64 heads fill its 16 columns exactly. The small CM shader puts 16 + // TOKENS in N and pays a flat 64 tiles no matter how small the batch. Routing the + // whole window to decode CM is 2.7x at batch 5 and 4.8x at batch 2-3 (526k source + // tokens, gfx1151); the shader body has no n_batch == 1 assumption, token is + // gl_WorkGroupID.y, so the old ne[2] == 1 gate was an artefact. + // + // That measurement only exists because of pepuscz's PR #6 and issue #10: their + // per-kernel table (43.8 us scalar vs 27.9 us small CM per 1k scanned rows at batch + // 5, the same ratio at every depth) is what showed the cost is a per-tile constant, + // which is what makes the tile count the thing to minimise. Their small CM route is + // kept as the opt-out arm rather than deleted: the tile-count argument is hardware + // independent, but decode CM re-reads the K tile once per token, and that part is + // bandwidth dependent, so the crossover need not sit here on other devices. + static const char * decode_cm_batch_env = getenv("GGML_VK_LIGHTNING_INDEXER_DECODE_CM_BATCH"); + static const bool decode_cm_batch_on = !decode_cm_batch_env || decode_cm_batch_env[0] != '0'; + const int64_t decode_cm_max = decode_cm_batch_on ? 15 : 1; + if (ctx->device->pipeline_lightning_indexer_decode_cm_f16[nhi] && src0->ne[2] <= decode_cm_max) { + return ctx->device->pipeline_lightning_indexer_decode_cm_f16[nhi]; + } + // PR #6 as its author proposed it, now default-on so that the kill switch above is + // enough on its own to reach it. + static const char * small_cm_env = getenv("GGML_VK_LIGHTNING_INDEXER_SMALL_CM"); + static const bool small_cm_on = !small_cm_env || small_cm_env[0] != '0'; + if (ctx->device->pipeline_lightning_indexer_cm_small_f16[nhi] && small_cm_on && + src0->ne[2] >= 4 && src0->ne[2] < 16) { + return ctx->device->pipeline_lightning_indexer_cm_small_f16[nhi]; + } + vk_pipeline cm = ctx->device->pipeline_lightning_indexer_cm_f16[nhi] ? + ctx->device->pipeline_lightning_indexer_cm_f16[nhi] : ctx->device->pipeline_lightning_indexer_cm_small_f16[nhi]; + return cm && src0->ne[2] >= 16 ? cm : ctx->device->pipeline_lightning_indexer_f16[nhi]; + } // only the k type selects a pipeline, the other types are fixed by ggml_lightning_indexer() if (ggml_vk_lightning_indexer_k_type_supported(src1->type)) { return ctx->device->pipeline_lightning_indexer_f32[src1->type]; @@ -12702,6 +14319,11 @@ static void ggml_vk_op_f32(ggml_backend_vk_context * ctx, vk_context& subctx, co case GGML_OP_GLU: case GGML_OP_CONV_2D_DW: { + // The tiled concat kernel is dispatched per 32x32 tile, not per element. + if (op == GGML_OP_CONCAT && pipeline == ctx->device->pipeline_concat_transpose_i32) { + elements = { (uint32_t)src1->ne[1], (uint32_t)src1->ne[0], 1 }; + break; + } uint32_t ne = ggml_nelements(dst); if (op == GGML_OP_CPY && ggml_is_quantized(src0->type) && ggml_is_quantized(dst->type)) { // Convert from number of logical elements to 2- or 4-byte units. @@ -13194,6 +14816,8 @@ static void ggml_vk_gated_linear_attn(ggml_backend_vk_context * ctx, vk_context& pc, { (uint32_t)(n_seqs * n_heads), 1, 1 }); } +static void ggml_vk_lightning_indexer_cm(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst, vk_pipeline pipeline); + static void ggml_vk_lightning_indexer(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) { const ggml_tensor * q = dst->src[0]; const ggml_tensor * k = dst->src[1]; @@ -13203,6 +14827,19 @@ static void ggml_vk_lightning_indexer(ggml_backend_vk_context * ctx, vk_context& vk_pipeline pipeline = ggml_vk_op_get_pipeline(ctx, q, k, w, dst, dst->op); GGML_ASSERT(pipeline != nullptr); + // the fork's wave64 f16 kernels take their own push-constant layout + bool fork_li = false; + for (size_t nhi = 0; nhi < LI_NH_COUNT && !fork_li; ++nhi) { + fork_li = pipeline == ctx->device->pipeline_lightning_indexer_f16[nhi] || + pipeline == ctx->device->pipeline_lightning_indexer_cm_f16[nhi] || + pipeline == ctx->device->pipeline_lightning_indexer_cm_small_f16[nhi] || + pipeline == ctx->device->pipeline_lightning_indexer_decode_cm_f16[nhi]; + } + if (fork_li) { + ggml_vk_lightning_indexer_cm(ctx, subctx, dst, pipeline); + return; + } + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); const uint32_t n_kv = k->ne[2]; @@ -13300,6 +14937,124 @@ static void ggml_vk_gated_delta_net(ggml_backend_vk_context * ctx, vk_context& s pc, { H, n_seqs, S_v }); } +static void ggml_vk_lightning_indexer_cm(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst, vk_pipeline pipeline) { + const ggml_tensor * q = dst->src[0]; + const ggml_tensor * k = dst->src[1]; + const ggml_tensor * w = dst->src[2]; + const ggml_tensor * m = dst->src[3]; + + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); + + const vk_op_lightning_indexer_cm_push_constants pc = { + (uint32_t) k->ne[2], + (uint32_t) q->ne[2], + (uint32_t) q->ne[3], + (uint32_t) m->ne[3], + (uint32_t) (dst->nb[1] / sizeof(float)), + (uint32_t) (dst->nb[3] / sizeof(float)), + (uint32_t) (q->nb[1] / sizeof(float)), + (uint32_t) (q->nb[2] / sizeof(float)), + (uint32_t) (q->nb[3] / sizeof(float)), + (uint32_t) (k->nb[2] / sizeof(ggml_fp16_t)), + (uint32_t) (k->nb[3] / sizeof(ggml_fp16_t)), + (uint32_t) (w->nb[1] / sizeof(float)), + (uint32_t) (w->nb[3] / sizeof(float)), + (uint32_t) (m->nb[1] / sizeof(ggml_fp16_t)), + (uint32_t) (m->nb[3] / sizeof(ggml_fp16_t)), + }; + + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, + {ggml_vk_tensor_subbuffer(ctx, q), ggml_vk_tensor_subbuffer(ctx, k), + ggml_vk_tensor_subbuffer(ctx, w), ggml_vk_tensor_subbuffer(ctx, m), + ggml_vk_tensor_subbuffer(ctx, dst)}, + pc, {(uint32_t) k->ne[2], (uint32_t) q->ne[2], (uint32_t) q->ne[3]}); +} + +// DSv4 fused hyper-connection ops — ports of ggml-cuda/dsv4-hc.cu. Strides are passed in +// f32 elements; grids mirror the CUDA launch geometry (flat 1D, 256 threads per workgroup). + +static void ggml_vk_dsv4_hc_pre(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) { + const ggml_tensor * x = dst->src[0]; + const ggml_tensor * w = dst->src[1]; + + const uint32_t n_embd = (uint32_t) x->ne[0]; + const uint32_t hc = (uint32_t) x->ne[1]; + const uint32_t n_tokens = (uint32_t) x->ne[2]; + const uint32_t nr = n_embd * n_tokens; + + vk_pipeline pipeline = ctx->device->pipeline_dsv4_hc_pre_f32; + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); + + const vk_op_dsv4_hc_pre_push_constants pc = { + n_embd, hc, nr, + (uint32_t)(x->nb[0] / sizeof(float)), (uint32_t)(x->nb[1] / sizeof(float)), (uint32_t)(x->nb[2] / sizeof(float)), + (uint32_t)(w->nb[0] / sizeof(float)), (uint32_t)(w->nb[1] / sizeof(float)), + (uint32_t)(dst->nb[0] / sizeof(float)), (uint32_t)(dst->nb[1] / sizeof(float)), + }; + + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, + {ggml_vk_tensor_subbuffer(ctx, x), ggml_vk_tensor_subbuffer(ctx, w), + ggml_vk_tensor_subbuffer(ctx, dst)}, + pc, {nr, 1, 1}); +} + +static void ggml_vk_dsv4_hc_comb(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) { + const ggml_tensor * mixes = dst->src[0]; + const ggml_tensor * scale = dst->src[1]; + const ggml_tensor * base = dst->src[2]; + + const uint32_t n_tokens = (uint32_t) mixes->ne[1]; + const float eps = ggml_get_op_params_f32(dst, 0); + const int32_t n_iter = ggml_get_op_params_i32(dst, 1); + + vk_pipeline pipeline = ctx->device->pipeline_dsv4_hc_comb_f32; + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); + + const vk_op_dsv4_hc_comb_push_constants pc = { + n_tokens, + (uint32_t)(mixes->nb[0] / sizeof(float)), (uint32_t)(mixes->nb[1] / sizeof(float)), + (uint32_t)(scale->nb[0] / sizeof(float)), + (uint32_t)(base->nb[0] / sizeof(float)), + (uint32_t)(dst->nb[0] / sizeof(float)), (uint32_t)(dst->nb[1] / sizeof(float)), (uint32_t)(dst->nb[2] / sizeof(float)), + eps, n_iter, + }; + + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, + {ggml_vk_tensor_subbuffer(ctx, mixes), ggml_vk_tensor_subbuffer(ctx, scale), + ggml_vk_tensor_subbuffer(ctx, base), ggml_vk_tensor_subbuffer(ctx, dst)}, + pc, {n_tokens, 1, 1}); +} + +static void ggml_vk_dsv4_hc_post(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) { + const ggml_tensor * x = dst->src[0]; + const ggml_tensor * residual = dst->src[1]; + const ggml_tensor * post = dst->src[2]; + const ggml_tensor * comb = dst->src[3]; + + const uint32_t n_embd = (uint32_t) x->ne[0]; + const uint32_t n_tokens = (uint32_t) x->ne[1]; + const uint32_t hc = (uint32_t) residual->ne[1]; + const uint32_t nr = n_embd * hc * n_tokens; + + vk_pipeline pipeline = ctx->device->pipeline_dsv4_hc_post_f32; + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); + + const vk_op_dsv4_hc_post_push_constants pc = { + n_embd, hc, nr, + (uint32_t)(x->nb[0] / sizeof(float)), (uint32_t)(x->nb[1] / sizeof(float)), + (uint32_t)(residual->nb[0] / sizeof(float)), (uint32_t)(residual->nb[1] / sizeof(float)), (uint32_t)(residual->nb[2] / sizeof(float)), + (uint32_t)(post->nb[0] / sizeof(float)), (uint32_t)(post->nb[1] / sizeof(float)), + (uint32_t)(comb->nb[0] / sizeof(float)), (uint32_t)(comb->nb[1] / sizeof(float)), (uint32_t)(comb->nb[2] / sizeof(float)), + (uint32_t)(dst->nb[0] / sizeof(float)), (uint32_t)(dst->nb[1] / sizeof(float)), (uint32_t)(dst->nb[2] / sizeof(float)), + }; + + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, + {ggml_vk_tensor_subbuffer(ctx, x), ggml_vk_tensor_subbuffer(ctx, residual), + ggml_vk_tensor_subbuffer(ctx, post), ggml_vk_tensor_subbuffer(ctx, comb), + ggml_vk_tensor_subbuffer(ctx, dst)}, + pc, {nr, 1, 1}); +} + static void ggml_vk_ssm_scan(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) { const ggml_tensor * src0 = dst->src[0]; const ggml_tensor * src1 = dst->src[1]; @@ -16259,6 +18014,22 @@ static bool ggml_vk_build_graph(ggml_backend_vk_context * ctx, ggml_cgraph * cgr break; } + // Fused silu(x)*y: run it as a swiglu split, writing straight to the MUL's destination. + if (ctx->num_additional_fused_ops == 1) { + ggml_tensor * mul = cgraph->nodes[node_idx + 1]; + ggml_tensor * other = (mul->src[0] == node) ? mul->src[1] : mul->src[0]; + + ggml_tensor fused = *mul; + fused.op = GGML_OP_GLU; + memset(fused.op_params, 0, sizeof(fused.op_params)); + ggml_set_op_params_i32(&fused, 0, (int32_t) GGML_GLU_OP_SWIGLU); + fused.src[0] = node->src[0]; + fused.src[1] = other; + + ggml_vk_glu(ctx, compute_ctx, fused.src[0], fused.src[1], &fused); + break; + } + switch (ggml_get_unary_op(node)) { case GGML_UNARY_OP_ELU: case GGML_UNARY_OP_EXP: @@ -16461,6 +18232,21 @@ static bool ggml_vk_build_graph(ggml_backend_vk_context * ctx, ggml_cgraph * cgr break; + case GGML_OP_DSV4_HC_PRE: + ggml_vk_dsv4_hc_pre(ctx, compute_ctx, node); + + break; + + case GGML_OP_DSV4_HC_COMB: + ggml_vk_dsv4_hc_comb(ctx, compute_ctx, node); + + break; + + case GGML_OP_DSV4_HC_POST: + ggml_vk_dsv4_hc_post(ctx, compute_ctx, node); + + break; + case GGML_OP_SSM_SCAN: ggml_vk_ssm_scan(ctx, compute_ctx, node); @@ -16613,8 +18399,11 @@ static void ggml_vk_cleanup(ggml_backend_vk_context * ctx) { ggml_vk_destroy_buffer(ctx->prealloc_y); ggml_vk_destroy_buffer(ctx->prealloc_split_k); ggml_vk_destroy_buffer(ctx->prealloc_add_rms_partials); + ggml_vk_destroy_buffer(ctx->fa_union_stat); ggml_vk_destroy_buffer(ctx->sync_staging); + memset(ctx->fa_union_est_ratio, 0, sizeof(ctx->fa_union_est_ratio)); + ctx->prealloc_y_last_pipeline_used = nullptr; ctx->prealloc_y_last_tensor_used = nullptr; ctx->prealloc_y_last_k_padded = false; @@ -17276,15 +19065,63 @@ static bool ggml_vk_can_fuse(const ggml_backend_vk_context * ctx, const struct g } } + // EXPERIMENT (GGML_VK_FUSE_UNARY_MUL=1): silu(x)*y is emitted as two nodes by the delta-net + // path, so the silu result makes a full round trip through memory. That is the same shape + // swiglu-split already computes in one pass, so route the pair to the existing GLU pipeline. + if (ops.size() == 2 && ops.begin()[0] == GGML_OP_UNARY && ops.begin()[1] == GGML_OP_MUL) { + static const char * env = getenv("GGML_VK_FUSE_UNARY_MUL"); + if (!(env && atoi(env) != 0)) { + return false; + } + const ggml_tensor * unary = cgraph->nodes[node_idx]; + const ggml_tensor * mul = cgraph->nodes[node_idx + 1]; + + if (ggml_get_unary_op(unary) != GGML_UNARY_OP_SILU) { + return false; + } + if (mul->src[0] != unary && mul->src[1] != unary) { + return false; + } + const ggml_tensor * other = (mul->src[0] == unary) ? mul->src[1] : mul->src[0]; + // The GLU split shader walks both inputs and the output with the same element count. + if (unary->type != GGML_TYPE_F32 || other->type != GGML_TYPE_F32 || mul->type != GGML_TYPE_F32) { + return false; + } + if (!ggml_are_same_shape(unary, other) || !ggml_are_same_shape(unary, mul)) { + return false; + } + if (!ggml_is_contiguous(unary->src[0]) || !ggml_is_contiguous(other) || !ggml_is_contiguous(mul)) { + return false; + } + return true; + } + auto const &mmid_mul_ok = [&](const ggml_tensor *mmid, const ggml_tensor *mul) { const ggml_tensor *scale = mul->src[1]; if (mmid != mul->src[0]) { return false; } - // mat-vec only + // EXPERIMENT (GGML_VK_MMID_SCALE_EPILOGUE=1): the tile shader can apply the scale as it + // writes out, which removes a full write+read of the matmul result at prefill. The + // coopmat2 shader has the binding but not the epilogue, so it stays on the old path. if (!ggml_vk_use_mul_mat_vec_id(cgraph, node_idx)) { - return false; + static const char * env = getenv("GGML_VK_MMID_SCALE_EPILOGUE"); + if (!(env && atoi(env) != 0) || ctx->device->coopmat2) { + return false; + } + // Shader indexes the scale as [token * nei0 + expert_slot]. + if (scale->type != GGML_TYPE_F32 || mul->type != GGML_TYPE_F32 || !ggml_is_contiguous(scale)) { + return false; + } + if (get_misalign_bytes(ctx, scale) != 0) { + return false; + } + // The shader indexes the scale as [token * nei0 + expert_slot] from row_ids, which + // carries no 4th dimension, so ne[3] must be 1 or later batches read the wrong scale. + return scale->ne[0] == 1 && mmid->ne[3] == 1 && scale->ne[3] == 1 && + scale->ne[1] == mmid->ne[1] && scale->ne[2] == mmid->ne[2] && + ggml_are_same_shape(mul, mmid); } // shaders assume the types match if (mmid->type != scale->type) { @@ -17884,14 +19721,22 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg ctx->query_fusion_node_count.resize(ctx->num_queries); ctx->query_nodes.resize(ctx->num_queries); ctx->query_node_idx.resize(ctx->num_queries); + ctx->query_sub_names.resize(ctx->num_queries); } - ctx->device->device.resetQueryPool(ctx->query_pool, 0, cgraph->n_nodes+1); + // Reset the whole pool, not just n_nodes+1: sub-op marks consume slots past that. + ctx->device->device.resetQueryPool(ctx->query_pool, 0, ctx->num_queries); std::fill(ctx->query_fusion_names.begin(), ctx->query_fusion_names.end(), nullptr); std::fill(ctx->query_fusion_node_count.begin(), ctx->query_fusion_node_count.end(), 0); std::fill(ctx->query_nodes.begin(), ctx->query_nodes.end(), nullptr); std::fill(ctx->query_node_idx.begin(), ctx->query_node_idx.end(), 0); + // Under partial offload the scheduler's async input copies between graph + // splits can leave commands in a pending compute ctx. Flush it so the + // timestamp stream starts on a fresh command buffer. + if (!ctx->compute_ctx.expired()) { + ggml_vk_synchronize(ctx); + } GGML_ASSERT(ctx->compute_ctx.expired()); compute_ctx = ggml_vk_get_compute_ctx(ctx); ctx->query_idx = 0; @@ -17918,6 +19763,7 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg uint32_t submitted_nodes = 0; uint32_t submit_count = 0; uint64_t batch_flops = 0; + uint64_t batch_bytes = 0; uint64_t total_flops = 0; uint64_t flops_cap = 200'000'000'000ULL; @@ -17955,6 +19801,7 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg first_node_in_batch = true; submitted_nodes = 0; batch_flops = 0; + batch_bytes = 0; if (submit_count < 3) { flops_per_submit *= 2; } @@ -17969,9 +19816,11 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg { auto node_flops = ggml_vk_get_node_flops(cgraph->nodes[i]); total_flops += node_flops; + auto node_bytes = ggml_vk_get_node_bytes(cgraph->nodes[i]); - // Flush the current batch before recording a node that would push it over the flop threshold - if (flops_per_submit != 0 && submitted_nodes > 0 && batch_flops + node_flops >= flops_per_submit) { + // Flush the current batch before recording a node that would push it over the flop or byte threshold + if ((flops_per_submit != 0 && submitted_nodes > 0 && batch_flops + node_flops >= flops_per_submit) || + (ctx->device->max_bytes_per_submit != 0 && submitted_nodes > 0 && batch_bytes + node_bytes >= ctx->device->max_bytes_per_submit)) { vk_context flush_ctx = ggml_vk_get_compute_ctx(ctx); ggml_vk_ctx_end(flush_ctx); flush_ctx->exit_tensor_idx = -1; @@ -17982,6 +19831,7 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg } batch_flops += node_flops; + batch_bytes += node_bytes; } // op_srcs_fused_elementwise indicates whether an op's srcs all contribute to @@ -18027,6 +19877,9 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg fusion_string = "MUL_MAT_ID_MUL"; op_srcs_fused_elementwise[0] = false; op_srcs_fused_elementwise[1] = true; + } else if (ggml_vk_can_fuse(ctx, cgraph, i, { GGML_OP_UNARY, GGML_OP_MUL })) { + ctx->num_additional_fused_ops = 1; + fusion_string = "SILU_MUL"; } else if (ggml_can_fuse_subgraph(cgraph, i, { GGML_OP_RMS_NORM, GGML_OP_MUL, GGML_OP_ROPE, GGML_OP_VIEW, GGML_OP_SET_ROWS }, { i + 4 }) && ggml_check_edges(cgraph, i, rms_norm_mul_rope_view_set_rows_edges) && ggml_vk_can_fuse_rms_norm_mul_rope(ctx, cgraph, i) && @@ -18211,6 +20064,7 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg bool almost_ready = (cgraph->n_nodes - i) < cgraph->n_nodes / 5; bool submit = (submitted_nodes >= ctx->device->max_nodes_per_submit) || (flops_per_submit != 0 && batch_flops >= flops_per_submit) || + (ctx->device->max_bytes_per_submit != 0 && batch_bytes >= ctx->device->max_bytes_per_submit) || (i + ctx->num_additional_fused_ops >= last_node) || (almost_ready && !ctx->almost_ready_fence_pending); @@ -18222,6 +20076,7 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg // track a single node/fusion for the current query ctx->query_nodes[ctx->query_idx] = cgraph->nodes[i]; ctx->query_fusion_names[ctx->query_idx] = fusion_string; + ctx->query_sub_names[ctx->query_idx] = nullptr; compute_ctx->s->buffer->buf.writeTimestamp(vk::PipelineStageFlagBits::eAllCommands, ctx->query_pool, ctx->query_idx++); ggml_vk_sync_buffers(ctx, compute_ctx); } else { @@ -18263,14 +20118,22 @@ static ggml_status ggml_backend_vk_graph_compute(ggml_backend_t backend, ggml_cg ctx->compute_ctx.reset(); // Get the results and pass them to the logger - std::vector timestamps(cgraph->n_nodes + 1); - VK_CHECK(ctx->device->device.getQueryPoolResults(ctx->query_pool, 0, ctx->query_idx, (cgraph->n_nodes + 1)*sizeof(uint64_t), timestamps.data(), sizeof(uint64_t), vk::QueryResultFlagBits::e64 | vk::QueryResultFlagBits::eWait), "get timestamp results", ctx->device); + // Sized to the pool, not n_nodes+1: sub-op marks push query_idx past the node count. + std::vector timestamps(ctx->num_queries); + VK_CHECK(ctx->device->device.getQueryPoolResults(ctx->query_pool, 0, ctx->query_idx, ctx->num_queries*sizeof(uint64_t), timestamps.data(), sizeof(uint64_t), vk::QueryResultFlagBits::e64 | vk::QueryResultFlagBits::eWait), "get timestamp results", ctx->device); if (!vk_perf_logger_concurrent) { // Log each op separately for (int i = 1; i < ctx->query_idx; i++) { + const uint64_t dt = uint64_t((timestamps[i] - timestamps[i-1]) * ctx->device->properties.limits.timestampPeriod); + if (ctx->query_sub_names[i] != nullptr) { + // sub-node interval (e.g. the FA K/V contiguize pass) - billed separately so + // it is not silently folded into the op that dispatched it + ctx->perf_logger->log_timing_named(ctx->query_sub_names[i], dt); + continue; + } auto node = ctx->query_nodes[i]; auto name = ctx->query_fusion_names[i]; - ctx->perf_logger->log_timing(node, name, uint64_t((timestamps[i] - timestamps[i-1]) * ctx->device->properties.limits.timestampPeriod)); + ctx->perf_logger->log_timing(node, name, dt); } } else { // Log each group of nodes @@ -19021,25 +20884,34 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm if (op->src[3] && op->src[3]->type != GGML_TYPE_F16) { return false; } - auto fa_kv_ok = [](ggml_type t) { - switch (t) { - case GGML_TYPE_F32: - case GGML_TYPE_F16: - case GGML_TYPE_BF16: - case GGML_TYPE_Q8_0: - case GGML_TYPE_Q5_1: - case GGML_TYPE_Q5_0: - case GGML_TYPE_Q4_1: - case GGML_TYPE_Q4_0: - case GGML_TYPE_IQ4_NL: - return true; - default: - return false; - } + auto fa_kv_ok = [&](ggml_type t) { + // ggml_vk_fa_kv_native is the single source of truth; on this base every + // admitted type is native, and the non-native hard-gate below is dormant + // hardening for any future type routed through the dequant-once scratch. + return ggml_vk_fa_kv_native(t, coopmat2); }; if (!fa_kv_ok(op->src[1]->type) || !fa_kv_ok(op->src[2]->type)) { return false; } + if (!ggml_vk_fa_kv_native(op->src[1]->type, coopmat2) || !ggml_vk_fa_kv_native(op->src[2]->type, coopmat2)) { + // Only correct through the dequant-once scratch path; admit only when every + // hard condition of the dispatch-time gate holds, so dispatch can never fall + // back to the native shader (it reads garbage for these types, not an error). + const ggml_tensor * k = op->src[1]; + const ggml_tensor * v = op->src[2]; + static const char * fa_dequant_env = getenv("GGML_VK_FA_DEQUANT"); + const bool fa_dequant_off = fa_dequant_env && fa_dequant_env[0] == '0'; + const uint64_t kv_f16_sz = ((uint64_t)ggml_nelements(k) + (uint64_t)ggml_nelements(v)) * sizeof(ggml_fp16_t); + if (fa_dequant_off || + op->src[0]->ne[1] < 64 || + device->pipeline_dequant_transpose[k->type] == nullptr || + device->pipeline_dequant_transpose[v->type] == nullptr || + k->nb[0] != ggml_type_size(k->type) || v->nb[0] != ggml_type_size(v->type) || + !ggml_is_contiguously_allocated(k) || !ggml_is_contiguously_allocated(v) || + kv_f16_sz > device->properties.limits.maxStorageBufferRange) { + return false; + } + } if ((op->src[1]->type == GGML_TYPE_BF16) != (op->src[2]->type == GGML_TYPE_BF16)) { return false; } @@ -19378,6 +21250,16 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm case GGML_OP_GATED_LINEAR_ATTN: // the shader block size is hardcoded to head_size 64 return op->src[0]->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32 && op->src[0]->ne[0] == 64; + case GGML_OP_DSV4_HC_PRE: + return op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F32 && + op->type == GGML_TYPE_F32; + case GGML_OP_DSV4_HC_COMB: + return op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F32 && + op->src[2]->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32; + case GGML_OP_DSV4_HC_POST: + return op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F32 && + op->src[2]->type == GGML_TYPE_F32 && op->src[3]->type == GGML_TYPE_F32 && + op->type == GGML_TYPE_F32; case GGML_OP_LIGHTNING_INDEXER: { const ggml_tensor * q = op->src[0]; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/concat_transpose.comp b/ggml/src/ggml-vulkan/vulkan-shaders/concat_transpose.comp new file mode 100644 index 000000000000..653aeaa011b2 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/concat_transpose.comp @@ -0,0 +1,43 @@ +#version 450 + +#include "types.glsl" +#include "generic_binary_head.glsl" + +layout(local_size_x = 32, local_size_y = 8, local_size_z = 1) in; + +// dim-0 concat whose src1 is transposed. The generic kernel reads src1 with the transposed +// stride, so neighbouring lanes touch different cache lines. Stage a 32x32 tile in shared +// memory instead, which keeps both the load and the store coalesced. 33 columns pads away +// the shared-memory bank conflicts. +shared A_TYPE tmp[32][33]; + +void main() { + const uint tx = gl_LocalInvocationID.x; + const uint ty = gl_LocalInvocationID.y; + + const uint row = gl_WorkGroupID.x * 32 + tx; + + // src0 is already contiguous, copy it straight through. + if (gl_WorkGroupID.y == 0 && row < p.ne01) { + for (uint i0 = ty; i0 < p.ne00; i0 += 8) { + data_d[get_doffset() + row*p.nb21 + i0*p.nb20] = D_TYPE(data_a[get_aoffset() + row*p.nb01 + i0*p.nb00]); + } + } + + [[unroll]] for (uint j = 0; j < 32; j += 8) { + const uint c = gl_WorkGroupID.y * 32 + ty + j; + if (c < p.ne10 && row < p.ne11) { + tmp[ty + j][tx] = A_TYPE(data_b[get_boffset() + c*p.nb10 + row*p.nb11]); + } + } + + barrier(); + + const uint col = gl_WorkGroupID.y * 32 + tx; + [[unroll]] for (uint j = 0; j < 32; j += 8) { + const uint r = gl_WorkGroupID.x * 32 + ty + j; + if (col < p.ne10 && r < p.ne11) { + data_d[get_doffset() + r*p.nb21 + (p.ne00 + col)*p.nb20] = D_TYPE(tmp[tx][ty + j]); + } + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_f16_transpose.comp b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_f16_transpose.comp new file mode 100644 index 000000000000..dda53d749acc --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_f16_transpose.comp @@ -0,0 +1,32 @@ +#version 450 + +#include "dequant_head.glsl" + +#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require + +layout(local_size_x = 256, local_size_y = 1, local_size_z = 1) in; + +layout (binding = 0) readonly buffer A {f16vec4 data_a[];}; +layout (binding = 1) writeonly buffer D {f16vec4 data_b[];}; + +// Strided-copy counterpart of the fused dequant+transpose shaders for FA f16 KV: +// source physical is [HS, NH, KV, NS] (p.M=HS, p.K=NH, p.stride_a=KV); write to +// per-head-contiguous dest [HS, KV, NH, NS] so the f16 FA reads KV coalesced. +// HS stays innermost in both layouts, so each invocation moves 8 HS-contiguous +// elements (two f16vec4) requiring HS % 8 == 0 (enforced by the host gate). +void main() { + const uint i = gl_GlobalInvocationID.x; + const uint e0 = i * 8; + if (e0 >= p.nel) { + return; + } + + const uint HS = p.M, NH = p.K, KVn = p.stride_a; + const uint dst = (e0 % HS) + + ((e0 / (HS * NH)) % KVn) * HS + + ((e0 / HS) % NH) * (HS * KVn) + + (e0 / (HS * NH * KVn)) * (HS * KVn * NH); + + data_b[dst / 4 ] = data_a[e0 / 4 ]; + data_b[dst / 4 + 1] = data_a[e0 / 4 + 1]; +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_iq4_nl.comp b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_iq4_nl.comp index 8f7833eab2e7..befcdc1c5100 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_iq4_nl.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_iq4_nl.comp @@ -10,8 +10,6 @@ layout (binding = 1) writeonly buffer D {D_TYPE data_b[];}; void main() { const uint i = gl_WorkGroupID.x * 4 + gl_LocalInvocationID.x / 64; - init_iq_shmem(gl_WorkGroupSize); - const uint tid = gl_LocalInvocationID.x % 64; const uint il = tid/32; const uint ir = tid%32; @@ -21,12 +19,27 @@ void main() { } const uint q_idx = 8*il; + +#ifdef DEQUANT_TRANSPOSE + // Fused dequant+transpose for FA quant-KV: source physical is [HS, NH, KV, NS] + // (p.M=HS, p.K=NH, p.stride_a=KV); write to per-head-contiguous dest [HS, KV, NH, NS] so the + // f16 FA reads KV coalesced. An iq4_nl block = 32 consecutive HS elements at fixed (head,kv) -> + // 32 contiguous dest positions (intra-block nibble offsets unchanged). + const uint HS = p.M, NH = p.K, KVn = p.stride_a; + const uint e0 = ib * 32; + const uint b_idx = (e0 % HS) + + ((e0 / (HS * NH)) % KVn) * HS + + ((e0 / HS) % NH) * (HS * KVn) + + (e0 / (HS * NH * KVn)) * (HS * KVn * NH) + + q_idx; +#else const uint b_idx = 1024*i + 32*ir + q_idx; +#endif const float d = float(data_a[ib].d); [[unroll]] for (uint l = 0; l < 8; ++l) { - data_b[b_idx + l + 0] = D_TYPE(d * kvalues_iq4nl[data_a[ib].qs[q_idx + l] & 0xF]); - data_b[b_idx + l + 16] = D_TYPE(d * kvalues_iq4nl[data_a[ib].qs[q_idx + l] >> 4]); + data_b[b_idx + l + 0] = D_TYPE(d * float(kvalues_iq4nl_const[data_a[ib].qs[q_idx + l] & 0xF])); + data_b[b_idx + l + 16] = D_TYPE(d * float(kvalues_iq4nl_const[data_a[ib].qs[q_idx + l] >> 4])); } } diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q4_0.comp b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q4_0.comp index b92b292135b4..51a6c89a60c8 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q4_0.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q4_0.comp @@ -19,7 +19,19 @@ void main() { } const uint q_idx = 8*il; + +#ifdef DEQUANT_TRANSPOSE + // fused dequant+transpose for FA quant-KV: per-head-contiguous f16 scratch (see dequant_q8_0.comp) + const uint HS = p.M, NH = p.K, KVn = p.stride_a; + const uint e0 = ib * 32; + const uint b_idx = (e0 % HS) + + ((e0 / (HS * NH)) % KVn) * HS + + ((e0 / HS) % NH) * (HS * KVn) + + (e0 / (HS * NH * KVn)) * (HS * KVn * NH) + + q_idx; +#else const uint b_idx = 1024*i + 32*ir + q_idx; +#endif const float d = float(data_a[ib].d); diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q4_1.comp b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q4_1.comp index 6b63cbe5833b..76f2d958cbab 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q4_1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q4_1.comp @@ -18,7 +18,18 @@ void main() { return; } +#ifdef DEQUANT_TRANSPOSE + // fused dequant+transpose for FA quant-KV: per-head-contiguous f16 scratch (see dequant_q8_0.comp) + const uint HS = p.M, NH = p.K, KVn = p.stride_a; + const uint e0 = ib * 32; + const uint b_idx = (e0 % HS) + + ((e0 / (HS * NH)) % KVn) * HS + + ((e0 / HS) % NH) * (HS * KVn) + + (e0 / (HS * NH * KVn)) * (HS * KVn * NH) + + 8*il; +#else const uint b_idx = 1024*i + 32*ir + 8*il; +#endif const float d = float(data_a[ib].d); const float m = float(data_a[ib].m); diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q5_0.comp b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q5_0.comp index f1b0bac87271..1402fa4292d8 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q5_0.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q5_0.comp @@ -18,7 +18,18 @@ void main() { return; } +#ifdef DEQUANT_TRANSPOSE + // fused dequant+transpose for FA quant-KV: per-head-contiguous f16 scratch (see dequant_q8_0.comp) + const uint HS = p.M, NH = p.K, KVn = p.stride_a; + const uint e0 = ib * 32; + const uint b_idx = (e0 % HS) + + ((e0 / (HS * NH)) % KVn) * HS + + ((e0 / HS) % NH) * (HS * KVn) + + (e0 / (HS * NH * KVn)) * (HS * KVn * NH) + + 8*il; +#else const uint b_idx = 1024*i + 32*ir + 8*il; +#endif const float d = float(data_a[ib].d); const uint qh = uint(data_a[ib].qh[1]) << 16 | data_a[ib].qh[0]; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q5_1.comp b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q5_1.comp index c495b31f1754..1fd6e2552af7 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q5_1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q5_1.comp @@ -18,7 +18,18 @@ void main() { return; } +#ifdef DEQUANT_TRANSPOSE + // fused dequant+transpose for FA quant-KV: per-head-contiguous f16 scratch (see dequant_q8_0.comp) + const uint HS = p.M, NH = p.K, KVn = p.stride_a; + const uint e0 = ib * 32; + const uint b_idx = (e0 % HS) + + ((e0 / (HS * NH)) % KVn) * HS + + ((e0 / HS) % NH) * (HS * KVn) + + (e0 / (HS * NH * KVn)) * (HS * KVn * NH) + + 8*il; +#else const uint b_idx = 1024*i + 32*ir + 8*il; +#endif const float d = float(data_a[ib].d); const float m = float(data_a[ib].m); diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_comb.comp b/ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_comb.comp new file mode 100644 index 000000000000..0449715bf458 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_comb.comp @@ -0,0 +1,101 @@ +#version 450 + +// DeepSeek V4 fused hyper-connection "comb": build the 4x4 stream-mixing matrix per token — +// per-source softmax over scaled+biased logits, then eps-stabilized alternating column/row +// (sinkhorn) normalization. Port of ggml-cuda/dsv4-hc.cu (hc_comb): one THREAD per token, +// the whole 4x4 lives in registers. At decode this is a single active thread by design — +// the fusion's value is collapsing the ~dozen decomposed graph ops (and their intermediate +// tensors) into one dispatch, not throughput. + +layout(local_size_x = 256, local_size_y = 1, local_size_z = 1) in; + +layout(binding = 0) readonly buffer MBuf { float data_mixes[]; }; +layout(binding = 1) readonly buffer SBuf { float data_scale[]; }; +layout(binding = 2) readonly buffer BBuf { float data_base[]; }; +layout(binding = 3) writeonly buffer DBuf { float data_d[]; }; + +layout(push_constant) uniform Parameters { + uint n_tokens; + uint sm0, sm1; + uint ss0; + uint sb0; + uint sd0, sd1, sd2; + float eps; + int n_iter; +} p; + +const uint HC = 4; +const uint COMB_OFFSET = 2 * HC; // comb logits start after the pre/post blocks of the mix vector + +void norm_cols(inout float comb[HC * HC]) { + for (uint idst = 0; idst < HC; ++idst) { + float sum = p.eps; + for (uint isrc = 0; isrc < HC; ++isrc) { + sum += comb[idst + HC * isrc]; + } + const float inv_sum = 1.0 / sum; + for (uint isrc = 0; isrc < HC; ++isrc) { + comb[idst + HC * isrc] *= inv_sum; + } + } +} + +void norm_rows(inout float comb[HC * HC]) { + for (uint isrc = 0; isrc < HC; ++isrc) { + float sum = p.eps; + for (uint idst = 0; idst < HC; ++idst) { + sum += comb[idst + HC * isrc]; + } + const float inv_sum = 1.0 / sum; + for (uint idst = 0; idst < HC; ++idst) { + comb[idst + HC * isrc] *= inv_sum; + } + } +} + +void main() { + const uint it = gl_GlobalInvocationID.x; + if (it >= p.n_tokens) { + return; + } + + const float scale_comb = data_scale[2 * p.ss0]; + float comb[HC * HC]; + + for (uint isrc = 0; isrc < HC; ++isrc) { + float vmax = uintBitsToFloat(0xff800000); // -inf + for (uint idst = 0; idst < HC; ++idst) { + const uint idx = idst + HC * isrc; + const float v = data_mixes[(COMB_OFFSET + idx) * p.sm0 + it * p.sm1] * scale_comb + + data_base[(COMB_OFFSET + idx) * p.sb0]; + comb[idx] = v; + vmax = max(vmax, v); + } + + float sum = 0.0; + for (uint idst = 0; idst < HC; ++idst) { + const uint idx = idst + HC * isrc; + const float v = exp(comb[idx] - vmax); + comb[idx] = v; + sum += v; + } + + const float inv_sum = 1.0 / sum; + for (uint idst = 0; idst < HC; ++idst) { + const uint idx = idst + HC * isrc; + comb[idx] = comb[idx] * inv_sum + p.eps; + } + } + + norm_cols(comb); + for (int i = 1; i < p.n_iter; ++i) { + norm_rows(comb); + norm_cols(comb); + } + + for (uint isrc = 0; isrc < HC; ++isrc) { + for (uint idst = 0; idst < HC; ++idst) { + data_d[idst * p.sd0 + isrc * p.sd1 + it * p.sd2] = comb[idst + HC * isrc]; + } + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_post.comp b/ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_post.comp new file mode 100644 index 000000000000..212109dfd434 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_post.comp @@ -0,0 +1,44 @@ +#version 450 + +// DeepSeek V4 fused hyper-connection "post": redistribute the layer output back into the HC +// streams with the sinkhorn-mixed residual, +// dst[i0, idst, it] = x[i0, it] * post[idst, it] + sum_isrc residual[i0, isrc, it] * comb[idst, isrc, it]. +// Port of ggml-cuda/dsv4-hc.cu (hc_post). Flat elementwise over n_embd * hc * n_tokens. + +layout(local_size_x = 256, local_size_y = 1, local_size_z = 1) in; + +layout(binding = 0) readonly buffer XBuf { float data_x[]; }; +layout(binding = 1) readonly buffer RBuf { float data_r[]; }; +layout(binding = 2) readonly buffer PBuf { float data_p[]; }; +layout(binding = 3) readonly buffer CBuf { float data_c[]; }; +layout(binding = 4) writeonly buffer DBuf { float data_d[]; }; + +layout(push_constant) uniform Parameters { + uint n_embd; + uint hc; + uint nr; // n_embd * hc * n_tokens + uint sx0, sx1; + uint sr0, sr1, sr2; + uint sp0, sp1; + uint sc0, sc1, sc2; + uint sd0, sd1, sd2; +} p; + +void main() { + const uint ir = gl_GlobalInvocationID.x; + if (ir >= p.nr) { + return; + } + + const uint i0 = ir % p.n_embd; + const uint idst = (ir / p.n_embd) % p.hc; + const uint it = ir / (p.n_embd * p.hc); + + float sum = data_x[i0 * p.sx0 + it * p.sx1] * data_p[idst * p.sp0 + it * p.sp1]; + for (uint isrc = 0; isrc < p.hc; ++isrc) { + sum += data_r[i0 * p.sr0 + isrc * p.sr1 + it * p.sr2] + * data_c[idst * p.sc0 + isrc * p.sc1 + it * p.sc2]; + } + + data_d[i0 * p.sd0 + idst * p.sd1 + it * p.sd2] = sum; +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_pre.comp b/ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_pre.comp new file mode 100644 index 000000000000..b6cc09fe2cfa --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dsv4_hc_pre.comp @@ -0,0 +1,37 @@ +#version 450 + +// DeepSeek V4 fused hyper-connection "pre": mix the HC input streams down to one embedding, +// dst[i0, it] = sum_ih x[i0, ih, it] * w[ih, it]. Port of ggml-cuda/dsv4-hc.cu (hc_pre). +// Flat elementwise kernel over n_embd * n_tokens; strides are in f32 elements. + +layout(local_size_x = 256, local_size_y = 1, local_size_z = 1) in; + +layout(binding = 0) readonly buffer XBuf { float data_x[]; }; +layout(binding = 1) readonly buffer WBuf { float data_w[]; }; +layout(binding = 2) writeonly buffer DBuf { float data_d[]; }; + +layout(push_constant) uniform Parameters { + uint n_embd; + uint hc; + uint nr; // n_embd * n_tokens + uint sx0, sx1, sx2; + uint sw0, sw1; + uint sd0, sd1; +} p; + +void main() { + const uint ir = gl_GlobalInvocationID.x; + if (ir >= p.nr) { + return; + } + + const uint i0 = ir % p.n_embd; + const uint it = ir / p.n_embd; + + float sum = data_x[i0 * p.sx0 + it * p.sx2] * data_w[it * p.sw1]; + for (uint ih = 1; ih < p.hc; ++ih) { + sum += data_x[i0 * p.sx0 + ih * p.sx1 + it * p.sx2] * data_w[ih * p.sw0 + it * p.sw1]; + } + + data_d[i0 * p.sd0 + it * p.sd1] = sum; +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp index 0c1b6d0673e9..5b19ff61093f 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp @@ -186,9 +186,9 @@ void main() { // FaBlockBytesK/V == 2 for f16, 16 for f32, ggml block byte size for quants. uint32_t k_offset = (ik2*p.nb12 + ik3*p.nb13) / FaBlockBytesK; uint32_t v_offset = (iv2*p.nb22 + iv3*p.nb23) / FaBlockBytesV; - uint32_t m_offset = gqa_iq1*KV; + uint32_t m_offset = gqa_iq1*m_row_len; if (p.nem2 != 1 || p.nem3 != 1) { - m_offset += ((iq3 % p.nem3) * p.nem2 + (iq2 % p.nem2)) * p.nem1 * KV; + m_offset += ((iq3 % p.nem3) * p.nem2 + (iq2 % p.nem2)) * p.nem1 * m_row_len; mo_offset += ((iq3 % p.nem3) * p.nem2 + (iq2 % p.nem2)) * CEIL_DIV(p.nem1, Br) * mo_stride; } @@ -210,7 +210,7 @@ void main() { } // Only load if the block is not all zeros if (mask_opt_bits != MASK_OPT_ALL_ZERO) { - bool nem1_bounds_check = !(p.gqa_ratio > 1) && (p.nem1 % Br) != 0; + bool nem1_bounds_check = !(gqa_ratio > 1) && (p.nem1 % Br) != 0; float max_mask = NEG_FLT_MAX_OVER_2; barrier(); @@ -420,7 +420,11 @@ void main() { acc += dotPacked4x8EXT(Qf[qib].qs[qiqs + d], k_quants[d]); } - Sf[r][c] += ACC_TYPE(acc) * ACC_TYPE(Qf[qib].ds.x) * k_dm.x; + // scale in fp32 before narrowing: acc is an int32 sum of dotPacked4x8EXT + // results, bounded by d_per_step*4*127*127 with q8_0 on both sides, which + // overflows f16 when ACC_TYPE is f16 (GGML_PREC_DEFAULT). Identical + // arithmetic when ACC_TYPE is float. + Sf[r][c] += ACC_TYPE(float(acc) * float(Qf[qib].ds.x) * float(k_dm.x)); if ((d_tid * (HSK_per_thread / 4) + d_block) % 8 == 0) { Sf[r][c] += k_dot_correction(qib, k_dm); } @@ -656,10 +660,10 @@ void main() { // If there is split_k, then the split_k resolve shader does the final // division by L. Store the intermediate O value and per-row m and L values. - if (p.k_num > 1) { - if (p.gqa_ratio > 1) { + if (partial_output || split_k_num > 1) { + if (gqa_ratio > 1) { // note: O and Q have swapped coord 1,2. - uint32_t o_offset = HSV * p.ne1 * (split_k_index + p.k_num * (gqa_iq1 + p.ne2 * iq3)) / 4; + uint32_t o_offset = HSV * p.ne1 * (split_k_index + output_k_num * (gqa_iq1 + p.ne2 * iq3)) / 4; [[unroll]] for (uint32_t r = 0; r < rows_per_thread; ++r) { const uint row = tile_row(r); @@ -670,7 +674,7 @@ void main() { } } - o_offset = HSV * p.ne1 * p.k_num * p.ne2 * p.ne3 + p.ne1 * 2 * (split_k_index + p.k_num * (gqa_iq1 + p.ne2 * iq3)); + o_offset = HSV * p.ne1 * output_k_num * p.ne2 * p.ne3 + p.ne1 * 2 * (split_k_index + output_k_num * (gqa_iq1 + p.ne2 * iq3)); [[unroll]] for (uint32_t r = 0; r < rows_per_thread; ++r) { const uint row = tile_row(r); if (row < N) { @@ -684,7 +688,7 @@ void main() { const uint global_row = i * Br + row; if (global_row < N) { - uint32_t o_offset = HSV * p.ne1 * (split_k_index + p.k_num * (global_row + p.ne2 * iq3)) / 4; + uint32_t o_offset = HSV * p.ne1 * (split_k_index + output_k_num * (global_row + p.ne2 * iq3)) / 4; [[unroll]] for (uint32_t d = 0; d < HSV_per_thread / 4; ++d) { data_ov4[o_offset + iq2 * HSV/4 + d * D_split + d_tid] = D_TYPEV4(Of[r][d]); @@ -692,7 +696,7 @@ void main() { } if (global_row < N && d_tid == 0 && col_tid == 0) { - uint32_t lm_offset = HSV * p.ne1 * p.k_num * p.ne2 * p.ne3 + p.ne1 * 2 * (split_k_index + p.k_num * (global_row + p.ne2 * iq3)); + uint32_t lm_offset = HSV * p.ne1 * output_k_num * p.ne2 * p.ne3 + p.ne1 * 2 * (split_k_index + output_k_num * (global_row + p.ne2 * iq3)); data_o[lm_offset + iq2] = D_TYPE(Lf[r]); data_o[lm_offset + p.ne1 + iq2] = D_TYPE(Mf[r]); } @@ -738,7 +742,7 @@ void main() { uint32_t o_offset = (gqa_iq1*p.ne1*HSV + iq3*p.ne2*p.ne1*HSV) / 4; - if (p.gqa_ratio > 1) { + if (gqa_ratio > 1) { [[unroll]] for (uint32_t r = 0; r < rows_per_thread; ++r) { const uint row = tile_row(r); if (row < N) { diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl index 0ce4503a8847..b562c5d78749 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl @@ -24,6 +24,11 @@ const bool USE_MASK_OPT = (Flags & 1) != 0; const bool MASK_ENABLE = (Flags & 2) != 0; const bool LOGIT_SOFTCAP = (Flags & 4) != 0; const bool OLD_AMD_WINDOWS = (Flags & 8) != 0; +// KV comes from a buffer instead of the push constant. Used by paths that compact K/V on the +// GPU, where the row count is only known after a dedup pass and so cannot be pushed. The +// workgroup counts derive from neq1/neq2/neq3 and never from KV, so no indirect dispatch is +// needed: only this loop bound changes. Folds away for every other pipeline. +const bool DYNAMIC_KV = (Flags & 16) != 0; // Round up head sizes to a multiple of 16, for coopmat1/coopmat2 paths const uint32_t HSK_pad = (HSK + 15) & ~15; @@ -81,6 +86,7 @@ layout (binding = 5) writeonly buffer O {D_TYPE data_o[];}; layout (binding = 5) writeonly buffer OV4 {D_TYPEV4 data_ov4[];}; layout (binding = 6) readonly buffer MO {uint32_t data_mask_opt[];}; +layout (binding = 7) readonly buffer KVB {uint32_t data_kv_dyn[];}; #define MASK_OPT_ALL_NEG_INF 1 #define MASK_OPT_ALL_ZERO 2 @@ -124,7 +130,7 @@ ACC_TYPE perElemOpStoreCol0(const in uint32_t r, const in uint32_t c, const in A // Load the slope matrix, indexed by Q's dimension 2. ACC_TYPE perElemOpComputeSlope(const in uint32_t r, const in uint32_t c, const in ACC_TYPE elem, const in uint32_t iq2) { - const uint32_t h = iq2 + (r % p.gqa_ratio); + const uint32_t h = iq2 + (r % (p.gqa_ratio & 0xffff)); uint32_t n_head_log2 = p.mask_n_head_log2 & N_LOG2_MASK; @@ -137,32 +143,40 @@ ACC_TYPE perElemOpComputeSlope(const in uint32_t r, const in uint32_t c, const i // Load the sink value, indexed by Q's dimension 2. ACC_TYPE perElemOpGetSink(const in uint32_t r, const in uint32_t c, const in ACC_TYPE elem, const in uint32_t iq2) { - const uint32_t h = iq2 + (r % p.gqa_ratio); + const uint32_t h = iq2 + (r % (p.gqa_ratio & 0xffff)); return ACC_TYPE(data_s[h]); } uint32_t i, N, KV, split_k_index, Tr, start_j, end_j, gqa_iq1, iq2, iq3, rk2, rk3, rv2, rv3, ik2, ik3, iv2, iv3, - q_stride, k_stride, v_stride, m_stride; + q_stride, k_stride, v_stride, m_stride, m_row_len, gqa_ratio, split_k_num, output_k_num; +bool partial_output; void init_indices() { N = p.N; - KV = p.KV; + KV = DYNAMIC_KV ? data_kv_dyn[0] : p.KV; + gqa_ratio = p.gqa_ratio & 0xffff; + split_k_num = p.k_num & 0xffff; + output_k_num = p.k_num >> 16; + partial_output = output_k_num != 0; + if (!partial_output) { + output_k_num = split_k_num; + } - if (p.k_num > 1) { - if (p.gqa_ratio > 1) { + if (split_k_num > 1) { + if (gqa_ratio > 1) { i = 0; // batch and split_k share gl_WorkGroupID.x - gqa_iq1 = gl_WorkGroupID.x / p.k_num; - split_k_index = gl_WorkGroupID.x % p.k_num; + gqa_iq1 = gl_WorkGroupID.x / split_k_num; + split_k_index = gl_WorkGroupID.x % split_k_num; } else { gqa_iq1 = 0; - split_k_index = gl_WorkGroupID.x % p.k_num; - i = gl_WorkGroupID.x / p.k_num; + split_k_index = gl_WorkGroupID.x % split_k_num; + i = gl_WorkGroupID.x / split_k_num; } - } else if (p.gqa_ratio > 1) { + } else if (gqa_ratio > 1) { i = 0; gqa_iq1 = gl_WorkGroupID.x; split_k_index = 0; @@ -179,7 +193,7 @@ void init_indices() // When not using grouped query attention, all rows share the same iq2, equal to gl_WorkGroupID.y. // When using grouped query attention, each workgroup does gqa_ratio consecutive values of iq2. - iq2 = gl_WorkGroupID.y * p.gqa_ratio; + iq2 = gl_WorkGroupID.y * gqa_ratio; iq3 = gl_WorkGroupID.z; // broadcast factors @@ -200,14 +214,24 @@ void init_indices() // nb?1 are already divided by the type size and are in units of elements. // When using grouped query attention, Q is indexed by iq2, so the stride // should be nb02 (which is in bytes). - q_stride = p.gqa_ratio > 1 ? (p.nb02 / 4) : p.nb01; + q_stride = gqa_ratio > 1 ? (p.nb02 / 4) : p.nb01; k_stride = p.nb11; v_stride = p.nb21; - // When using grouped query attention, all rows use the same mask (stride 0). - // "p.gqa_ratio >> 16" is just a roundabout way of writing zero - // that prevents the compiler from folding the "&" through the select - // and breaking the alignment detection. - m_stride = (p.gqa_ratio > 1) ? (p.gqa_ratio >> 16) : KV; + // Bit 31 of gqa_ratio means "the mask row stride is in split_kv", used by the DeepSeek V4 + // sparse split path where the mask spans the full K range but this dispatch only covers the + // raw prefix, so m_stride != KV. That path always sets split_k_num == 1, which is what keeps + // split_kv free to carry the stride; ggml_vk_flash_attn_top_k asserts the invariant. + // Otherwise: when using grouped query attention all rows share the same mask (stride 0). + // "p.gqa_ratio >> 16" is just a roundabout way of writing zero that prevents the compiler + // from folding the "&" through the select and breaking the alignment detection. + const bool mask_stride_in_split_kv = (p.gqa_ratio & 0x80000000u) != 0; + m_stride = mask_stride_in_split_kv ? p.split_kv : ((gqa_ratio > 1) ? (p.gqa_ratio >> 16) : KV); + // Distinct from m_stride: m_stride is the row-to-row step INSIDE this tile (0 under GQA, + // where every row shares one mask row), while m_row_len is the mask tensor's actual row + // length, used to step between tokens and between streams. They differ under GQA and + // under the sparse split path, where this dispatch only covers the raw prefix (KV) but + // the mask rows span the whole K range. + m_row_len = mask_stride_in_split_kv ? p.split_kv : KV; } // Bias applied to softmax to stay in fp16 range. diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp index 057ed739aa8d..195901a7d90d 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp @@ -44,8 +44,11 @@ shared float tmpsh[row_split]; const uint32_t qstride = HSK_pad / 4 + 2; shared FLOAT_TYPEV4 Qf[Br * qstride]; -const uint psh_stride = Br / 4 + 2; -shared FLOAT_TYPEV4 Psh[Bc * psh_stride]; +// P is stored query-major with KV contiguous so the GEMM2 UseA load can be RowMajor: +// RADV flips the requested layout for UseA, and only the resulting internal column-major +// case gets an alignment hint, which is what lets the 16 element loads vectorize. +const uint psh_stride = Bc / 4 + 2; +shared FLOAT_TYPEV4 Psh[Br * psh_stride]; // Avoid padding for hsk==256 to make it fit in 48KB shmem. const uint32_t sfshstride = (HSK <= 128) ? (Br / 4 + 2) : Br / 4; @@ -139,9 +142,9 @@ void main() { // FaBlockBytesK/V == 2 for f16 (sizeof f16) and == 16 for f32 (vec4) and == ggml block size for quants. uint32_t k_offset = (ik2*p.nb12 + ik3*p.nb13) / FaBlockBytesK; uint32_t v_offset = (iv2*p.nb22 + iv3*p.nb23) / FaBlockBytesV; - uint32_t m_offset = gqa_iq1*KV; + uint32_t m_offset = gqa_iq1*m_row_len; if (p.nem2 != 1 || p.nem3 != 1) { - m_offset += ((iq3 % p.nem3) * p.nem2 + (iq2 % p.nem2)) * p.nem1 * KV; + m_offset += ((iq3 % p.nem3) * p.nem2 + (iq2 % p.nem2)) * p.nem1 * m_row_len; mo_offset += ((iq3 % p.nem3) * p.nem2 + (iq2 % p.nem2)) * CEIL_DIV(p.nem1, Br) * mo_stride; } @@ -169,7 +172,7 @@ void main() { } // Only load if the block is not all zeros if (mask_opt_bits != MASK_OPT_ALL_ZERO) { - bool nem1_bounds_check = !(p.gqa_ratio > 1) && (p.nem1 % Br) != 0; + bool nem1_bounds_check = !(gqa_ratio > 1) && (p.nem1 % Br) != 0; float max_mask = NEG_FLT_MAX_OVER_2; [[unroll]] for (uint32_t idx = 0; idx < Bc * Br / 4; idx += gl_WorkGroupSize.x) { @@ -382,15 +385,19 @@ void main() { [[unroll]] for (uint32_t r = 0; r < rows_per_thread; r += 4) { const uint row = tile_row(r); + const uint pcol_vec = col / 4; + const uint pcol_comp = col % 4; if (KV_bounds_check && j * Bc + col >= KV) { - Psh[col * psh_stride + row / 4] = FLOAT_TYPEV4(0.0f); + [[unroll]] for (uint32_t vec_idx = 0; vec_idx < 4; ++vec_idx) { + Psh[(row + vec_idx) * psh_stride + pcol_vec][pcol_comp] = FLOAT_TYPE(0.0f); + } } else { const vec4 mfvec = vec4(Mf[r], Mf[r + 1], Mf[r + 2], Mf[r + 3]); const FLOAT_TYPEV4 Pf = FLOAT_TYPEV4(exp(vec4(sfsh[row / 4 + col * sfshstride]) - mfvec)); [[unroll]] for (uint32_t vec_idx = 0; vec_idx < 4; ++vec_idx) { Lf[r + vec_idx] += Pf[vec_idx]; + Psh[(row + vec_idx) * psh_stride + pcol_vec][pcol_comp] = Pf[vec_idx]; } - Psh[col * psh_stride + row / 4] = Pf; } } } @@ -423,6 +430,13 @@ void main() { const uint num_hsv_tiles = (HSV + MatBc * row_split - 1) / (MatBc * row_split); // round up + // Psh is not written again until the next KV block, so the P fragments are the same + // for every hsv_tile. Load them once instead of re-reading LDS per tile. + coopmat PMat[Bc / MatBc]; + [[unroll]] for (uint32_t bc_chunk = 0; bc_chunk < Bc / MatBc; ++bc_chunk) { + coopMatLoad(PMat[bc_chunk], Psh, bc_chunk * (MatBc / 4), psh_stride, gl_CooperativeMatrixLayoutRowMajor); + } + // Each subgroup handles HSV/4 columns [[unroll]] for (uint32_t hsv_tile = 0; hsv_tile < num_hsv_tiles; ++hsv_tile) { const uint hsv_offset = (hsv_tile * row_split + gl_SubgroupID) * 16; @@ -476,8 +490,6 @@ void main() { if (hsv_offset < HSV_pad) { [[unroll]] for (uint32_t bc_chunk = 0; bc_chunk < Bc / MatBc; ++bc_chunk) { - coopMatLoad(KMat, Psh, bc_chunk * MatBc * psh_stride, psh_stride, gl_CooperativeMatrixLayoutColumnMajor); - if (SHMEM_STAGING == 0) { if (!USE_DECODE_V && !KV_bounds_check) { // F16/BF16 values can be loaded directly from global memory @@ -493,7 +505,7 @@ void main() { coopMatLoad(QMat, kvsh, v_tile_offset, kvsh_stride, gl_CooperativeMatrixLayoutRowMajor); } - PVMat = coopMatMulAdd(KMat, QMat, PVMat); + PVMat = coopMatMulAdd(PMat[bc_chunk], QMat, PVMat); } // Store PVMat to pvsh and load into Of @@ -533,10 +545,10 @@ void main() { // If there is split_k, then the split_k resolve shader does the final // division by L. Store the intermediate O value and per-row m and L values. - if (p.k_num > 1) { - if (p.gqa_ratio > 1) { + if (partial_output || split_k_num > 1) { + if (gqa_ratio > 1) { // note: O and Q have swapped coord 1,2. - uint32_t o_offset = HSV * p.ne1 * (split_k_index + p.k_num * (gqa_iq1 + p.ne2 * iq3)) / 4; + uint32_t o_offset = HSV * p.ne1 * (split_k_index + output_k_num * (gqa_iq1 + p.ne2 * iq3)) / 4; [[unroll]] for (uint32_t r = 0; r < rows_per_thread; ++r) { if (tile_row(r) < N) { @@ -549,7 +561,7 @@ void main() { } } - o_offset = HSV * p.ne1 * p.k_num * p.ne2 * p.ne3 + p.ne1 * 2 * (split_k_index + p.k_num * (gqa_iq1 + p.ne2 * iq3)); + o_offset = HSV * p.ne1 * output_k_num * p.ne2 * p.ne3 + p.ne1 * 2 * (split_k_index + output_k_num * (gqa_iq1 + p.ne2 * iq3)); [[unroll]] for (uint32_t r = 0; r < rows_per_thread; ++r) { if (tile_row(r) < N) { perElemOpStoreCol0(tile_row(r), 0u, ACC_TYPE(Lf[r]), o_offset, iq2, N); @@ -562,7 +574,7 @@ void main() { const uint global_row = i * Br + row; if (global_row < N) { - uint32_t o_offset = HSV * p.ne1 * (split_k_index + p.k_num * (global_row + p.ne2 * iq3)) / 4; + uint32_t o_offset = HSV * p.ne1 * (split_k_index + output_k_num * (global_row + p.ne2 * iq3)) / 4; [[unroll]] for (uint32_t d0 = 0; d0 < HSV / 4; d0 += threads_per_rowgroup) { const uint d = d0 + col_tid; @@ -572,7 +584,7 @@ void main() { } if (global_row < N && col_tid == 0) { - uint32_t lm_offset = HSV * p.ne1 * p.k_num * p.ne2 * p.ne3 + p.ne1 * 2 * (split_k_index + p.k_num * (global_row + p.ne2 * iq3)); + uint32_t lm_offset = HSV * p.ne1 * output_k_num * p.ne2 * p.ne3 + p.ne1 * 2 * (split_k_index + output_k_num * (global_row + p.ne2 * iq3)); data_o[lm_offset + iq2] = D_TYPE(Lf[r]); data_o[lm_offset + p.ne1 + iq2] = D_TYPE(Mf[r]); } @@ -621,7 +633,7 @@ void main() { uint32_t o_offset = (gqa_iq1*p.ne1*HSV + iq3*p.ne2*p.ne1*HSV) / 4; - if (p.gqa_ratio > 1) { + if (gqa_ratio > 1) { [[unroll]] for (uint32_t r = 0; r < rows_per_thread; ++r) { if (tile_row(r) < N) { [[unroll]] for (uint32_t d0 = 0; d0 < HSV / 4; d0 += threads_per_rowgroup) { diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm2.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm2.comp index 317411153087..54be1e6daa99 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm2.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm2.comp @@ -151,7 +151,7 @@ D_TYPE perElemOpNonGqaSplitKStore(const in uint32_t r, const in uint32_t c, cons uint32_t global_row = i * Br + r; if (global_row < N && c < HSV) { uint32_t o_off = HSV * p.ne1 - * (split_k_index + p.k_num * (global_row + p.ne2 * iq3)); + * (split_k_index + output_k_num * (global_row + p.ne2 * iq3)); data_o[o_off + iq2 * HSV + c] = D_TYPE(elem); } return elem; @@ -161,8 +161,8 @@ D_TYPE perElemOpNonGqaSplitKStore(const in uint32_t r, const in uint32_t c, cons ACC_TYPE perElemOpNonGqaSplitKStoreCol0(const in uint32_t r, const in uint32_t c, const in ACC_TYPE elem, const in uint32_t lm_base, const in uint32_t iq2, const in uint32_t N) { uint32_t global_row = i * Br + r; if (global_row < N && c == 0) { - uint32_t lm_off = HSV * p.ne1 * p.k_num * p.ne2 * p.ne3 - + p.ne1 * 2 * (split_k_index + p.k_num * (global_row + p.ne2 * iq3)); + uint32_t lm_off = HSV * p.ne1 * output_k_num * p.ne2 * p.ne3 + + p.ne1 * 2 * (split_k_index + output_k_num * (global_row + p.ne2 * iq3)); data_o[lm_off + lm_base + iq2] = D_TYPE(elem); } return elem; @@ -242,9 +242,9 @@ void main() { // mo_offset will point to the tile starting at row i*Br and col 0 uint32_t mo_offset = mo_stride * i; - uint32_t m_offset = gqa_iq1*KV * 2 /*sizeof(float16_t)*/; + uint32_t m_offset = gqa_iq1*m_row_len * 2 /*sizeof(float16_t)*/; if (p.nem2 != 1 || p.nem3 != 1) { - m_offset += ((iq3 % p.nem3) * p.nem2 + (iq2 % p.nem2)) * p.nem1 * KV * 2 /*sizeof(float16_t)*/; + m_offset += ((iq3 % p.nem3) * p.nem2 + (iq2 % p.nem2)) * p.nem1 * m_row_len * 2 /*sizeof(float16_t)*/; mo_offset += ((iq3 % p.nem3) * p.nem2 + (iq2 % p.nem2)) * CEIL_DIV(p.nem1, Br) * mo_stride; } @@ -268,7 +268,7 @@ void main() { } // Only load if the block is not all zeros if (mask_opt_bits != MASK_OPT_ALL_ZERO) { - bool nem1_bounds_check = !(p.gqa_ratio > 1) && (p.nem1 % Br) != 0; + bool nem1_bounds_check = !(gqa_ratio > 1) && (p.nem1 % Br) != 0; if (nem1_bounds_check) { tensorLayoutNV<2, gl_CooperativeMatrixClampModeConstantNV> tensorLayoutM = createTensorLayoutNV(2, gl_CooperativeMatrixClampModeConstantNV); @@ -287,7 +287,7 @@ void main() { } else { tensorLayoutNV<2, Clamp> tensorLayoutM = createTensorLayoutNV(2, Clamp); // Don't clamp against nem1 when GQA is enabled - uint32_t m_height = p.gqa_ratio > 1 ? ~0 : p.nem1; + uint32_t m_height = gqa_ratio > 1 ? ~0 : p.nem1; tensorLayoutM = setTensorLayoutDimensionNV(tensorLayoutM, m_height, KV); tensorLayoutM = setTensorLayoutStrideNV(tensorLayoutM, m_stride, 1); @@ -406,15 +406,15 @@ void main() { // If there is split_k, then the split_k resolve shader does the final // division by L. Store the intermediate O value and per-row m and L values. - if (p.k_num > 1) { + if (partial_output || split_k_num > 1) { coopmat O_D = coopmat(O); - if (p.gqa_ratio > 1) { + if (gqa_ratio > 1) { // note: O and Q have swapped coord 1,2. - uint32_t o_offset = HSV * p.ne1 * (split_k_index + p.k_num * (gqa_iq1 + p.ne2 * iq3)); + uint32_t o_offset = HSV * p.ne1 * (split_k_index + output_k_num * (gqa_iq1 + p.ne2 * iq3)); coopMatPerElementNV(O_D, O_D, perElemOpGqaStore, o_offset, iq2, N); - o_offset = HSV * p.ne1 * p.k_num * p.ne2 * p.ne3 + p.ne1 * 2 * (split_k_index + p.k_num * (gqa_iq1 + p.ne2 * iq3)); + o_offset = HSV * p.ne1 * output_k_num * p.ne2 * p.ne3 + p.ne1 * 2 * (split_k_index + output_k_num * (gqa_iq1 + p.ne2 * iq3)); coopMatPerElementNV(L, L, perElemOpStoreCol0, o_offset, iq2, N); coopMatPerElementNV(M, M, perElemOpStoreCol0, o_offset + p.ne1, iq2, N); } else { @@ -473,7 +473,7 @@ void main() { uint32_t o_offset = gqa_iq1*p.ne1*HSV + iq3*p.ne2*p.ne1*HSV; - if (p.gqa_ratio > 1) { + if (gqa_ratio > 1) { coopMatPerElementNV(O_D, O_D, perElemOpGqaStore, o_offset, iq2, N); } else { tensorLayoutNV<3, gl_CooperativeMatrixClampModeConstantNV> tensorLayoutD = createTensorLayoutNV(3, gl_CooperativeMatrixClampModeConstantNV); diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather.comp new file mode 100644 index 000000000000..4e16b2a9cece --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather.comp @@ -0,0 +1,117 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : require +#extension GL_EXT_shader_16bit_storage : require +#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require + +// Gathers the active KV rows of a top-k sparse attention (DeepSeek V4 CSA) into a compact +// contiguous scratch: rows [0, n_kv_raw) of the source (the dense prefix), then, for each of +// the n_batch query tokens, that token's n_top_k selected rows, then zero padding up to kv_c. +// +// n_batch > 1 (small-batch decode, e.g. speculative drafts) gives every token its own block +// rather than deduplicating into a union. That costs n_batch*n_top_k gathered rows instead of +// |union|, but needs no dedup pass and no atomics, and the size is bounded by the worst case +// the caller already allocates for. Cross-token rows are neutralised through the mask: token t +// sees -inf on any block that is not its own, so nothing is double counted in the softmax. The gathered mask row keeps the +// per-key mask values so causality/validity survive compaction; invalid top-k indices and +// padding get -inf mask and zeroed K (softmax-neutral either way, zeroed so no NaN*0). +// One workgroup per compact row; V is the K latent (V==K), so a single gather serves both. + +layout(local_size_x = 64, local_size_y = 1, local_size_z = 1) in; + +// K is addressed as raw 4-byte WORDS, never as values: a gather relocates rows and does not +// need to know whether they hold f16 elements or quantised blocks. Only the MASK is typed. +layout(binding = 0) readonly buffer KBuf { uint data_k[]; }; +layout(binding = 1) readonly buffer TopBuf { int data_top[]; }; +layout(binding = 2) readonly buffer MBuf { float16_t data_m[]; }; +layout(binding = 3) writeonly buffer KcBuf { uint data_kc[]; }; +layout(binding = 4) writeonly buffer McBuf { float16_t data_mc[]; }; + +layout(push_constant) uniform Parameters { + uint n_kv; // total source KV rows + uint n_kv_raw; // dense prefix length + uint n_top_k; // selected rows for the (single) query token + uint kv_c; // padded compact row count == dispatch row range + uint nbk1; // K source row stride, WORDS + uint nbk3; // K source stream stride, WORDS + uint nbt1; // top_k row (per query token) stride, elements + uint nbt3; // top_k stream stride, elements + uint nbm1; // mask source row (per query token) stride, elements + uint nbm3; // mask source stream stride, elements + uint nem3; // mask ne[3], for stream broadcast + uint n_batch; // query tokens sharing this gather; <= LANES + uint row_words; // bytes per K row / 4 +} p; + +const uint LANES = 64; + +void main() { + const uint row = gl_WorkGroupID.x; + const uint stream = gl_WorkGroupID.z; + const uint tid = gl_LocalInvocationIndex; + + // map compact row -> source row; p.n_kv is the invalid sentinel. + // owner is the token whose block this row belongs to, or ALL_TOKENS for the shared prefix. + const uint ALL_TOKENS = 0xffffffffu; + uint src = p.n_kv; + uint owner = ALL_TOKENS; + if (row < p.n_kv_raw) { + src = row; + } else { + const uint off = row - p.n_kv_raw; + const uint tok = off / p.n_top_k; + const uint slot = off - tok * p.n_top_k; + if (tok < p.n_batch) { + owner = tok; + const int idx = data_top[stream * p.nbt3 + tok * p.nbt1 + slot]; + if (idx >= 0 && uint(idx) < p.n_kv - p.n_kv_raw) { + src = p.n_kv_raw + uint(idx); + } + } + } + + // Zeroing writes zero BYTES, which decode to zero for every block-quantised type here (a + // zero scale zeroes the block) as well as for f16. The row is -inf in the mask either way; + // zeroing only keeps a garbage dot product from reaching the softmax as a NaN. + const uint dst_base = (stream * p.kv_c + row) * p.row_words; + // row_words is a push constant, so this cannot be [[unroll]]ed the way it was when the row + // was a fixed 512 f16 elements. Stepping 4*LANES recovers that: each access stays contiguous + // across the lanes, and an f16 row (256 words == 4*LANES) is one iteration with no loop + // overhead. The tail carries the quantised row sizes, which are not multiples of 4*LANES. + if (src < p.n_kv) { + const uint src_base = stream * p.nbk3 + src * p.nbk1; + uint i = tid; + for (; i + 3 * LANES < p.row_words; i += 4 * LANES) { + data_kc[dst_base + i] = data_k[src_base + i]; + data_kc[dst_base + i + LANES] = data_k[src_base + i + LANES]; + data_kc[dst_base + i + 2 * LANES] = data_k[src_base + i + 2 * LANES]; + data_kc[dst_base + i + 3 * LANES] = data_k[src_base + i + 3 * LANES]; + } + for (; i < p.row_words; i += LANES) { + data_kc[dst_base + i] = data_k[src_base + i]; + } + } else { + uint i = tid; + for (; i + 3 * LANES < p.row_words; i += 4 * LANES) { + data_kc[dst_base + i] = 0u; + data_kc[dst_base + i + LANES] = 0u; + data_kc[dst_base + i + 2 * LANES] = 0u; + data_kc[dst_base + i + 3 * LANES] = 0u; + } + for (; i < p.row_words; i += LANES) { + data_kc[dst_base + i] = 0u; + } + } + + // Compact mask is token-major [n_batch][kv_c], which is what the GQA mask path expects + // (m_stride 0, rows stepped by gqa_iq1 * m_row_len). One lane per token; n_batch <= LANES. + const float NEG_INF = uintBitsToFloat(0xff800000); + if (tid < p.n_batch) { + const uint mc_idx = (stream * p.n_batch + tid) * p.kv_c + row; + float mv = NEG_INF; + if (src < p.n_kv && (owner == ALL_TOKENS || owner == tid)) { + mv = float(data_m[(stream % p.nem3) * p.nbm3 + tid * p.nbm1 + src]); + } + data_mc[mc_idx] = float16_t(mv); + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_dq.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_dq.comp new file mode 100644 index 000000000000..f2711539641a --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_dq.comp @@ -0,0 +1,106 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : require +#extension GL_EXT_shader_16bit_storage : require +#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require +#extension GL_EXT_shader_explicit_arithmetic_types_int8 : require +#extension GL_EXT_shader_explicit_arithmetic_types_int16 : require +#extension GL_EXT_shader_explicit_arithmetic_types_int32 : require + +#define FLOAT_TYPE float +#include "types.glsl" + +// Gather + DEQUANTISE for the per-token compact form: same row mapping as the verbatim +// flash_attn_gather.comp, same decode as flash_attn_gather_union_dq.comp. One variant per K +// type, via DATA_A_*. +// +// The union carries the decode only for n_batch > 1, because it needs more than one token to +// deduplicate. Single-token decode is the common case and it lands here; a verbatim gather +// leaves flash attention to decode inline, measured 1.98x the f16 op at q8_0 and 2.26x at +// q4_0, kv 11008. + +layout(local_size_x = 64, local_size_y = 1, local_size_z = 1) in; + +layout(binding = 0) readonly buffer A { A_TYPE data_a[]; }; +layout(binding = 1) readonly buffer TopBuf { int data_top[]; }; +layout(binding = 2) readonly buffer MBuf { float16_t data_m[]; }; +layout(binding = 3) writeonly buffer KcBuf { float16_t data_kc[]; }; +layout(binding = 4) writeonly buffer McBuf { float16_t data_mc[]; }; + +layout(push_constant) uniform Parameters { + uint n_kv; + uint n_kv_raw; + uint n_top_k; + uint kv_c; + uint nbk1; // K source row stride, BLOCKS + uint nbk3; // K source stream stride, BLOCKS + uint nbt1; + uint nbt3; + uint nbm1; + uint nbm3; + uint nem3; + uint n_batch; + uint row_elems; // elements per K row; 64 lanes x 8 covers the 512 the gate pins +} p; + +void main() { + const uint row = gl_WorkGroupID.x; + const uint stream = gl_WorkGroupID.z; + const uint tid = gl_LocalInvocationIndex; + + const uint ALL_TOKENS = 0xffffffffu; + uint src = p.n_kv; + uint owner = ALL_TOKENS; + if (row < p.n_kv_raw) { + src = row; + } else { + const uint off = row - p.n_kv_raw; + const uint tok = off / p.n_top_k; + const uint slot = off - tok * p.n_top_k; + if (tok < p.n_batch) { + owner = tok; + const int idx = data_top[stream * p.nbt3 + tok * p.nbt1 + slot]; + if (idx >= 0 && uint(idx) < p.n_kv - p.n_kv_raw) { + src = p.n_kv_raw + uint(idx); + } + } + } + + const uint dst_base = (stream * p.kv_c + row) * p.row_elems; + const uint e0 = tid * 8; + if (src < p.n_kv) { + const uint ib = stream * p.nbk3 + src * p.nbk1 + e0 / QUANT_K; + const uint iqs = e0 % QUANT_K; + const float d = float(data_a[ib].d); +#if defined(DATA_A_Q8_0) + // element e is qs[e] + [[unroll]] for (uint l = 0; l < 8; ++l) { + data_kc[dst_base + e0 + l] = float16_t(d * float(data_a[ib].qs[iqs + l])); + } +#elif defined(DATA_A_Q4_0) + // byte j carries element j in its low nibble and element j+16 in its high one + const uint shift = (iqs >> 4) * 4; + const uint byte0 = iqs & 0xF; + [[unroll]] for (uint l = 0; l < 8; ++l) { + const float q = float((data_a[ib].qs[byte0 + l] >> shift) & 0xF) - 8.0f; + data_kc[dst_base + e0 + l] = float16_t(d * q); + } +#else +#error "no element mapping written for this K type" +#endif + } else { + [[unroll]] for (uint l = 0; l < 8; ++l) { + data_kc[dst_base + e0 + l] = float16_t(0.0); + } + } + + const float NEG_INF = uintBitsToFloat(0xff800000); + if (tid < p.n_batch) { + const uint mc_idx = (stream * p.n_batch + tid) * p.kv_c + row; + float mv = NEG_INF; + if (src < p.n_kv && (owner == ALL_TOKENS || owner == tid)) { + mv = float(data_m[(stream % p.nem3) * p.nbm3 + tid * p.nbm1 + src]); + } + data_mc[mc_idx] = float16_t(mv); + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union.comp new file mode 100644 index 000000000000..306372ba8af6 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union.comp @@ -0,0 +1,105 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : require +#extension GL_EXT_shader_16bit_storage : require +#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require + +// Gathers K/V and the mask for the DeepSeek V4 small-batch decode path, using the deduplicated +// union produced by flash_attn_union.comp: rows [0, n_kv_raw) of the source, then one row per +// distinct selected compressed row, then padding. +// +// Simpler than the block-per-token form it replaces, because a union row appears exactly once: +// there is no owner to track and no cross-token masking to apply. Each token just reads its own +// mask value for the gathered source row, which is already -inf where that token did not select +// it, so the softmax still cannot double count. +// +// Row count is a runtime value in data_c[0]; rows past the union are zeroed K and -inf mask, +// which is softmax-neutral and keeps the padded tail harmless. + +layout(local_size_x = 64, local_size_y = 1, local_size_z = 1) in; + +// K is addressed as raw 4-byte WORDS, never as values. A gather relocates rows; it does not +// need to know whether they hold f16 elements or quantised blocks, so the same shader serves +// every K type whose row is a whole number of words. Only the MASK is typed, and a mask is +// always f16. +layout(binding = 0) readonly buffer KBuf { uint data_k[]; }; +layout(binding = 1) readonly buffer UBuf { uint data_u[]; }; +layout(binding = 2) readonly buffer MBuf { float16_t data_m[]; }; +layout(binding = 3) writeonly buffer KcBuf { uint data_kc[]; }; +layout(binding = 4) writeonly buffer McBuf { float16_t data_mc[]; }; +layout(binding = 5) readonly buffer CBuf { uint data_c[]; }; + +layout(push_constant) uniform Parameters { + uint n_kv; + uint n_kv_raw; + uint kv_c_max; + uint nbk1; // K source row stride, WORDS + uint nbm1; + uint n_batch; + uint row_words; // bytes per K row / 4 +} p; + +const uint LANES = 64; + +void main() { + const uint row = gl_WorkGroupID.x; + const uint tid = gl_LocalInvocationIndex; + + const uint kv_c = data_c[0]; // padded compact rows, the FA's runtime KV + const uint n_uni = data_c[1]; // unpadded union size + + if (row >= kv_c) { + return; + } + + uint src = p.n_kv; // sentinel: invalid + if (row < p.n_kv_raw) { + src = row; + } else if (row - p.n_kv_raw < n_uni) { + src = p.n_kv_raw + data_u[row - p.n_kv_raw]; + } + + // Zeroing an unused row writes zero BYTES, which decode to zero for every block-quantised + // type here (a zero scale zeroes the block) as well as for f16. The row is -inf in the mask + // either way; zeroing only keeps a garbage dot product from reaching the softmax as a NaN. + const uint dst_base = row * p.row_words; + // Stepped 4*LANES for the same reason as flash_attn_gather.comp: row_words is a push + // constant, so the copy cannot be [[unroll]]ed, and an f16 row is one iteration this way. + if (src < p.n_kv) { + const uint src_base = src * p.nbk1; + uint i = tid; + for (; i + 3 * LANES < p.row_words; i += 4 * LANES) { + data_kc[dst_base + i] = data_k[src_base + i]; + data_kc[dst_base + i + LANES] = data_k[src_base + i + LANES]; + data_kc[dst_base + i + 2 * LANES] = data_k[src_base + i + 2 * LANES]; + data_kc[dst_base + i + 3 * LANES] = data_k[src_base + i + 3 * LANES]; + } + for (; i < p.row_words; i += LANES) { + data_kc[dst_base + i] = data_k[src_base + i]; + } + } else { + uint i = tid; + for (; i + 3 * LANES < p.row_words; i += 4 * LANES) { + data_kc[dst_base + i] = 0u; + data_kc[dst_base + i + LANES] = 0u; + data_kc[dst_base + i + 2 * LANES] = 0u; + data_kc[dst_base + i + 3 * LANES] = 0u; + } + for (; i < p.row_words; i += LANES) { + data_kc[dst_base + i] = 0u; + } + } + + // Compact mask is token-major [n_batch][kv_c]. The stride must be the RUNTIME kv_c, not + // kv_c_max: the FA derives m_row_len from KV, which is now that same runtime value, and a + // mismatch would step the mask by the wrong amount for every token past the first. The + // buffer is allocated for kv_c_max, so a smaller stride simply leaves a tail unused. + const float NEG_INF = uintBitsToFloat(0xff800000); + if (tid < p.n_batch) { + float mv = NEG_INF; + if (src < p.n_kv) { + mv = float(data_m[tid * p.nbm1 + src]); + } + data_mc[tid * kv_c + row] = float16_t(mv); + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union_dq.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union_dq.comp new file mode 100644 index 000000000000..fd3d0708b2f0 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_gather_union_dq.comp @@ -0,0 +1,123 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : require +#extension GL_EXT_shader_16bit_storage : require +#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require +#extension GL_EXT_shader_explicit_arithmetic_types_int8 : require +#extension GL_EXT_shader_explicit_arithmetic_types_int16 : require +#extension GL_EXT_shader_explicit_arithmetic_types_int32 : require + +// iq4_nl's value table lives in shared memory and is typed FLOAT_TYPE +#define FLOAT_TYPE float +#include "types.glsl" + +// Gather + DEQUANTISE in one pass, for the DeepSeek V4 small-batch decode path with a +// quantised KV cache. One variant per K type, via DATA_A_*. +// +// The sibling flash_attn_gather_union.comp relocates rows verbatim, which leaves flash +// attention to decode them in its inner loop - once per query block that reads a row, rather +// than once per row. Measured, that is a constant per KV row attended: +// +// q8_0 0.148 - 0.153 us/row q4_0 0.172 - 0.179 us/row +// +// flat across batch 2/4/8 and across the dense and gathered regimes alike. It is what made +// quantised attention slower than f16 despite reading fewer bytes, and q4_0 worse than q8_0 +// despite reading fewer bytes still: unpacking nibbles costs more ALU than the bytes save. +// +// The gather already touches each selected row exactly once, so decoding here converts per-use +// work into per-row work. The compact scratch becomes f16 and flash attention takes its f16 +// path. The pass itself already existed, so the only added cost is writing f16 rather than +// blocks into the scratch. + +layout(local_size_x = 64, local_size_y = 1, local_size_z = 1) in; + +layout (binding = 0) readonly buffer A {A_TYPE data_a[];}; + +layout(binding = 1) readonly buffer UBuf { uint data_u[]; }; +layout(binding = 2) readonly buffer MBuf { float16_t data_m[]; }; +layout(binding = 3) writeonly buffer KcBuf { float16_t data_kc[]; }; +layout(binding = 4) writeonly buffer McBuf { float16_t data_mc[]; }; +layout(binding = 5) readonly buffer CBuf { uint data_c[]; }; + +layout(push_constant) uniform Parameters { + uint n_kv; + uint n_kv_raw; + uint kv_c_max; + uint nbk1; // K source row stride, BLOCKS + uint nbm1; + uint n_batch; + uint row_elems; // elements per K row (== head size) +} p; + +const uint LANES = 64; + +void main() { +#ifdef NEEDS_INIT_IQ_SHMEM + // barrier inside; must run before any divergent return below + init_iq_shmem(gl_WorkGroupSize); +#endif + const uint row = gl_WorkGroupID.x; + const uint tid = gl_LocalInvocationIndex; + + const uint kv_c = data_c[0]; // padded compact rows, the FA's runtime KV + const uint n_uni = data_c[1]; // unpadded union size + + if (row >= kv_c) { + return; + } + + uint src = p.n_kv; // sentinel: invalid + if (row < p.n_kv_raw) { + src = row; + } else if (row - p.n_kv_raw < n_uni) { + src = p.n_kv_raw + data_u[row - p.n_kv_raw]; + } + + // 64 lanes x 8 elements covers a 512-element row exactly, and QUANT_K is 32, so a lane's 8 + // elements sit wholly inside one block AND wholly inside one nibble half. + // + // Deliberately NOT dequantize4() from dequant_funcs.glsl. That helper is permutation + // AGNOSTIC: mul_mat_vec only ever feeds it into a dot product, where any consistent element + // order gives the same answer, so for q4_0 it returns nibbles in packed order rather than + // element order. Materialising a row to memory is the one use where the order matters, and + // using it here produced NMSE 1.17 on q4_0 while q8_0 passed - q8_0's layout just happens to + // be contiguous. Each type's true element mapping is written out below instead. + const uint dst_base = row * p.row_elems; + const uint e0 = tid * 8; + if (src < p.n_kv) { + const uint ib = src * p.nbk1 + e0 / QUANT_K; + const uint iqs = e0 % QUANT_K; + const float d = float(data_a[ib].d); +#if defined(DATA_A_Q8_0) + // element e is qs[e] + [[unroll]] for (uint l = 0; l < 8; ++l) { + data_kc[dst_base + e0 + l] = float16_t(d * float(data_a[ib].qs[iqs + l])); + } +#elif defined(DATA_A_Q4_0) + // byte j carries element j in its low nibble and element j+16 in its high one + const uint shift = (iqs >> 4) * 4; // 0 for elements 0..15, 4 for 16..31 + const uint byte0 = iqs & 0xF; + [[unroll]] for (uint l = 0; l < 8; ++l) { + const float q = float((data_a[ib].qs[byte0 + l] >> shift) & 0xF) - 8.0f; + data_kc[dst_base + e0 + l] = float16_t(d * q); + } +#else +#error "no element mapping written for this K type" +#endif + } else { + [[unroll]] for (uint j = 0; j < 8; ++j) { + data_kc[dst_base + e0 + j] = float16_t(0.0); + } + } + + // Compact mask is token-major [n_batch][kv_c], stride the RUNTIME kv_c: the FA derives + // m_row_len from KV, which is that same runtime value. + const float NEG_INF = uintBitsToFloat(0xff800000); + if (tid < p.n_batch) { + float mv = NEG_INF; + if (src < p.n_kv) { + mv = float(data_m[tid * p.nbm1 + src]); + } + data_mc[tid * kv_c + row] = float16_t(mv); + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_split_k_reduce.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_split_k_reduce.comp index 68917fc0bb02..69342ba3ceea 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_split_k_reduce.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_split_k_reduce.comp @@ -14,6 +14,7 @@ layout (push_constant) uniform parameter { uint D; uint ne1; uint ne2; + uint dst_ne2; uint ne3; uint k_num; uint sinks; @@ -116,6 +117,6 @@ void main() { const float FLT_MAX = uintBitsToFloat(0x7F7FFFFF); O = clamp(O, -FLT_MAX, FLT_MAX); - data_d[(i3 * p.ne2 + i2) * p.ne1 * D + D * n + d] = O; + data_d[(i3 * p.dst_ne2 + i2) * p.ne1 * D + D * n + d] = O; } } diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k.comp new file mode 100644 index 000000000000..930777b5f0ea --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k.comp @@ -0,0 +1,155 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : require +#extension GL_EXT_shader_16bit_storage : require +#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require +#extension GL_KHR_shader_subgroup_arithmetic : require +#extension GL_KHR_shader_subgroup_basic : require + +layout(constant_id = 0) const uint WORKGROUP_SIZE = 512; +layout(constant_id = 1) const uint SUBGROUP_SIZE = 64; +layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; + +layout(binding = 0) readonly buffer QBuf { float data_q[]; }; +layout(binding = 1) readonly buffer KBuf { float16_t data_k[]; }; +layout(binding = 2) readonly buffer MaskBuf { float16_t data_m[]; }; +layout(binding = 3) readonly buffer SinkBuf { float data_s[]; }; +layout(binding = 4) readonly buffer TopBuf { int data_top[]; }; +layout(binding = 5) writeonly buffer DstBuf { float data_dst[]; }; + +layout(push_constant) uniform Parameters { + uint n_batch; + uint n_kv; + uint n_kv_raw; + uint n_top_k; + uint n_head; + uint nbq1; + uint nbq2; + uint nbq3; + uint nbk1; + uint nbk3; + uint nbm1; + uint nbm3; + uint nbt1; + uint nbt3; + uint nb1; + uint nb2; + uint nb3; + float scale; + uint has_sinks; + uint split_mode; +} p; + +// Shape constants pinned by the dispatch gate in ggml_vk_flash_attn_top_k: DeepSeek V4 +// CSA attention only (hd 512, 64 heads MQA over an f16 K==V latent). The gate also pins +// the pipeline to subgroup size 64, so per-lane sizing uses LANES rather than the +// SUBGROUP_SIZE specialization constant (a spec-constant array bound would be folded at +// compile time from the default and silently break under a different runtime subgroup). +const uint HEAD_SIZE = 512; +const uint HEADS_PER_GROUP = 8; +const uint KEYS_PER_BLOCK = 16; +const uint LANES = 64; +// f16 -inf (or the lowest f16 normal some mask writers use in its place) marks a +// masked-out key. +const float MASK_NEG_INF = -65500.0; + +shared float16_t key_sh[KEYS_PER_BLOCK * HEAD_SIZE]; +shared uint key_idx[KEYS_PER_BLOCK]; + +void main() { + const uint tid = gl_LocalInvocationIndex; + const uint lane = gl_SubgroupInvocationID; + const uint head = gl_WorkGroupID.y * HEADS_PER_GROUP + gl_SubgroupID; + const uint token = gl_WorkGroupID.x; + const uint stream = gl_WorkGroupID.z; + + // No bounds check: the grid is sized exactly (x = n_batch, y * HEADS_PER_GROUP covers + // n_head == 64). An early return here would skip the barriers below for part of the + // workgroup, which is undefined behavior — do not reintroduce one without restructuring + // the barrier flow. + + float accum[HEAD_SIZE / LANES]; + [[unroll]] for (uint i = 0; i < HEAD_SIZE / LANES; ++i) { + accum[i] = 0.0; + } + + float row_max = uintBitsToFloat(0xff800000); + float row_sum = 0.0; + const uint q_base = stream * p.nbq3 + head * p.nbq2 + token * p.nbq1; + const uint mask_base = stream * p.nbm3 + token * p.nbm1; + const uint top_base = stream * p.nbt3 + token * p.nbt1; + const uint total_keys = p.n_kv_raw + p.n_top_k; + + for (uint kb = 0; kb < total_keys; kb += KEYS_PER_BLOCK) { + if (tid < KEYS_PER_BLOCK) { + const uint selected = kb + tid; + uint key = p.n_kv; + if (selected < p.n_kv_raw) { + key = selected; + } else if (selected < total_keys) { + const int compressed = data_top[top_base + selected - p.n_kv_raw]; + if (compressed >= 0 && uint(compressed) < p.n_kv - p.n_kv_raw) { + key = p.n_kv_raw + uint(compressed); + } + } + key_idx[tid] = key; + } + barrier(); + + for (uint idx = tid; idx < KEYS_PER_BLOCK * HEAD_SIZE; idx += WORKGROUP_SIZE) { + const uint col = idx / HEAD_SIZE; + const uint dim = idx % HEAD_SIZE; + const uint key = key_idx[col]; + key_sh[idx] = key < p.n_kv ? data_k[stream * p.nbk3 + key * p.nbk1 + dim] : float16_t(0.0); + } + barrier(); + + [[unroll]] for (uint col = 0; col < KEYS_PER_BLOCK; ++col) { + const uint selected = kb + col; + const uint key = key_idx[col]; + if (selected >= total_keys || key >= p.n_kv) { + continue; + } + + float partial = 0.0; + [[unroll]] for (uint i = 0; i < HEAD_SIZE / LANES; ++i) { + const uint dim = lane + i * LANES; + partial += data_q[q_base + dim] * float(key_sh[col * HEAD_SIZE + dim]); + } + const float mask = float(data_m[mask_base + key]); + const float score = subgroupAdd(partial) * p.scale + mask; + if (mask < MASK_NEG_INF) { + continue; + } + + const float new_max = max(row_max, score); + const float old_scale = row_sum == 0.0 ? 0.0 : exp(row_max - new_max); + const float value_scale = exp(score - new_max); + row_sum = row_sum * old_scale + value_scale; + row_max = new_max; + + [[unroll]] for (uint i = 0; i < HEAD_SIZE / LANES; ++i) { + const uint dim = lane + i * LANES; + accum[i] = accum[i] * old_scale + value_scale * float(key_sh[col * HEAD_SIZE + dim]); + } + } + barrier(); + } + + if (p.has_sinks != 0) { + const float sink = data_s[head]; + const float new_max = max(row_max, sink); + const float old_scale = row_sum == 0.0 ? 0.0 : exp(row_max - new_max); + const float sink_scale = exp(sink - new_max); + row_sum = row_sum * old_scale + sink_scale; + [[unroll]] for (uint i = 0; i < HEAD_SIZE / LANES; ++i) { + accum[i] *= old_scale; + } + } + + const uint dst_base = stream * p.nb3 + token * p.nb2 + head * p.nb1; + const float inv_sum = row_sum == 0.0 ? 0.0 : 1.0 / row_sum; + [[unroll]] for (uint i = 0; i < HEAD_SIZE / LANES; ++i) { + data_dst[dst_base + lane + i * LANES] = accum[i] * inv_sum; + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp new file mode 100644 index 000000000000..7a328938777a --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_top_k_cm.comp @@ -0,0 +1,296 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : require +#extension GL_EXT_shader_16bit_storage : require +#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require +#extension GL_KHR_cooperative_matrix : require +#extension GL_KHR_memory_scope_semantics : require +#extension GL_KHR_shader_subgroup_basic : require +#extension GL_KHR_shader_subgroup_shuffle : require + +layout(constant_id = 0) const uint WORKGROUP_SIZE = 512; +layout(constant_id = 1) const uint SUBGROUP_SIZE = 64; +layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; + +layout(binding = 0) readonly buffer QBuf { float data_q[]; }; +layout(binding = 1) readonly buffer KBuf { float16_t data_k[]; }; +layout(binding = 2) readonly buffer MaskBuf { float16_t data_m[]; }; +layout(binding = 3) readonly buffer SinkBuf { float data_s[]; }; +layout(binding = 4) readonly buffer TopBuf { int data_top[]; }; +layout(binding = 5) writeonly buffer DstBuf { float data_dst[]; }; + +layout(push_constant) uniform Parameters { + uint n_batch; + uint n_kv; + uint n_kv_raw; + uint n_top_k; + uint n_head; + uint nbq1; + uint nbq2; + uint nbq3; + uint nbk1; + uint nbk3; + uint nbm1; + uint nbm3; + uint nbt1; + uint nbt3; + uint nb1; + uint nb2; + uint nb3; + float scale; + uint has_sinks; + uint split_mode; +} p; + +// This shader is hard-wired to a 512-thread workgroup of eight 64-wide subgroups: +// the softmax segments SUBGROUP_SIZE/4 = 16 lanes per head (so head_local stays < 32), +// and the QK/PV tiling assumes gl_NumSubgroups == 8 via gl_SubgroupID % 4 / gl_SubgroupID / 4. +// A different subgroup size indexes past row_max_sh/q_p_sh and corrupts the score tiles, so +// the pipeline is only created under `device->subgroup_size == 64` in ggml_vk_load_shaders. +// Keep that gate in sync with these constants. +const uint TILE = 16; +const uint HEAD_SIZE = 512; +const uint HEADS_PER_GROUP = 32; +const uint KEYS_PER_BLOCK = 64; +const uint DIMS_PER_BLOCK = 64; +const uint QK_STRIDE = TILE / 4 + 2; +const uint SCORE_STRIDE = HEADS_PER_GROUP / 4 + 1; +const uint P_STRIDE = KEYS_PER_BLOCK / 4 + 2; +const uint V_STRIDE = DIMS_PER_BLOCK / 4 + 2; +const uint PV_STRIDE = DIMS_PER_BLOCK / 4; +const float MASK_NEG_INF = -65500.0; + +shared uint key_idx[KEYS_PER_BLOCK]; +shared float key_mask[KEYS_PER_BLOCK]; +shared f16vec4 q_p_sh[HEADS_PER_GROUP * P_STRIDE]; +shared f16vec4 k_v_sh[KEYS_PER_BLOCK * V_STRIDE]; +shared vec4 matrix_sh[KEYS_PER_BLOCK * SCORE_STRIDE]; +shared float old_scale_sh[HEADS_PER_GROUP]; +shared float row_max_sh[HEADS_PER_GROUP]; +shared float row_sum_sh[HEADS_PER_GROUP]; + +void main() { + const uint tid = gl_LocalInvocationIndex; + const uint token = gl_WorkGroupID.x; + const uint head_base = gl_WorkGroupID.y * HEADS_PER_GROUP; + const uint stream = gl_WorkGroupID.z; + const uint mask_base = stream * p.nbm3 + token * p.nbm1; + const uint top_base = stream * p.nbt3 + token * p.nbt1; + const uint total_keys = p.split_mode != 0 ? p.n_top_k : p.n_kv_raw + p.n_top_k; + + float accum[HEADS_PER_GROUP * HEAD_SIZE / WORKGROUP_SIZE]; + [[unroll]] for (uint i = 0; i < accum.length(); ++i) { + accum[i] = 0.0; + } + + if (tid < HEADS_PER_GROUP) { + row_max_sh[tid] = uintBitsToFloat(0xff800000); + row_sum_sh[tid] = 0.0; + } + barrier(); + + for (uint kb = 0; kb < total_keys; kb += KEYS_PER_BLOCK) { + if (tid < KEYS_PER_BLOCK) { + const uint selected = kb + tid; + uint key = p.n_kv; + if (p.split_mode == 0 && selected < p.n_kv_raw) { + key = selected; + } else if (selected < total_keys) { + const uint top_pos = p.split_mode != 0 ? selected : selected - p.n_kv_raw; + const int compressed = data_top[top_base + top_pos]; + if (compressed >= 0 && uint(compressed) < p.n_kv - p.n_kv_raw) { + key = p.n_kv_raw + uint(compressed); + } + } + key_idx[tid] = key; + key_mask[tid] = key < p.n_kv ? float(data_m[mask_base + key]) : uintBitsToFloat(0xff800000); + } + barrier(); + + coopmat scores = + coopmat(0.0); + coopmat kmat; + coopmat qmat; + + [[unroll]] for (uint d = 0; d < HEAD_SIZE; d += TILE) { + if (tid < KEYS_PER_BLOCK * (TILE / 4)) { + const uint key_local = tid / (TILE / 4); + const uint d4 = tid % (TILE / 4); + const uint key = key_idx[key_local]; + f16vec4 value = f16vec4(0.0); + if (key < p.n_kv) { + const uint offset = stream * p.nbk3 + key * p.nbk1 + d + d4 * 4; + value = f16vec4(data_k[offset], data_k[offset + 1], data_k[offset + 2], data_k[offset + 3]); + } + k_v_sh[key_local * QK_STRIDE + d4] = value; + } + if (tid < HEADS_PER_GROUP * (TILE / 4)) { + const uint head_local = tid / (TILE / 4); + const uint d4 = tid % (TILE / 4); + const uint head = head_base + head_local; + const uint offset = stream * p.nbq3 + head * p.nbq2 + token * p.nbq1 + d + d4 * 4; + q_p_sh[head_local * QK_STRIDE + d4] = f16vec4( + data_q[offset], data_q[offset + 1], data_q[offset + 2], data_q[offset + 3]); + } + barrier(); + + const uint key_chunk = gl_SubgroupID % (KEYS_PER_BLOCK / TILE); + const uint head_tile = gl_SubgroupID / (KEYS_PER_BLOCK / TILE); + coopMatLoad(kmat, k_v_sh, key_chunk * TILE * QK_STRIDE, + QK_STRIDE, gl_CooperativeMatrixLayoutRowMajor); + coopMatLoad(qmat, q_p_sh, head_tile * TILE * QK_STRIDE, + QK_STRIDE, gl_CooperativeMatrixLayoutColumnMajor); + scores = coopMatMulAdd(kmat, qmat, scores); + barrier(); + } + + const uint score_key_chunk = gl_SubgroupID % (KEYS_PER_BLOCK / TILE); + const uint score_head_tile = gl_SubgroupID / (KEYS_PER_BLOCK / TILE); + coopMatStore(scores, matrix_sh, + score_key_chunk * TILE * SCORE_STRIDE + score_head_tile * (TILE / 4), + SCORE_STRIDE, gl_CooperativeMatrixLayoutRowMajor); + barrier(); + + { + const uint head_local = tid / (SUBGROUP_SIZE / 4); + const uint softmax_lane = tid % (SUBGROUP_SIZE / 4); + float block_max = uintBitsToFloat(0xff800000); + [[unroll]] for (uint i = 0; i < KEYS_PER_BLOCK / (SUBGROUP_SIZE / 4); ++i) { + const uint key_local = softmax_lane + i * (SUBGROUP_SIZE / 4); + const float mask = key_mask[key_local]; + const float score = float(matrix_sh[key_local * SCORE_STRIDE + head_local / 4][head_local % 4]) * p.scale + mask; + block_max = mask < MASK_NEG_INF ? block_max : max(block_max, score); + } + [[unroll]] for (uint delta = 1; delta < SUBGROUP_SIZE / 4; delta *= 2) { + block_max = max(block_max, subgroupShuffleXor(block_max, delta)); + } + + const float old_row_max = row_max_sh[head_local]; + const float old_row_sum = row_sum_sh[head_local]; + const float new_max = max(old_row_max, block_max); + const float old_scale = old_row_sum == 0.0 ? 0.0 : exp(old_row_max - new_max); + float block_sum = 0.0; + [[unroll]] for (uint i = 0; i < KEYS_PER_BLOCK / (SUBGROUP_SIZE / 4); ++i) { + const uint key_local = softmax_lane + i * (SUBGROUP_SIZE / 4); + const float mask = key_mask[key_local]; + float weight = 0.0; + if (mask >= MASK_NEG_INF) { + const float score = float(matrix_sh[key_local * SCORE_STRIDE + head_local / 4][head_local % 4]) * p.scale + mask; + weight = exp(score - new_max); + block_sum += weight; + } + q_p_sh[head_local * P_STRIDE + key_local / 4][key_local % 4] = float16_t(weight); + } + [[unroll]] for (uint delta = 1; delta < SUBGROUP_SIZE / 4; delta *= 2) { + block_sum += subgroupShuffleXor(block_sum, delta); + } + + if (softmax_lane == 0) { + row_sum_sh[head_local] = old_row_sum * old_scale + block_sum; + row_max_sh[head_local] = new_max; + old_scale_sh[head_local] = old_scale; + } + } + barrier(); + + [[unroll]] for (uint i = 0; i < accum.length(); ++i) { + const uint out_idx = tid + i * WORKGROUP_SIZE; + const uint head_local = out_idx / HEAD_SIZE; + accum[i] *= old_scale_sh[head_local]; + } + + coopmat pmats[KEYS_PER_BLOCK / TILE]; + [[unroll]] for (uint key_chunk = 0; key_chunk < KEYS_PER_BLOCK / TILE; ++key_chunk) { + const uint pv_head_tile = gl_SubgroupID / (DIMS_PER_BLOCK / TILE); + coopMatLoad(pmats[key_chunk], q_p_sh, + pv_head_tile * TILE * P_STRIDE + key_chunk * (TILE / 4), + P_STRIDE, gl_CooperativeMatrixLayoutRowMajor); + } + + [[unroll]] for (uint dim_base = 0; dim_base < HEAD_SIZE; dim_base += DIMS_PER_BLOCK) { + [[unroll]] for (uint idx = tid; idx < KEYS_PER_BLOCK * (DIMS_PER_BLOCK / 4); idx += WORKGROUP_SIZE) { + const uint key_local = idx / (DIMS_PER_BLOCK / 4); + const uint d4 = idx % (DIMS_PER_BLOCK / 4); + const uint key = key_idx[key_local]; + f16vec4 value = f16vec4(0.0); + if (key < p.n_kv) { + const uint offset = stream * p.nbk3 + key * p.nbk1 + dim_base + d4 * 4; + value = f16vec4(data_k[offset], data_k[offset + 1], data_k[offset + 2], data_k[offset + 3]); + } + k_v_sh[key_local * V_STRIDE + d4] = value; + } + barrier(); + + coopmat pv = + coopmat(0.0); + coopmat vmat; + + const uint pv_head_tile = gl_SubgroupID / (DIMS_PER_BLOCK / TILE); + const uint pv_dim_tile = gl_SubgroupID % (DIMS_PER_BLOCK / TILE); + [[unroll]] for (uint key_chunk = 0; key_chunk < KEYS_PER_BLOCK / TILE; ++key_chunk) { + coopMatLoad(vmat, k_v_sh, + key_chunk * TILE * V_STRIDE + pv_dim_tile * (TILE / 4), + V_STRIDE, gl_CooperativeMatrixLayoutRowMajor); + pv = coopMatMulAdd(pmats[key_chunk], vmat, pv); + } + + coopMatStore(pv, matrix_sh, + pv_head_tile * TILE * PV_STRIDE + pv_dim_tile * (TILE / 4), PV_STRIDE, + gl_CooperativeMatrixLayoutRowMajor); + barrier(); + + [[unroll]] for (uint i = 0; i < accum.length(); ++i) { + const uint out_idx = tid + i * WORKGROUP_SIZE; + const uint head_local = out_idx / HEAD_SIZE; + const uint dim = out_idx % HEAD_SIZE; + if (dim >= dim_base && dim < dim_base + DIMS_PER_BLOCK) { + accum[i] += matrix_sh[head_local * PV_STRIDE + (dim - dim_base) / 4][dim % 4]; + } + } + barrier(); + } + } + + if (p.split_mode == 0 && p.has_sinks != 0 && tid < HEADS_PER_GROUP) { + const float sink = data_s[head_base + tid]; + const float new_max = max(row_max_sh[tid], sink); + const float old_scale = row_sum_sh[tid] == 0.0 ? 0.0 : exp(row_max_sh[tid] - new_max); + row_sum_sh[tid] = row_sum_sh[tid] * old_scale + exp(sink - new_max); + row_max_sh[tid] = new_max; + old_scale_sh[tid] = old_scale; + } + barrier(); + + if (p.split_mode == 0 && p.has_sinks != 0) { + [[unroll]] for (uint i = 0; i < accum.length(); ++i) { + const uint out_idx = tid + i * WORKGROUP_SIZE; + accum[i] *= old_scale_sh[out_idx / HEAD_SIZE]; + } + } + [[unroll]] for (uint i = 0; i < accum.length(); ++i) { + const uint out_idx = tid + i * WORKGROUP_SIZE; + const uint head_local = out_idx / HEAD_SIZE; + const uint dim = out_idx % HEAD_SIZE; + if (p.split_mode != 0) { + const uint part_idx = 1; + // Split-K buffer layout, shared with flash_attn_cm1.comp and + // flash_attn_split_k_reduce.comp: all O matrices [HSV, ne1, k_num, ne2, ne3] + // first, then the L/M rows [ne1, k_num, ne2, ne3]. Here ne1 = n_head, + // ne2 = n_batch and ne3 = the stream count, so the L/M base must span every + // stream -- omitting gl_NumWorkGroups.z lands L/M inside the O region as soon + // as there is more than one sequence. + const uint n_streams = gl_NumWorkGroups.z; + const uint matrix_base = HEAD_SIZE * p.n_head * (part_idx + 2 * (token + p.n_batch * stream)); + const uint lm_base = HEAD_SIZE * p.n_head * p.n_batch * n_streams * 2 + p.n_head * 2 * (part_idx + 2 * (token + p.n_batch * stream)); + data_dst[matrix_base + (head_base + head_local) * HEAD_SIZE + dim] = accum[i]; + if (dim == 0) { + data_dst[lm_base + head_base + head_local] = row_sum_sh[head_local]; + data_dst[lm_base + p.n_head + head_base + head_local] = row_max_sh[head_local]; + } + } else { + const uint dst_base = stream * p.nb3 + token * p.nb2 + (head_base + head_local) * p.nb1; + const float inv_sum = row_sum_sh[head_local] == 0.0 ? 0.0 : 1.0 / row_sum_sh[head_local]; + data_dst[dst_base + dim] = accum[i] * inv_sum; + } + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_union.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_union.comp new file mode 100644 index 000000000000..513375709300 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_union.comp @@ -0,0 +1,134 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : require +#extension GL_KHR_shader_subgroup_arithmetic : require +#extension GL_KHR_shader_subgroup_basic : require + +// Builds the deduplicated union of the compressed rows selected by any of the n_batch query +// tokens, for the DeepSeek V4 small-batch decode path. +// +// Marks a bitmap from the top-k index lists, then compacts by scanning bitmap WORDS. The +// marking is O(n_batch * n_top_k), independent of context depth, and the compaction touches +// R/32 words instead of R rows. An earlier version scanned the mask row by row instead: that +// is simpler, but it made dedup cost scale with depth and measured 0.71x (ie a loss) at +// kv=133888, which defeats the purpose. Do not go back to it. +// +// Ascending source order falls out of the bitmap scan, so the result is deterministic. +// +// Emits the index list plus the PADDED compact row count. Padding to a multiple of pad_to +// (a multiple of every FA block width) is what lets flash-attention keep its "aligned" +// pipeline variant even though the row count is now a runtime value. +// +// Also emits the raw union size, the candidate count and the batch it came from. The host +// reads those back to decide whether compaction is worth it at all, and files the sample under +// the batch reported here rather than the one it is pricing: it cannot tell which call last +// wrote the slot, and the overlap it is measuring depends on the batch. +// +// count_only skips the list writes. The scan still has to run to produce the count, but with +// nothing to write the index buffer need not exist, which is what lets the host price the +// union on a call it is not going to compact. +// +// One workgroup: the bitmap and the running offset both live in shared memory. + +layout(local_size_x = 256, local_size_y = 1, local_size_z = 1) in; + +layout(binding = 0) readonly buffer TBuf { int data_top[]; }; +layout(binding = 1) writeonly buffer UBuf { uint data_u[]; }; +layout(binding = 2) writeonly buffer CBuf { uint data_c[]; }; + +layout(push_constant) uniform Parameters { + uint n_kv; + uint n_kv_raw; + uint n_batch; + uint n_top_k; + uint max_union; + uint nbt1; + uint max_words; // capacity of the shared bitmap, host-checked + uint pad_to; + uint count_only; // 1: produce the count, write no index list +} p; + +// 12288 words = 393216 compressed rows, ~48 KiB of shared memory. +const uint MAX_WORDS = 12288; + +shared uint bitmap[MAX_WORDS]; +shared uint wave_totals[16]; +shared uint base_sh; + +void main() { + const uint tid = gl_LocalInvocationIndex; + const uint lane = gl_SubgroupInvocationID; + const uint wave = gl_SubgroupID; + const uint nwave = gl_NumSubgroups; + const uint R = p.n_kv - p.n_kv_raw; + const uint words = min((R + 31) / 32, p.max_words); + + // phase 1: clear + for (uint w = tid; w < words; w += gl_WorkGroupSize.x) { + bitmap[w] = 0; + } + if (tid == 0) { + base_sh = 0; + } + barrier(); + + // phase 2: mark. Depth-independent: one pass over the top-k lists. + const uint n_cand = p.n_batch * p.n_top_k; + for (uint c = tid; c < n_cand; c += gl_WorkGroupSize.x) { + const uint t = c / p.n_top_k; + const uint j = c - t * p.n_top_k; + const int idx = data_top[t * p.nbt1 + j]; + if (idx >= 0 && uint(idx) < R) { + atomicOr(bitmap[uint(idx) >> 5], 1u << (uint(idx) & 31u)); + } + } + barrier(); + + // phase 3: compact. R/32 iterations, ascending, exact running offset in shared memory. + for (uint chunk = 0; chunk < words; chunk += gl_WorkGroupSize.x) { + const uint w = chunk + tid; + const uint bits = w < words ? bitmap[w] : 0u; + const uint cnt = bitCount(bits); + + const uint wave_off = subgroupExclusiveAdd(cnt); + const uint wave_tot = subgroupAdd(cnt); + if (lane == 0) { + wave_totals[wave] = wave_tot; + } + barrier(); + + uint prefix = 0; + for (uint i = 0; i < wave; ++i) { + prefix += wave_totals[i]; + } + uint total = 0; + for (uint i = 0; i < nwave; ++i) { + total += wave_totals[i]; + } + + uint slot = base_sh + prefix + wave_off; + uint rem = p.count_only != 0 ? 0u : bits; + while (rem != 0) { + const uint b = findLSB(rem); + rem &= rem - 1; + if (slot < p.max_union) { + data_u[slot] = w * 32 + b; + } + ++slot; + } + barrier(); + if (tid == 0) { + base_sh += total; + } + barrier(); + } + + if (tid == 0) { + const uint u = min(base_sh, p.max_union); + const uint rows = p.n_kv_raw + u; + data_c[0] = ((rows + p.pad_to - 1) / p.pad_to) * p.pad_to; + data_c[1] = u; + data_c[2] = n_cand; + data_c[3] = p.n_batch; + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp new file mode 100644 index 000000000000..a2004bd93d7f --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_cm.comp @@ -0,0 +1,168 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : require +#extension GL_EXT_shader_16bit_storage : require +#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require +#extension GL_KHR_cooperative_matrix : require +#extension GL_KHR_memory_scope_semantics : require +#extension GL_KHR_shader_subgroup_basic : require + +layout(constant_id = 0) const uint SUBGROUP_SIZE = 64; +// Head count is a SPECIALIZATION constant, not a push constant: the head loop's trip +// count must stay visible to the compiler. As a push constant it cannot unroll and +// decode costs ~55% more (measured, gfx1151, kv=131584). +layout(constant_id = 1) const uint N_HEAD = 64; +#if N_WAVES == 1 +layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; +#else +layout(local_size_x = 512, local_size_y = 1, local_size_z = 1) in; +#endif + +layout(binding = 0) readonly buffer QBuf { float data_q[]; }; +layout(binding = 1) readonly buffer KBuf { float16_t data_k[]; }; +layout(binding = 2) readonly buffer WBuf { float data_w[]; }; +layout(binding = 3) readonly buffer MBuf { float16_t data_m[]; }; +layout(binding = 4) writeonly buffer DstBuf { float data_dst[]; }; + +layout(push_constant) uniform Parameters { + uint n_kv; + uint n_batch; + uint n_stream; + uint nem3; + uint nb1; + uint nb3; + uint nbq1; + uint nbq2; + uint nbq3; + uint nbk2; + uint nbk3; + uint nbw1; + uint nbw3; + uint nbm1; + uint nbm3; +} p; + +const uint TILE = 16; +const uint HEAD_SIZE = 128; +const uint VEC_PER_HEAD = HEAD_SIZE / 4; +const uint TILE_STRIDE = VEC_PER_HEAD + 2; +const uint SCORE_STRIDE = TILE / 4 + 1; + +// Shared-memory budget, N_WAVES=8 / HEADS_PER_TILE=4 (TILE 16, HEAD_SIZE 128): +// k_sh 8 * 16 * 34 * 8 B = 34816 +// q_sh 4 * 16 * 34 * 8 B = 17408 +// score_sh 8 * 16 * 5 * 16 B = 10240 +// weight_sh 4 * 16 * 4 B = 256 +// total = 62720 of 65536 (95.7%) +// Only one workgroup fits per CU at that size, which is the intended trade. Raising either +// constant overruns: HEADS_PER_TILE=8 needs 80128 B and the pipeline then fails to create, +// silently falling back to the small variant. The host gates on +// maxComputeSharedMemorySize >= 64 KiB; keep that gate and this arithmetic in sync. +shared f16vec4 q_sh[HEADS_PER_TILE][TILE * TILE_STRIDE]; +shared f16vec4 k_sh[N_WAVES][TILE * TILE_STRIDE]; +shared vec4 score_sh[N_WAVES][TILE * SCORE_STRIDE]; +shared float weight_sh[HEADS_PER_TILE][TILE]; + +void main() { + const uint tid = gl_LocalInvocationIndex; + const uint lane = gl_SubgroupInvocationID; + const uint wave = gl_SubgroupID; + const uint kv_base = (gl_WorkGroupID.x * N_WAVES + wave) * TILE; + const uint token_base = gl_WorkGroupID.y * TILE; + const uint stream = gl_WorkGroupID.z; + + float totals[4]; + [[unroll]] for (uint i = 0; i < 4; ++i) { + totals[i] = 0.0; + } + + for (uint idx = tid; idx < N_WAVES * TILE * VEC_PER_HEAD; idx += gl_WorkGroupSize.x) { + const uint load_wave = idx / (TILE * VEC_PER_HEAD); + const uint wave_idx = idx % (TILE * VEC_PER_HEAD); + const uint key = wave_idx / VEC_PER_HEAD; + const uint d4 = wave_idx % VEC_PER_HEAD; + const uint kv = (gl_WorkGroupID.x * N_WAVES + load_wave) * TILE + key; + f16vec4 value = f16vec4(0.0); + if (kv < p.n_kv) { + const uint offset = stream * p.nbk3 + kv * p.nbk2 + d4 * 4; + value = f16vec4(data_k[offset], data_k[offset + 1], data_k[offset + 2], data_k[offset + 3]); + } + k_sh[load_wave][key * TILE_STRIDE + d4] = value; + } + barrier(); + + // k_sh[wave] is written once above and never touched again, but coopMatLoad sits inside both + // the head_base and head_local loops, so the same 8 A fragments are re-read from shared + // memory N_HEAD times per wave (512 loads where 8 would do). Each wave64 fragment load moves + // 1 KiB of LDS traffic because the 16x16 f16 fragment is replicated 4x across the subgroup. + coopmat kmat_h[HEAD_SIZE / TILE]; + [[unroll]] for (uint d = 0; d < HEAD_SIZE; d += TILE) { + coopMatLoad(kmat_h[d / TILE], k_sh[wave], d / 4, TILE_STRIDE, gl_CooperativeMatrixLayoutRowMajor); + } + + for (uint head_base = 0; head_base < N_HEAD; head_base += HEADS_PER_TILE) { + for (uint idx = tid; idx < HEADS_PER_TILE * TILE * VEC_PER_HEAD; idx += gl_WorkGroupSize.x) { + const uint head_local = idx / (TILE * VEC_PER_HEAD); + const uint head_idx = idx % (TILE * VEC_PER_HEAD); + const uint token_local = head_idx / VEC_PER_HEAD; + const uint d4 = head_idx % VEC_PER_HEAD; + const uint token = token_base + token_local; + f16vec4 value = f16vec4(0.0); + if (token < p.n_batch && (N_HEAD % HEADS_PER_TILE == 0 || head_base + head_local < N_HEAD)) { + const uint offset = stream * p.nbq3 + token * p.nbq2 + (head_base + head_local) * p.nbq1 + d4 * 4; + value = f16vec4(data_q[offset], data_q[offset + 1], data_q[offset + 2], data_q[offset + 3]); + } + q_sh[head_local][token_local * TILE_STRIDE + d4] = value; + } + if (tid < HEADS_PER_TILE * TILE) { + const uint head_local = tid / TILE; + const uint token_local = tid % TILE; + const uint token = token_base + token_local; + weight_sh[head_local][token_local] = (token < p.n_batch && (N_HEAD % HEADS_PER_TILE == 0 || head_base + head_local < N_HEAD)) + ? data_w[stream * p.nbw3 + token * p.nbw1 + head_base + head_local] : 0.0; + } + barrier(); + + [[unroll]] for (uint head_local = 0; head_local < HEADS_PER_TILE; ++head_local) { + coopmat scores = + coopmat(0.0); + coopmat qmat; + + [[unroll]] for (uint d = 0; d < HEAD_SIZE; d += TILE) { + coopMatLoad(qmat, q_sh[head_local], d / 4, TILE_STRIDE, gl_CooperativeMatrixLayoutColumnMajor); + scores = coopMatMulAdd(kmat_h[d / TILE], qmat, scores); + } + + coopMatStore(scores, score_sh[wave], 0, SCORE_STRIDE, gl_CooperativeMatrixLayoutRowMajor); + controlBarrier(gl_ScopeSubgroup, gl_ScopeSubgroup, gl_StorageSemanticsShared, gl_SemanticsAcquireRelease); + + [[unroll]] for (uint i = 0; i < 4; ++i) { + const uint idx = lane + i * SUBGROUP_SIZE; + const uint key = idx / TILE; + const uint token_local = idx % TILE; + const uint token = token_base + token_local; + if (token < p.n_batch && kv_base + key < p.n_kv) { + const float score = score_sh[wave][key * SCORE_STRIDE + token_local / 4][token_local % 4]; + totals[i] += max(score, 0.0) * weight_sh[head_local][token_local]; + } + } + if (head_local + 1 < HEADS_PER_TILE) { + controlBarrier(gl_ScopeSubgroup, gl_ScopeSubgroup, gl_StorageSemanticsShared, gl_SemanticsAcquireRelease); + } + } + barrier(); + } + + [[unroll]] for (uint i = 0; i < 4; ++i) { + const uint idx = lane + i * SUBGROUP_SIZE; + const uint key = idx / TILE; + const uint token_local = idx % TILE; + const uint kv = kv_base + key; + const uint token = token_base + token_local; + if (kv < p.n_kv && token < p.n_batch) { + const uint mask_base = (stream % p.nem3) * p.nbm3 + token * p.nbm1; + const uint dst_base = stream * p.nb3 + token * p.nb1; + data_dst[dst_base + kv] = totals[i] + float(data_m[mask_base + kv]); + } + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_decode_cm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_decode_cm.comp new file mode 100644 index 000000000000..5d32e51361b3 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_decode_cm.comp @@ -0,0 +1,137 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : require +#extension GL_EXT_shader_16bit_storage : require +#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require +#extension GL_KHR_cooperative_matrix : require +#extension GL_KHR_memory_scope_semantics : require + +layout(constant_id = 0) const uint SUBGROUP_SIZE = 64; +// Head count is a SPECIALIZATION constant, not a push constant: the head loop's trip +// count must stay visible to the compiler. As a push constant it cannot unroll and +// decode costs ~55% more (measured, gfx1151, kv=131584). +layout(constant_id = 1) const uint N_HEAD = 64; +layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; + +layout(binding = 0) readonly buffer QBuf { float data_q[]; }; +layout(binding = 1) readonly buffer KBuf { float16_t data_k[]; }; +layout(binding = 2) readonly buffer WBuf { float data_w[]; }; +layout(binding = 3) readonly buffer MBuf { float16_t data_m[]; }; +layout(binding = 4) writeonly buffer DstBuf { float data_dst[]; }; + +layout(push_constant) uniform Parameters { + uint n_kv; + uint n_batch; + uint n_stream; + uint nem3; + uint nb1; + uint nb3; + uint nbq1; + uint nbq2; + uint nbq3; + uint nbk2; + uint nbk3; + uint nbw1; + uint nbw3; + uint nbm1; + uint nbm3; +} p; + +const uint TILE = 16; +const uint HEAD_SIZE = 128; +const uint VEC_PER_HEAD = HEAD_SIZE / 4; +const uint TILE_STRIDE = VEC_PER_HEAD + 2; +const uint SCORE_STRIDE = TILE / 4 + 1; + +shared f16vec4 q_sh[TILE * TILE_STRIDE]; +shared f16vec4 k_sh[TILE * TILE_STRIDE]; +shared vec4 score_sh[TILE * SCORE_STRIDE]; + +void main() { + const uint tid = gl_LocalInvocationIndex; + const uint kv_base = gl_WorkGroupID.x * TILE; + const uint token = gl_WorkGroupID.y; + const uint stream = gl_WorkGroupID.z; + + for (uint idx = tid; idx < TILE * VEC_PER_HEAD; idx += SUBGROUP_SIZE) { + const uint key = idx / VEC_PER_HEAD; + const uint d4 = idx % VEC_PER_HEAD; + const uint kv = kv_base + key; + f16vec4 value = f16vec4(0.0); + if (kv < p.n_kv) { + const uint offset = stream * p.nbk3 + kv * p.nbk2 + d4 * 4; + value = f16vec4(data_k[offset], data_k[offset + 1], data_k[offset + 2], data_k[offset + 3]); + } + k_sh[key * TILE_STRIDE + d4] = value; + } + barrier(); + + // The K tile is invariant across the head loop, but coopMatLoad is inside it, so the shipped + // shader re-reads the same 8 A fragments from shared memory once per head tile (4x at + // N_HEAD=64). Each wave64 fragment load moves 1 KiB of LDS traffic because the 16x16 f16 + // fragment is replicated 4x across the subgroup, so those re-reads are the largest single + // item of per-tile cost after the multiplies themselves. Hoisting costs 8 A fragments of + // register state (16 f16 per lane each). + coopmat kmat[HEAD_SIZE / TILE]; + [[unroll]] for (uint d = 0; d < HEAD_SIZE; d += TILE) { + coopMatLoad(kmat[d / TILE], k_sh, d / 4, TILE_STRIDE, gl_CooperativeMatrixLayoutRowMajor); + } + + float total = 0.0; + for (uint head_base = 0; head_base < N_HEAD; head_base += TILE) { + for (uint idx = tid; idx < TILE * VEC_PER_HEAD; idx += SUBGROUP_SIZE) { + const uint head_local = idx / VEC_PER_HEAD; + const uint d4 = idx % VEC_PER_HEAD; + const uint head = head_base + head_local; + // n_head need not be a multiple of TILE (Qwen4-exp indexers are nh=4), so the tail + // tile stages zeros rather than reading past the end of Q. Zero q gives score 0, + // and relu(0) * w = 0, so the padded heads drop out of the sum on their own. + // The guard exists only when N_HEAD is not tile-aligned (nh=4). The condition is a + // specialization-constant expression, so at nh=64 it folds to true and this compiles + // to the original unguarded load - measured: relying on the compiler to range-prove + // head < N_HEAD instead costs ~45% at nh=64. + f16vec4 value = f16vec4(0.0); + if (N_HEAD % TILE == 0 || head < N_HEAD) { + const uint offset = stream * p.nbq3 + token * p.nbq2 + head * p.nbq1 + d4 * 4; + value = f16vec4(data_q[offset], data_q[offset + 1], data_q[offset + 2], data_q[offset + 3]); + } + q_sh[head_local * TILE_STRIDE + d4] = value; + } + barrier(); + + coopmat scores = + coopmat(0.0); + coopmat qmat; + + [[unroll]] for (uint d = 0; d < HEAD_SIZE; d += TILE) { + coopMatLoad(qmat, q_sh, d / 4, TILE_STRIDE, gl_CooperativeMatrixLayoutColumnMajor); + scores = coopMatMulAdd(kmat[d / TILE], qmat, scores); + } + + coopMatStore(scores, score_sh, 0, SCORE_STRIDE, gl_CooperativeMatrixLayoutRowMajor); + barrier(); + + if (tid < TILE && kv_base + tid < p.n_kv) { + [[unroll]] for (uint head_local = 0; head_local < TILE; ++head_local) { + // Padded heads staged q as zero, so their score is 0 and the term dies in the + // multiply. Only the weight READ needs bounding, and clamping the index does that + // without a branch - a dynamic break here forces the loop rolled and costs ~60%. + const uint head = N_HEAD % TILE == 0 ? head_base + head_local + : min(head_base + head_local, N_HEAD - 1); + const float score = score_sh[tid * SCORE_STRIDE + head_local / 4][head_local % 4]; + const float weight = data_w[stream * p.nbw3 + token * p.nbw1 + head]; + total += max(score, 0.0) * weight; + } + } + barrier(); + } + + if (tid < TILE) { + const uint kv = kv_base + tid; + if (kv < p.n_kv) { + const uint mask_base = (stream % p.nem3) * p.nbm3 + token * p.nbm1; + const uint dst_base = stream * p.nb3 + token * p.nb1; + data_dst[dst_base + kv] = total + float(data_m[mask_base + kv]); + } + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_scalar64.comp b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_scalar64.comp new file mode 100644 index 000000000000..83fd923fe738 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer_scalar64.comp @@ -0,0 +1,95 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : require +#extension GL_EXT_shader_16bit_storage : require +#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require +#extension GL_KHR_shader_subgroup_arithmetic : require +#extension GL_KHR_shader_subgroup_basic : require + +layout(constant_id = 0) const uint SUBGROUP_SIZE = 64; +// Head count is a SPECIALIZATION constant, not a push constant: the head loop's trip +// count must stay visible to the compiler. As a push constant it cannot unroll and +// decode costs ~55% more (measured, gfx1151, kv=131584). +layout(constant_id = 1) const uint N_HEAD = 64; +layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; + +layout(binding = 0) readonly buffer QBuf { float data_q[]; }; +layout(binding = 1) readonly buffer KBuf { float16_t data_k[]; }; +layout(binding = 2) readonly buffer WBuf { float data_w[]; }; +layout(binding = 3) readonly buffer MaskBuf { float16_t data_m[]; }; +layout(binding = 4) writeonly buffer DstBuf { float data_dst[]; }; + +layout(push_constant) uniform Parameters { + uint n_kv; + uint n_batch; + uint n_stream; + uint nem3; + uint nb1; + uint nb3; + uint nbq1; + uint nbq2; + uint nbq3; + uint nbk2; + uint nbk3; + uint nbw1; + uint nbw3; + uint nbm1; + uint nbm3; +} p; + +const uint K_PER_GROUP = 8; + +void main() { + const uint lane = gl_SubgroupInvocationID; + const uint token = gl_WorkGroupID.y; + const uint stream = gl_WorkGroupID.z; + const uint kv_base = gl_WorkGroupID.x * K_PER_GROUP; + + if (token >= p.n_batch || stream >= p.n_stream) { + return; + } + + float k0[K_PER_GROUP]; + float k1[K_PER_GROUP]; + [[unroll]] for (uint j = 0; j < K_PER_GROUP; ++j) { + const uint kv = kv_base + j; + if (kv < p.n_kv) { + const uint k_base = stream * p.nbk3 + kv * p.nbk2; + k0[j] = float(data_k[k_base + lane]); + k1[j] = float(data_k[k_base + lane + SUBGROUP_SIZE]); + } else { + k0[j] = 0.0; + k1[j] = 0.0; + } + } + + float score[K_PER_GROUP]; + [[unroll]] for (uint j = 0; j < K_PER_GROUP; ++j) { + score[j] = 0.0; + } + + for (uint head = 0; head < N_HEAD; ++head) { + const uint q_base = stream * p.nbq3 + token * p.nbq2 + head * p.nbq1; + const float q0 = data_q[q_base + lane]; + const float q1 = data_q[q_base + lane + SUBGROUP_SIZE]; + const float weight = data_w[stream * p.nbw3 + token * p.nbw1 + head]; + + [[unroll]] for (uint j = 0; j < K_PER_GROUP; ++j) { + const float qk = subgroupAdd(q0 * k0[j] + q1 * k1[j]); + if (lane == 0) { + score[j] += max(qk, 0.0) * weight; + } + } + } + + if (lane == 0) { + const uint mask_base = (stream % p.nem3) * p.nbm3 + token * p.nbm1; + const uint dst_base = stream * p.nb3 + token * p.nb1; + [[unroll]] for (uint j = 0; j < K_PER_GROUP; ++j) { + const uint kv = kv_base + j; + if (kv < p.n_kv) { + data_dst[dst_base + kv] = score[j] + float(data_m[mask_base + kv]); + } + } + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp index 63c4aaebcb1a..7d32124b969b 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp @@ -68,6 +68,9 @@ layout (binding = 2) writeonly buffer D {D_TYPE data_d[];}; #ifdef MUL_MAT_ID layout (binding = 3) readonly buffer IDS {int data_ids[];}; layout (binding = 4) readonly buffer Counts {int data_expert_count[];}; +// Fused MUL epilogue: one scale per (expert slot, token), broadcast down the M dimension. +// Always bound; p.fusion_flags == 0 means ignore it. +layout (binding = 5) readonly buffer FusedScale {float data_fscale[];}; #endif layout (push_constant) uniform parameter @@ -90,6 +93,7 @@ layout (push_constant) uniform parameter uint ne11; uint n_experts; uint hoist_row_ids; + uint fusion_flags; #else uint base_work_group_z; uint num_batches; @@ -394,9 +398,13 @@ void main() { if (row_i >= _ne1) break; const u16vec2 row_idx = row_ids[row_i - ic * BN]; - if (dr + cm_row * TM + store_r < p.M) { - data_d[row_idx.y * p.batch_stride_d + row_idx.x * p.stride_d + dr + cm_row * TM + store_r] = D_TYPE(coopmat_stage[warp_i * TM * TN + (col + store_c) * TM + store_r]); + const uint didx = row_idx.y * p.batch_stride_d + row_idx.x * p.stride_d + dr + cm_row * TM + store_r; + if (p.fusion_flags != 0) { + data_d[didx] = D_TYPE(float(coopmat_stage[warp_i * TM * TN + (col + store_c) * TM + store_r]) * data_fscale[row_idx.y * p.nei0 + row_idx.x]); + } else { + data_d[didx] = D_TYPE(coopmat_stage[warp_i * TM * TN + (col + store_c) * TM + store_r]); + } } } barrier(); @@ -449,15 +457,19 @@ void main() { if (row_i >= _ne1) break; const u16vec2 row_idx = row_ids[row_i - ic * BN]; + const bool do_scale = p.fusion_flags != 0; + const float fscale = do_scale ? data_fscale[row_idx.y * p.nei0 + row_idx.x] : 1.0f; #endif // MUL_MAT_ID [[unroll]] for (uint cr = 0; cr < TM / 2; cr++) { const uint sums_idx = (wsic * TN + cc) * WMITER * (TM / 2) + wsir * (TM / 2) + cr; #ifdef MUL_MAT_ID if (dr_warp + 2 * cr < p.M) { - data_d[row_idx.y * p.batch_stride_d + row_idx.x * p.stride_d + dr_warp + 2 * cr] = D_TYPE(sums[sums_idx].x); + const uint didx = row_idx.y * p.batch_stride_d + row_idx.x * p.stride_d + dr_warp + 2 * cr; + data_d[didx] = do_scale ? D_TYPE(float(sums[sums_idx].x) * fscale) : D_TYPE(sums[sums_idx].x); } if (dr_warp + 2 * cr + 1 < p.M) { - data_d[row_idx.y * p.batch_stride_d + row_idx.x * p.stride_d + dr_warp + 2 * cr + 1] = D_TYPE(sums[sums_idx].y); + const uint didx = row_idx.y * p.batch_stride_d + row_idx.x * p.stride_d + dr_warp + 2 * cr + 1; + data_d[didx] = do_scale ? D_TYPE(float(sums[sums_idx].y) * fscale) : D_TYPE(sums[sums_idx].y); } #else if (dr_warp + 2 * cr < p.M && dc_warp + cc < p.N) { diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp index 27f3178e7f26..9bfad031d1b6 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp @@ -106,6 +106,9 @@ layout (binding = 1) readonly buffer B4 {B_TYPEV4 data_b_v4[];}; #ifdef MUL_MAT_ID layout (binding = 3) readonly buffer IDS {int data_ids[];}; layout (binding = 4) readonly buffer Counts {int data_expert_count[];}; +// Bound for descriptor-layout parity with the other mul_mat_id shaders. The fused MUL +// epilogue is not implemented here, so the host never enables it on coopmat2 devices. +layout (binding = 5) readonly buffer FusedScale {float data_fscale[];}; shared u16vec4 row_ids[BN]; diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq.comp index 1fbcbf6c9332..b36d056c1c9c 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq.comp @@ -36,6 +36,8 @@ layout (binding = 2) writeonly buffer D {D_TYPE data_d[];}; #ifdef MUL_MAT_ID layout (binding = 3) readonly buffer IDS {int data_ids[];}; layout (binding = 4) readonly buffer Counts {int data_expert_count[];}; +// Fused MUL epilogue, see mul_mm.comp. +layout (binding = 5) readonly buffer FusedScale {float data_fscale[];}; #endif layout (push_constant) uniform parameter @@ -58,6 +60,7 @@ layout (push_constant) uniform parameter uint ne11; uint n_experts; uint hoist_row_ids; + uint fusion_flags; #else uint base_work_group_z; uint num_batches; @@ -303,7 +306,10 @@ void main() { const uint sums_idx = (wsic * TN + cc) * WMITER * TM + wsir * TM + cr; #ifdef MUL_MAT_ID if (dr_warp + cr < p.M) { - data_d[row_idx.y * p.batch_stride_d + row_idx.x * p.stride_d + dr_warp + cr] = D_TYPE(sums[sums_idx].x); + const uint didx = row_idx.y * p.batch_stride_d + row_idx.x * p.stride_d + dr_warp + cr; + data_d[didx] = p.fusion_flags != 0 + ? D_TYPE(float(sums[sums_idx].x) * data_fscale[row_idx.y * p.nei0 + row_idx.x]) + : D_TYPE(sums[sums_idx].x); } #else if (dr_warp + cr < p.M && dc_warp + cc < p.N) { diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp index 1a57fdc24e8c..295e17b7b088 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -791,9 +791,13 @@ void process_shaders() { string_to_spv("dequant_" + tname, "dequant_" + tname + ".comp", merge_maps(base_dict, {{data_a_key, "1"}, {"D_TYPE", "float16_t"}})); } // Fused dequant+transpose variant for FA quant-KV (per-head-contiguous f16 scratch). - if (tname == "q8_0") { + if (tname == "q8_0" || tname == "iq4_nl" || tname == "q4_0" || tname == "q4_1" || tname == "q5_0" || tname == "q5_1") { string_to_spv("dequant_" + tname + "_transpose", "dequant_" + tname + ".comp", merge_maps(base_dict, {{data_a_key, "1"}, {"D_TYPE", "float16_t"}, {"DEQUANT_TRANSPOSE", "1"}})); } + // Strided-copy counterpart for f16 KV (contiguize the head-interleaved cache layout). + if (tname == "f16") { + string_to_spv("dequant_f16_transpose", "dequant_f16_transpose.comp", {}); + } shader = (tname == "f32" || tname == "f16" || tname == "bf16") ? "get_rows.comp" : "get_rows_quant.comp"; @@ -807,6 +811,33 @@ void process_shaders() { string_to_spv("get_rows_i32", "get_rows.comp", {{"TEMP_TYPE", "uint"}, {"A_TYPE", "uint"}, {"B_TYPE", "int"}, {"D_TYPE", "uint"}}); + string_to_spv("lightning_indexer_f16", "lightning_indexer_scalar64.comp", {}); +#if defined(GGML_VULKAN_COOPMAT_GLSLC_SUPPORT) + string_to_spv("lightning_indexer_cm_f16", "lightning_indexer_cm.comp", {{"N_WAVES", "8"}, {"HEADS_PER_TILE", "4"}}); + string_to_spv("lightning_indexer_cm_small_f16", "lightning_indexer_cm.comp", {{"N_WAVES", "1"}, {"HEADS_PER_TILE", "1"}}); + string_to_spv("lightning_indexer_decode_cm_f16", "lightning_indexer_decode_cm.comp", {}); +#endif + string_to_spv("flash_attn_top_k_f16", "flash_attn_top_k.comp", {}); +#if defined(GGML_VULKAN_COOPMAT_GLSLC_SUPPORT) + string_to_spv("flash_attn_top_k_cm_f16", "flash_attn_top_k_cm.comp", {}); +#endif + string_to_spv("flash_attn_gather_f16", "flash_attn_gather.comp", {}); + string_to_spv("flash_attn_union_f16", "flash_attn_union.comp", {}); + string_to_spv("flash_attn_gather_union_f16", "flash_attn_gather_union.comp", {}); + // one decoder per quantised K type flash-attention supports natively; f16/bf16/f32 need no + // decode and take the verbatim gather + // q8_0 and q4_0 only: each needs its true element mapping written out (see the shader), and + // these are the two types actually used as a KV cache. The rest take the verbatim gather. + for (const auto& tname : {"q4_0", "q8_0"}) { + string_to_spv("flash_attn_gather_union_dq_" + std::string(tname), "flash_attn_gather_union_dq.comp", + {{"DATA_A_" + to_uppercase(tname), "1"}}); + string_to_spv("flash_attn_gather_dq_" + std::string(tname), "flash_attn_gather_dq.comp", + {{"DATA_A_" + to_uppercase(tname), "1"}}); + } + string_to_spv("dsv4_hc_pre_f32", "dsv4_hc_pre.comp", {}); + string_to_spv("dsv4_hc_comb_f32", "dsv4_hc_comb.comp", {}); + string_to_spv("dsv4_hc_post_f32", "dsv4_hc_post.comp", {}); + string_to_spv("mul_mat_vec_p021_f16_f32_subgroup_add", "mul_mat_vec_p021.comp", {{"A_TYPE", "float16_t"}, {"A_TYPEV4", "f16vec4"}, {"B_TYPE", "float"}, {"B_TYPEV4", "vec4"}, {"D_TYPE", "float"}, {"USE_SUBGROUP_ADD", "1"}}); string_to_spv("mul_mat_vec_p021_f16_f32", "mul_mat_vec_p021.comp", {{"A_TYPE", "float16_t"}, {"A_TYPEV4", "f16vec4"}, {"B_TYPE", "float"}, {"B_TYPEV4", "vec4"}, {"D_TYPE", "float"}}); string_to_spv("mul_mat_vec_nc_f16_f32", "mul_mat_vec_nc.comp", {{"A_TYPE", "float16_t"}, {"A_TYPEV4", "f16vec4"}, {"B_TYPE", "float"}, {"B_TYPEV4", "vec4"}, {"D_TYPE", "float"}}); @@ -920,6 +951,7 @@ void process_shaders() { string_to_spv("concat_i16", "concat.comp", {{"A_TYPE", "uint16_t"}, {"B_TYPE", "uint16_t"}, {"D_TYPE", "uint16_t"}}); string_to_spv("concat_i32", "concat.comp", {{"A_TYPE", "uint"}, {"B_TYPE", "uint"}, {"D_TYPE", "uint"}}); string_to_spv("concat_i64", "concat.comp", {{"A_TYPE", "uvec2"}, {"B_TYPE", "uvec2"}, {"D_TYPE", "uvec2"}}); + string_to_spv("concat_transpose_i32", "concat_transpose.comp", {{"A_TYPE", "uint"}, {"B_TYPE", "uint"}, {"D_TYPE", "uint"}}); string_to_spv("upscale_f32", "upscale.comp", {{"A_TYPE", "float"}, {"B_TYPE", "float"}, {"D_TYPE", "float"}}); diff --git a/ggml/src/ggml.c b/ggml/src/ggml.c index 5fcb6c2f300a..c5093a039aa5 100644 --- a/ggml/src/ggml.c +++ b/ggml/src/ggml.c @@ -5597,6 +5597,20 @@ void ggml_flash_attn_ext_add_sinks( a->src[4] = sinks; } +void ggml_flash_attn_ext_add_top_k( + struct ggml_tensor * a, + struct ggml_tensor * top_k, + int64_t n_kv_raw) { + GGML_ASSERT(a->op == GGML_OP_FLASH_ATTN_EXT); + GGML_ASSERT(a->src[5] == NULL); + GGML_ASSERT(top_k->type == GGML_TYPE_I32); + GGML_ASSERT(top_k->ne[1] == a->src[0]->ne[1]); + GGML_ASSERT(n_kv_raw >= 0 && n_kv_raw <= a->src[1]->ne[1]); + + a->src[5] = top_k; + ggml_set_op_params_i32(a, 4, (int32_t) n_kv_raw); +} + // ggml_flash_attn_back struct ggml_tensor * ggml_flash_attn_back( diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp index b137c6cfe8a0..ed56fdbb1040 100644 --- a/src/llama-graph.cpp +++ b/src/llama-graph.cpp @@ -2554,7 +2554,9 @@ ggml_tensor * llm_graph_context::build_attn_mha( ggml_tensor * v_mla, int64_t n_kv_max, float kq_scale, - int il) const { + int il, + ggml_tensor * top_k, + int64_t n_kv_raw) const { const bool v_trans = v->nb[1] > v->nb[2]; // split the batch into streams if needed @@ -2592,6 +2594,9 @@ ggml_tensor * llm_graph_context::build_attn_mha( ggml_flash_attn_ext_add_sinks(cur, sinks); GGML_ASSERT(n_kv_max >= 0 && n_kv_max <= INT32_MAX); ggml_flash_attn_ext_set_n_kv_max(cur, static_cast(n_kv_max)); + if (top_k) { + ggml_flash_attn_ext_add_top_k(cur, top_k, n_kv_raw); + } ggml_flash_attn_ext_set_prec (cur, GGML_PREC_F32); if (v_mla) { diff --git a/src/llama-graph.h b/src/llama-graph.h index dddfdac7b51e..862142b4fd7b 100644 --- a/src/llama-graph.h +++ b/src/llama-graph.h @@ -1173,7 +1173,9 @@ struct llm_graph_context { ggml_tensor * v_mla, // [n_embd_head_v_mla, n_embd_head_v, n_head_v] int64_t n_kv_max, float kq_scale, - int il) const; + int il, + ggml_tensor * top_k = nullptr, + int64_t n_kv_raw = 0) const; llm_graph_input_attn_no_cache * build_attn_inp_no_cache() const; diff --git a/src/llama-kv-cache-dsa.cpp b/src/llama-kv-cache-dsa.cpp index 96cb045d2e5d..e926e34314b4 100644 --- a/src/llama-kv-cache-dsa.cpp +++ b/src/llama-kv-cache-dsa.cpp @@ -47,8 +47,11 @@ llama_kv_cache_dsa::llama_kv_cache_dsa( LLAMA_LOG_INFO("%s: creating indexer KV cache, size = %u cells\n", __func__, kv_size); + // keep indexer keys f16 regardless of type_k: the fused indexer kernels read + // f16 only, and quantizing this small cache (128 dims) saves little while + // forcing the much slower decomposed indexer path kv_lid = std::make_unique( - model, hparams_lid, type_k, type_v, + model, hparams_lid, GGML_TYPE_F16, type_v, v_trans, offload, unified, kv_size, n_seq_max, n_pad, n_swa, swa_type, nullptr, filter_lid, reuse, nullptr); } diff --git a/src/llama-kv-cache-dsv4.cpp b/src/llama-kv-cache-dsv4.cpp index 9e82a6198a09..2097271a0f18 100644 --- a/src/llama-kv-cache-dsv4.cpp +++ b/src/llama-kv-cache-dsv4.cpp @@ -1305,8 +1305,11 @@ llama_kv_cache_dsv4::llama_kv_cache_dsv4( LLAMA_LOG_INFO("%s: creating DSV4 lightning-indexer KV cache, size = %u cells\n", __func__, dsv4_comp_size(kv_size, DSV4_CSA_RATIO)); + // keep indexer keys f16 regardless of type_k: the fused indexer kernels read + // f16 only, and quantizing this small cache (128 dims) saves little while + // forcing the much slower decomposed indexer path kv_lid = std::make_unique( - model, hparams_lid, type_k, type_v, + model, hparams_lid, GGML_TYPE_F16, type_v, v_trans, offload, unified_compressed, GGML_PAD(dsv4_comp_size(kv_size, DSV4_CSA_RATIO), 256u), n_seq_max, n_pad, 0, LLAMA_SWA_TYPE_NONE, nullptr, filter_csa, nullptr, nullptr); diff --git a/src/models/deepseek4.cpp b/src/models/deepseek4.cpp index 6bf9d3444942..494e643276c5 100644 --- a/src/models/deepseek4.cpp +++ b/src/models/deepseek4.cpp @@ -758,7 +758,7 @@ ggml_tensor * llama_model_deepseek4::graph::build_csa_lid_attention( cb(kq_mask, "csa_lid_kq_mask", il); const int64_t n_kv_max = std::min(raw_mask->ne[0], hparams.n_swa) + top_k->ne[0]; - ggml_tensor * out = build_attn_mha(q, k_all, k_all, nullptr, kq_mask, sinks, nullptr, n_kv_max, kq_scale, il); + ggml_tensor * out = build_attn_mha(q, k_all, k_all, nullptr, kq_mask, sinks, nullptr, n_kv_max, kq_scale, il, top_k, raw_k->ne[2]); if (k_rot) { out = llama_mul_mat_hadamard(ctx0, out, k_rot); } @@ -1209,6 +1209,11 @@ ggml_tensor * llama_model_deepseek4::graph::build_attention_impl( out = ggml_reshape_3d(ctx0, out, o_group_dim, n_groups, nt); out = ggml_permute(ctx0, out, 0, 2, 1, 3); + // small multi-token batches (speculative verify) hit a pathological strided-B + // path in the grouped matmul below; a contiguous copy is much cheaper + if (nt > 1 && nt <= 8) { + out = ggml_cont(ctx0, out); + } ggml_tensor * oa = ggml_mul_mat(ctx0, layer.wo_a, out); cb(oa, "attn_wo_a", il); oa = ggml_permute(ctx0, oa, 0, 2, 1, 3); diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 3919e555273f..964c51728667 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -7915,7 +7915,7 @@ struct test_flash_attn_ext : public test_case { const ggml_type type_K; const ggml_type type_V; std::array permute; - const bool kv_view; // create K/V as views of a larger buffer (like a KV cache) + const bool kv_view; // create K/V as views of a larger buffer (like a KV cache); false = dense permuted like the model KV cache const bool v_is_view_of_k; const int64_t n_kv_max; @@ -8028,6 +8028,130 @@ struct test_flash_attn_ext : public test_case { } }; +// GGML_OP_FLASH_ATTN_EXT with a top-k sparse selection hint (DeepSeek V4 CSA shape). +// The kq_mask encodes the same selection as the top_k indices, so a backend that ignores +// the hint (CPU) computes the identical result densely — this is exactly the contract +// that keeps sparse and dense paths interchangeable, and what this test verifies. +struct test_flash_attn_ext_top_k : public test_case { + const int64_t kv; // total KV size (compressed region + dense prefix) + const int64_t nb; // batch size (query tokens) + const int64_t n_kv_raw; // dense prefix always attended + const int64_t n_top_k; // selected keys per query token + const bool sinks; + const int64_t ns; // sequences (ne3); >1 exercises the split-K stream stride + const int64_t ov; // % of each token's picks shared with its neighbours (dedup-union realism) + const ggml_type type_K; // K/V cache type; V is the same tensor, so one type covers both + + static constexpr int64_t hs = 512; // V4 CSA head size, K == V latent + static constexpr int64_t nh = 64; // V4 CSA query heads (MQA) + + std::string vars() override { + return VARS_TO_STR8(kv, nb, n_kv_raw, n_top_k, sinks, ns, ov, type_K); + } + + double max_nmse_err() override { + return 5e-4; + } + + uint64_t op_flops(ggml_tensor * t) override { + GGML_UNUSED(t); + // only the active keys contribute compute on a sparse backend; count those so + // perf mode reports the useful-work rate + return 2 * nh * nb * ns * (hs + hs) * (n_kv_raw + n_top_k); + } + + test_flash_attn_ext_top_k(int64_t kv = 768, int64_t nb = 8, int64_t n_kv_raw = 64, int64_t n_top_k = 128, bool sinks = false, int64_t ns = 1, int64_t ov = 0, + ggml_type type_K = GGML_TYPE_F16) + : kv(kv), nb(nb), n_kv_raw(n_kv_raw), n_top_k(n_top_k), sinks(sinks), ns(ns), ov(ov), type_K(type_K) {} + + ggml_tensor * build_graph(ggml_context * ctx) override { + ggml_tensor * q = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, hs, nb, nh, ns); + ggml_set_name(q, "q"); + + ggml_tensor * k = ggml_new_tensor_4d(ctx, type_K, hs, kv, 1, ns); + ggml_set_name(k, "k"); + + // V4 CSA attends over the K latent itself: V is the same cache tensor + ggml_tensor * v = ggml_view_4d(ctx, k, hs, kv, 1, ns, k->nb[1], k->nb[2], k->nb[3], 0); + ggml_set_name(v, "v"); + + ggml_tensor * m = ggml_new_tensor_4d(ctx, GGML_TYPE_F16, kv, nb, 1, ns); + ggml_set_name(m, "m"); + + ggml_tensor * t = ggml_new_tensor_4d(ctx, GGML_TYPE_I32, n_top_k, nb, 1, ns); + ggml_set_name(t, "top_k"); + + ggml_tensor * s = nullptr; + if (sinks) { + s = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, nh); + ggml_set_name(s, "s"); + } + + ggml_tensor * out = ggml_flash_attn_ext(ctx, q, k, v, m, 1.0f/sqrtf(hs), 0.0f, 0.0f); + ggml_flash_attn_ext_add_sinks(out, s); + ggml_flash_attn_ext_add_top_k(out, t, n_kv_raw); + ggml_flash_attn_ext_set_prec (out, GGML_PREC_F32); + ggml_set_name(out, "out"); + + return out; + } + + void initialize_tensors(ggml_context * ctx) override { + const int64_t range = kv - n_kv_raw; // size of the selectable compressed region + + for (ggml_tensor * t = ggml_get_first_tensor(ctx); t != NULL; t = ggml_get_next_tensor(ctx, t)) { + if (strcmp(t->name, "top_k") == 0 || strcmp(t->name, "m") == 0) { + continue; // filled together below + } + if (strcmp(t->name, "s") == 0) { + init_tensor_uniform(t, -10.0f, 10.0f); + } else { + init_tensor_uniform(t); + } + } + + // build a consistent (top_k, mask) pair: a deterministic per-token selection, + // strided so adjacent tokens select overlapping-but-different keys, with one + // deliberately invalid index (-1) whose mask slot stays -inf + std::vector top(n_top_k * nb * ns); + std::vector mask(kv * nb * ns); + const ggml_fp16_t minus_inf = ggml_fp32_to_fp16(-INFINITY); + const ggml_fp16_t zero = ggml_fp32_to_fp16(0.0f); + + for (int64_t s = 0; s < ns; ++s) { + for (int64_t b = 0; b < nb; ++b) { + const int64_t mrow = (s * nb + b) * kv; + const int64_t trow = (s * nb + b) * n_top_k; + for (int64_t i = 0; i < kv; ++i) { + mask[mrow + i] = i < n_kv_raw ? zero : minus_inf; + } + for (int64_t j = 0; j < n_top_k; ++j) { + // offset the selection by the stream too, so a dropped stream stride + // reads another sequence's keys and shows up as a mismatch + const bool shared = (int64_t) j * 100 < n_top_k * ov; + int32_t idx = shared + ? (int32_t) ((j * range) / n_top_k + s * 7) % (int32_t) range + : (int32_t) ((j * range) / n_top_k + b + s * 7) % (int32_t) range; + if (j == n_top_k - 1 && b == 0 && s == 0) { + idx = -1; // exercise the ignore-invalid-index path + } else { + mask[mrow + n_kv_raw + idx] = zero; + } + top[trow + j] = idx; + } + } + } + + for (ggml_tensor * t = ggml_get_first_tensor(ctx); t != NULL; t = ggml_get_next_tensor(ctx, t)) { + if (strcmp(t->name, "top_k") == 0) { + ggml_backend_tensor_set(t, top.data(), 0, top.size() * sizeof(int32_t)); + } else if (strcmp(t->name, "m") == 0) { + ggml_backend_tensor_set(t, mask.data(), 0, mask.size() * sizeof(ggml_fp16_t)); + } + } + } +}; + // GGML_OP_CROSS_ENTROPY_LOSS struct test_cross_entropy_loss : public test_case { const ggml_type type; @@ -10391,6 +10515,19 @@ static std::vector> make_test_cases_eval() { int k = 256; test_cases.emplace_back(new test_mul_mat_id(type_a, type_b, n_mats, n_used, b, m, n, k)); } + // MMQ tile-boundary cases. The MMQ J config is picked from n, and a wave-partitioning error in + // one J tile only shows up when n sits on that tile's boundary: the neighbouring n selects a + // different J and passes, which hides it. The general cases above stop at 129 and the MUL_MAT + // set jumps 64 -> 4096, so no existing case lands on 256 or 512. + // MMQ tile-boundary sweep: n on and either side of the 256 / 512 J boundaries. + for (ggml_type type_a : {GGML_TYPE_Q8_0, GGML_TYPE_Q4_0, GGML_TYPE_Q4_K, GGML_TYPE_Q5_K, GGML_TYPE_Q6_K}) { + for (int64_t n : {255, 256, 257, 511, 512, 513}) { + test_cases.emplace_back(new test_mul_mat(type_a, GGML_TYPE_F32, 1024, n, 256, {1, 1}, {1, 1})); + test_cases.emplace_back(new test_mul_mat_id(type_a, GGML_TYPE_F32, 8, 2, false, 1024, n, 256)); + } + } + + } } } @@ -10986,6 +11123,19 @@ static std::vector> make_test_cases_eval() { } } + // dense-permuted K/V (model KV-cache layout, engages the f16 contiguize path at nb>=64) + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 4, {8, 1}, 1024, 128, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 2, 1, 3}, false)); + test_cases.emplace_back(new test_flash_attn_ext(96, 96, 8, {4, 1}, 512, 80, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 2, 1, 3}, false)); + test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {8, 1}, 512, 75, true, false, 8.0f, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 2, 1, 3}, false)); + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 4, {8, 1}, 512, 96, true, false, 0, 30.0f, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 2, 1, 3}, false)); + + // dense-permuted iq4_nl K/V at prefill batch sizes: iq4_nl has no native FA shader, so these + // exercise the only supported route (the dequant-once path), incl. sinks and mixed-with-f16 + test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {1, 1}, 512, 75, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_IQ4_NL, GGML_TYPE_IQ4_NL, {0, 2, 1, 3}, false)); + test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {1, 1}, 512, 75, true, true, 0, 0, GGML_PREC_F32, GGML_TYPE_IQ4_NL, GGML_TYPE_IQ4_NL, {0, 2, 1, 3}, false)); + test_cases.emplace_back(new test_flash_attn_ext(72, 72, 4, {4, 1}, 113, 75, true, true, 0, 0, GGML_PREC_F32, GGML_TYPE_IQ4_NL, GGML_TYPE_IQ4_NL, {0, 2, 1, 3}, false)); + test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {1, 1}, 512, 75, true, true, 0, 0, GGML_PREC_F32, GGML_TYPE_IQ4_NL, GGML_TYPE_F16, {0, 2, 1, 3}, false)); + // mixed quant and Q1_0 test cases test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {1, 1}, 128, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q4_0)); test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {1, 1}, 128, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q4_0, GGML_TYPE_F16)); @@ -11237,7 +11387,7 @@ static std::vector> make_test_cases_eval() { // lightning_indexer for (int kv : { 256 }) { - for (int bs : { 1, 512 }) { + for (int bs : { 1, 4, 8, 15, 512 }) { for (int nh : { 32, 64 }) { for (auto [ns, nm] : { std::pair{1, 1}, std::pair{4, 4}, std::pair{4, 1} }) { for (ggml_type type_K : {GGML_TYPE_F32, GGML_TYPE_F16, GGML_TYPE_BF16, GGML_TYPE_Q8_0, GGML_TYPE_Q5_1, GGML_TYPE_Q5_0, GGML_TYPE_Q4_1, GGML_TYPE_Q4_0, GGML_TYPE_IQ4_NL}) { @@ -11247,6 +11397,13 @@ static std::vector> make_test_cases_eval() { } } } + // batch 1 = Vulkan decode-cm variant, 4/15 = scalar subgroup variant (below the cm + // threshold of 16, 15 is the boundary), 17/512 = cm prefill variant + test_cases.emplace_back(new test_lightning_indexer(128, 64, 257, 1, 1, 1, GGML_TYPE_F16)); + test_cases.emplace_back(new test_lightning_indexer(128, 64, 257, 4, 1, 1, GGML_TYPE_F16)); + test_cases.emplace_back(new test_lightning_indexer(128, 64, 257, 15, 1, 1, GGML_TYPE_F16)); + test_cases.emplace_back(new test_lightning_indexer(128, 64, 257, 17, 1, 1, GGML_TYPE_F16)); + test_cases.emplace_back(new test_lightning_indexer(128, 64, 512, 512, 1, 1, GGML_TYPE_F16)); for (int kv : { 1, 7, 8, 63, 64, 65 }) { for (ggml_type type_K : {GGML_TYPE_F32, GGML_TYPE_F16, GGML_TYPE_BF16, GGML_TYPE_Q8_0, GGML_TYPE_Q5_1, GGML_TYPE_Q5_0, GGML_TYPE_Q4_1, GGML_TYPE_Q4_0}) { @@ -11254,6 +11411,60 @@ static std::vector> make_test_cases_eval() { } } + // sparse top-k FA: (kv, nb, n_kv_raw, n_top_k, sinks). The Vulkan sparse path engages + // when kv >= 3*(n_kv_raw + n_top_k) AND nb >= 64 (prefill-only); the nb < 64 cases + // and the kv=512 case verify dense-fallback parity with the hint attached, the + // nb=64/128 cases exercise the sparse shader itself. + test_cases.emplace_back(new test_flash_attn_ext_top_k(4096, 1, 256, 512, false)); + test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 8, 64, 128, false)); + test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 17, 64, 128, false)); + test_cases.emplace_back(new test_flash_attn_ext_top_k( 512, 4, 64, 128, false)); + test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 64, 64, 128, false)); + test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 64, 64, 128, true)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(1024, 64, 65, 128, false)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(4096, 128, 256, 512, false)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(4096, 257, 256, 512, false)); + // ns > 1: the split-K partial-output path indexes O and L/M by stream, so these cover + // the stream stride in both regions (single tile and multi-tile). + // small-batch decode (speculative drafts): each token gets its own gathered top-k block, + // so cross-token rows must be masked out or the softmax double counts. kv must be large + // enough that compaction is worth it (the gather gates on kv >= 2*kv_c). + test_cases.emplace_back(new test_flash_attn_ext_top_k(8192, 2, 1024, 512, false)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(8192, 3, 1024, 512, false)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(8192, 4, 1024, 512, false)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(8192, 8, 1024, 512, true)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(32768, 16, 2304, 512, false)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(65536, 63, 2304, 512, false)); + // overlapping selections at the shapes where the compaction gate is tightest: with a + // deduplicated union these are admitted on the estimated union size rather than the + // worst case, so they cover the estimator's gate as well as the union itself. + test_cases.emplace_back(new test_flash_attn_ext_top_k(11008, 8, 2304, 512, false, 1, 60)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(11008, 16, 2304, 512, false, 1, 60)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(11008, 16, 2304, 512, false, 1, 86)); + // quantised K/V: the gather relocates rows verbatim, so it should serve any type whose row + // is a whole number of 4-byte words. These are the shapes a DSv4 decode with -ctk q8_0 hits, + // which took the dense fallback entirely before the gather learned to address rows as bytes. + for (ggml_type tk : { GGML_TYPE_Q8_0, GGML_TYPE_Q4_0 }) { + test_cases.emplace_back(new test_flash_attn_ext_top_k(8192, 4, 1024, 512, false, 1, 0, tk)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(11008, 8, 2304, 512, false, 1, 60, tk)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(11008, 16, 2304, 512, false, 1, 86, tk)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(8192, 1, 1024, 512, false, 1, 0, tk)); + // prefill widths (nb >= 64), where the sparse shaders run on a dequantised f16 scratch + // instead of the cache. ns=2 covers the scratch's stream stride, and the 4096 case is + // wide enough for the raw/selected split form. + test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 64, 64, 128, false, 1, 0, tk)); + test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 64, 64, 128, false, 2, 0, tk)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(4096, 128, 256, 512, false, 1, 0, tk)); + } + test_cases.emplace_back(new test_flash_attn_ext_top_k(8192, 4, 1024, 512, false, 2)); + test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 64, 64, 128, false, 2)); + test_cases.emplace_back(new test_flash_attn_ext_top_k( 768, 64, 64, 128, true, 2)); + test_cases.emplace_back(new test_flash_attn_ext_top_k(4096, 300, 256, 512, false, 3)); + + for (int kv : { 127, 128, 129 }) { + test_cases.emplace_back(new test_lightning_indexer(128, 64, kv, 32, 4, 1, GGML_TYPE_F16)); + } + return test_cases; } #ifdef _MSC_VER @@ -11289,6 +11500,11 @@ static std::vector> make_test_cases_perf() { GGML_TYPE_F32, {n_kv, 512, 64, 1}, false, {2, 1, 0, 3})); } + for (ggml_type type_a : { GGML_TYPE_IQ2_XS, GGML_TYPE_IQ3_XXS }) { + test_cases.emplace_back(new test_mul_mat_id(type_a, GGML_TYPE_F32, 256, 6, false, 2048, 512, 4096)); + test_cases.emplace_back(new test_mul_mat_id(type_a, GGML_TYPE_F32, 256, 6, false, 4096, 512, 2048)); + } + // Conv2d: K=CRS=NPQ=4096 matmul performance uint32_t iwh_idx = 0; uint32_t kwh_idx = 1; @@ -11534,6 +11750,23 @@ static std::vector> make_test_cases_perf() { // Qwen3-VL-8B https://github.com/ggml-org/llama.cpp/issues/17012 test_cases.emplace_back(new test_flash_attn_ext(72, 72, 16, {1, 1}, 5776, 5776, false, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + // Qwen3-Coder-30B-A3B prefill at depth: hd128, 4 KV heads, GQA 8, ub 2048 + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 4, {8, 1}, 2048, 2048, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 4, {8, 1}, 6144, 2048, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 4, {8, 1}, 10240, 2048, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 4, {8, 1}, 10240, 512, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + // MALL-spill probe: same FLOPs, 32 distinct KV heads (no GQA) -> K/V footprint 8x (168MB > 32MB MALL) + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 32, {1, 1}, 10240, 2048, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + // KV-cache layout probe: same shape, K/V strided token-major like the real cache (heads interleaved) + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 4, {8, 1}, 10240, 2048, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 2, 1, 3})); + // Same, dense-permuted (exact model KV-cache layout; eligible for the f16 contiguize path) + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 4, {8, 1}, 10240, 2048, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 2, 1, 3}, false)); + // L2-residence probe: 1 KV head x GQA 32 (K/V 5.2MB fits L2) - distinguishes cache-BW-bound from issue-bound + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 1, {32, 1}, 10240, 2048, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + // cost-partition probes: no mask; f16 accumulate + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 4, {8, 1}, 10240, 2048, false, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + test_cases.emplace_back(new test_flash_attn_ext(128, 128, 4, {8, 1}, 10240, 2048, true, false, 0, 0, GGML_PREC_DEFAULT, GGML_TYPE_F16, GGML_TYPE_F16)); + test_cases.emplace_back(new test_flash_attn_ext(64, 64, 8, {8, 1}, 7680, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); test_cases.emplace_back(new test_flash_attn_ext(64, 64, 8, {8, 1}, 7680, 4, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); test_cases.emplace_back(new test_flash_attn_ext(64, 64, 8, {8, 1}, 7680, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q4_0, GGML_TYPE_Q4_0)); @@ -11710,7 +11943,7 @@ static std::vector> make_test_cases_perf() { test_cases.emplace_back(new test_gated_delta_net(GGML_TYPE_F32, 32, 128, 2048, 1, 1, false, true)); // lightning_indexer - for (int kv : { 256, 4096, 65536 }) { + for (int kv : { 256, 512, 4096, 65536 }) { for (int bs : { 1, 512, 2048 }) { for (int nh : { 32, 64 }) { for (int ns : { 1, 4 }) { @@ -11721,6 +11954,80 @@ static std::vector> make_test_cases_perf() { } } } + // DSV4 PP2048 indexer rows after filling source contexts from 8k through 512k. + // The zero-depth kv=512 shape is covered above. + for (int kv : { 2560, 4608, 8704, 16896, 33280, 66048, 131584 }) { + test_cases.emplace_back(new test_lightning_indexer(128, 64, kv, 2048, 1, 1, GGML_TYPE_F16)); + } + + // DSpark verify-step indexer rows: batch 1-32 across the 4-15 small-CM routing window, + // at the compressed key counts a 128k/491k source context produces (source/4). + for (int kv : { 8704, 33280, 131584 }) { + for (int bs : { 1, 2, 3, 4, 5, 6, 8, 12, 15, 16, 32 }) { + test_cases.emplace_back(new test_lightning_indexer(128, 64, kv, bs, 1, 1, GGML_TYPE_F16)); + } + } + + // sparse top-k FA at V4 decode/prefill shapes — the A/B instrument for the + // gather-to-compact work (n_active = n_kv_raw + n_top_k stays fixed as kv grows). + // nb 1/8 currently takes the DENSE path (the sparse shader gates on nb >= 64): + // those rows measure the decode cost gather-to-compact must beat. nb 64/512 + // measures the existing sparse prefill shader. + for (int kv : { 8192, 32768, 65536 }) { + for (int nb : { 1, 8, 64, 512 }) { + test_cases.emplace_back(new test_flash_attn_ext_top_k(kv, nb, 1024, 512, false)); + } + } + // PP2048 compressed-K rows for source context depths 32k through 512k. + for (int kv : { 11008, 19200, 35584, 68352, 133888 }) { + test_cases.emplace_back(new test_flash_attn_ext_top_k(kv, 2048, 2304, 512, false)); + } + // small-batch decode at depth: the speculative-draft regime (batch 2-8), where the old + // path fell through to dense attention over the whole compressed KV. + for (int kv : { 11008, 35584, 133888 }) { + for (int nb : { 1, 2, 4, 8 }) { + test_cases.emplace_back(new test_flash_attn_ext_top_k(kv, nb, 2304, 512, false)); + } + } + // Same shapes with realistic adjacent-token overlap. Measured on DeepSeek-V4-Flash the + // real overlap is 60% over 4 adjacent tokens and 76% over 8; the default generator is + // near 0%, which would make a deduplicated union look worthless by construction. + // kv=11008 (~32k source) and nb=16 are where the compaction gate is tightest: the + // worst-case compact set 2304 + nb*512 crosses kv/2 at nb=6, so those cells measure + // whether the gate can be opened by the union rather than by the worst case. + for (int kv : { 11008, 35584, 133888 }) { + for (int nb : { 1, 2, 4, 8, 16 }) { + test_cases.emplace_back(new test_flash_attn_ext_top_k(kv, nb, 2304, 512, false, 1, 60)); + } + } + // ov is a per-token share, not the union/selected ratio the model was measured by: at nb + // tokens it gives a union of (ov + (1-ov)*nb)/nb of the selections, so ov=60 is 0.475 at + // nb=8 where the model measured 0.243. ov=86 is the setting that reproduces the model, and + // at kv=11008 it is the difference between a union that fits under the gate and one that + // does not. + for (int nb : { 8, 16 }) { + test_cases.emplace_back(new test_flash_attn_ext_top_k(11008, nb, 2304, 512, false, 1, 86)); + } + // q8_0 K/V at the same shapes: DSv4 with -ctk q8_0 took the dense fallback before the + // gather became type-agnostic, so this is the cell that says whether it now pays there. + // nb=1 is plain autoregressive decode: the union needs nb > 1 to dedup, so this width + // takes the per-token gather. Its f16 row is in the grid above. + for (ggml_type tk : { GGML_TYPE_Q8_0, GGML_TYPE_Q4_0 }) { + for (int nb : { 1, 2, 4, 8, 16 }) { + test_cases.emplace_back(new test_flash_attn_ext_top_k(11008, nb, 2304, 512, false, 1, 60, tk)); + } + } + // PREFILL widths with quantised K/V, which the sparse path now serves through a one-shot + // dequant into the f16 scratch; quantised should sit within ~1% of f16 at every kv here. + // nb=1024 is the reporting user's --ubatch-size; the kv list is their four source depths + // (17k/33k/67k/134k) in compressed-K rows. GGML_VK_FA_DEQUANT=0 reproduces the old dense + // fallback, whose gap GROWS with kv (dense is O(kv), sparse O(n_kv_raw + n_top_k)); f16 + // under GGML_VK_FA_TOPK=0 is the falsification arm for attributing that gap to the gate. + for (ggml_type tk : { GGML_TYPE_F16, GGML_TYPE_Q8_0, GGML_TYPE_Q4_0 }) { + for (int kv : { 5504, 11008, 19200, 35584 }) { + test_cases.emplace_back(new test_flash_attn_ext_top_k(kv, 1024, 2304, 512, false, 1, 0, tk)); + } + } return test_cases; }