whispercpp 1.3.6 → 1.3.8

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (965) hide show
  1. checksums.yaml +4 -4
  2. data/.document +3 -0
  3. data/.rdoc_options +2 -0
  4. data/README.md +43 -9
  5. data/Rakefile +18 -3
  6. data/ext/dependencies.rb +10 -4
  7. data/ext/dependencies_for_windows.rb +17 -0
  8. data/ext/extconf.rb +20 -8
  9. data/ext/options.rb +54 -14
  10. data/ext/options_for_windows.rb +51 -0
  11. data/ext/ruby_whisper.c +35 -42
  12. data/ext/ruby_whisper.h +141 -0
  13. data/ext/ruby_whisper_context.c +157 -29
  14. data/ext/ruby_whisper_log_queue.c +180 -0
  15. data/ext/ruby_whisper_log_settable.h +46 -0
  16. data/ext/ruby_whisper_parakeet.c +49 -0
  17. data/ext/ruby_whisper_parakeet_context.c +304 -0
  18. data/ext/ruby_whisper_parakeet_context_params.c +117 -0
  19. data/ext/ruby_whisper_parakeet_model.c +84 -0
  20. data/ext/ruby_whisper_parakeet_params.c +548 -0
  21. data/ext/ruby_whisper_parakeet_segment.c +157 -0
  22. data/ext/ruby_whisper_parakeet_token.c +188 -0
  23. data/ext/ruby_whisper_parakeet_transcribe.cpp +58 -0
  24. data/ext/ruby_whisper_params.c +265 -73
  25. data/ext/ruby_whisper_segment.c +6 -6
  26. data/ext/ruby_whisper_transcribe.cpp +23 -15
  27. data/ext/ruby_whisper_vad_context.c +30 -10
  28. data/ext/ruby_whisper_vad_context_detect.cpp +8 -9
  29. data/ext/ruby_whisper_vad_params.c +4 -4
  30. data/ext/ruby_whisper_vad_segment.c +2 -2
  31. data/ext/sources/CMakeLists.txt +42 -3
  32. data/ext/sources/CMakePresets.json +95 -0
  33. data/ext/sources/cmake/parakeet-config.cmake.in +30 -0
  34. data/ext/sources/cmake/parakeet.pc.in +10 -0
  35. data/ext/sources/cmake/whisper.pc.in +2 -2
  36. data/ext/sources/examples/CMakeLists.txt +4 -2
  37. data/ext/sources/examples/bench/bench.cpp +1 -1
  38. data/ext/sources/examples/cli/cli.cpp +52 -10
  39. data/ext/sources/examples/common-ggml.cpp +4 -0
  40. data/ext/sources/examples/common-whisper.cpp +139 -67
  41. data/ext/sources/examples/common-whisper.h +11 -0
  42. data/ext/sources/examples/ffmpeg-transcode.cpp +211 -341
  43. data/ext/sources/examples/parakeet-cli/CMakeLists.txt +8 -0
  44. data/ext/sources/examples/parakeet-cli/parakeet-cli.cpp +243 -0
  45. data/ext/sources/examples/parakeet-quantize/CMakeLists.txt +7 -0
  46. data/ext/sources/examples/parakeet-quantize/parakeet-quantize.cpp +230 -0
  47. data/ext/sources/examples/server/server.cpp +199 -163
  48. data/ext/sources/examples/vad-speech-segments/speech.cpp +3 -2
  49. data/ext/sources/ggml/CMakeLists.txt +21 -14
  50. data/ext/sources/ggml/cmake/FindNCCL.cmake +36 -0
  51. data/ext/sources/ggml/cmake/ggml-config.cmake.in +12 -2
  52. data/ext/sources/ggml/include/ggml-alloc.h +1 -0
  53. data/ext/sources/ggml/include/ggml-backend.h +72 -10
  54. data/ext/sources/ggml/include/ggml-cuda.h +2 -2
  55. data/ext/sources/ggml/include/ggml-rpc.h +3 -3
  56. data/ext/sources/ggml/include/ggml-sycl.h +8 -0
  57. data/ext/sources/ggml/include/ggml.h +103 -9
  58. data/ext/sources/ggml/include/gguf.h +10 -2
  59. data/ext/sources/ggml/src/CMakeLists.txt +30 -6
  60. data/ext/sources/ggml/src/ggml-alloc.c +5 -1
  61. data/ext/sources/ggml/src/ggml-backend-impl.h +22 -2
  62. data/ext/sources/ggml/src/ggml-backend-meta.cpp +2266 -0
  63. data/ext/sources/ggml/src/ggml-backend-reg.cpp +12 -0
  64. data/ext/sources/ggml/src/ggml-backend.cpp +110 -9
  65. data/ext/sources/ggml/src/ggml-blas/ggml-blas.cpp +4 -0
  66. data/ext/sources/ggml/src/ggml-cann/aclnn_ops.cpp +672 -257
  67. data/ext/sources/ggml/src/ggml-cann/aclnn_ops.h +71 -0
  68. data/ext/sources/ggml/src/ggml-cann/common.h +20 -10
  69. data/ext/sources/ggml/src/ggml-cann/ggml-cann.cpp +211 -30
  70. data/ext/sources/ggml/src/ggml-common.h +24 -2
  71. data/ext/sources/ggml/src/ggml-cpu/CMakeLists.txt +59 -30
  72. data/ext/sources/ggml/src/ggml-cpu/amx/amx.cpp +2 -0
  73. data/ext/sources/ggml/src/ggml-cpu/amx/mmq.cpp +21 -22
  74. data/ext/sources/ggml/src/ggml-cpu/arch/arm/quants.c +194 -11
  75. data/ext/sources/ggml/src/ggml-cpu/arch/arm/repack.cpp +65 -0
  76. data/ext/sources/ggml/src/ggml-cpu/arch/loongarch/quants.c +151 -1
  77. data/ext/sources/ggml/src/ggml-cpu/arch/powerpc/quants.c +0 -1
  78. data/ext/sources/ggml/src/ggml-cpu/arch/riscv/quants.c +4279 -1292
  79. data/ext/sources/ggml/src/ggml-cpu/arch/riscv/repack.cpp +5 -35
  80. data/ext/sources/ggml/src/ggml-cpu/arch/s390/quants.c +0 -1
  81. data/ext/sources/ggml/src/ggml-cpu/arch/wasm/quants.c +72 -1
  82. data/ext/sources/ggml/src/ggml-cpu/arch/x86/quants.c +319 -31
  83. data/ext/sources/ggml/src/ggml-cpu/arch/x86/repack.cpp +1 -1
  84. data/ext/sources/ggml/src/ggml-cpu/arch-fallback.h +12 -2
  85. data/ext/sources/ggml/src/ggml-cpu/cmake/FindSMTIME.cmake +32 -0
  86. data/ext/sources/ggml/src/ggml-cpu/ggml-cpu-impl.h +10 -0
  87. data/ext/sources/ggml/src/ggml-cpu/ggml-cpu.c +109 -5
  88. data/ext/sources/ggml/src/ggml-cpu/ggml-cpu.cpp +2 -0
  89. data/ext/sources/ggml/src/ggml-cpu/kleidiai/kleidiai.cpp +146 -134
  90. data/ext/sources/ggml/src/ggml-cpu/llamafile/sgemm.cpp +107 -82
  91. data/ext/sources/ggml/src/ggml-cpu/ops.cpp +501 -119
  92. data/ext/sources/ggml/src/ggml-cpu/ops.h +3 -0
  93. data/ext/sources/ggml/src/ggml-cpu/quants.c +106 -0
  94. data/ext/sources/ggml/src/ggml-cpu/quants.h +6 -0
  95. data/ext/sources/ggml/src/ggml-cpu/repack.cpp +3 -0
  96. data/ext/sources/ggml/src/ggml-cpu/simd-gemm.h +91 -1
  97. data/ext/sources/ggml/src/ggml-cpu/simd-mappings.h +14 -16
  98. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime.cpp +1402 -687
  99. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime.h +8 -0
  100. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime1_kernels.cpp +597 -2766
  101. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime2_kernels.cpp +5768 -0
  102. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime_env.cpp +320 -0
  103. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime_env.h +55 -0
  104. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime_kernels.h +182 -19
  105. data/ext/sources/ggml/src/ggml-cpu/spacemit/repack.cpp +1795 -0
  106. data/ext/sources/ggml/src/ggml-cpu/spacemit/repack.h +14 -0
  107. data/ext/sources/ggml/src/ggml-cpu/spacemit/rvv_kernels.cpp +3178 -0
  108. data/ext/sources/ggml/src/ggml-cpu/spacemit/rvv_kernels.h +95 -0
  109. data/ext/sources/ggml/src/ggml-cpu/spacemit/spine_barrier.h +34 -0
  110. data/ext/sources/ggml/src/ggml-cpu/spacemit/spine_mem_pool.cpp +760 -0
  111. data/ext/sources/ggml/src/ggml-cpu/spacemit/spine_mem_pool.h +32 -0
  112. data/ext/sources/ggml/src/ggml-cpu/spacemit/spine_tcm.h +409 -0
  113. data/ext/sources/ggml/src/ggml-cpu/vec.cpp +39 -55
  114. data/ext/sources/ggml/src/ggml-cpu/vec.h +225 -240
  115. data/ext/sources/ggml/src/ggml-cuda/CMakeLists.txt +17 -7
  116. data/ext/sources/ggml/src/ggml-cuda/allreduce.cu +971 -0
  117. data/ext/sources/ggml/src/ggml-cuda/allreduce.cuh +29 -0
  118. data/ext/sources/ggml/src/ggml-cuda/argsort.cu +62 -26
  119. data/ext/sources/ggml/src/ggml-cuda/binbcast.cu +134 -64
  120. data/ext/sources/ggml/src/ggml-cuda/binbcast.cuh +1 -0
  121. data/ext/sources/ggml/src/ggml-cuda/col2im-1d.cu +81 -0
  122. data/ext/sources/ggml/src/ggml-cuda/col2im-1d.cuh +3 -0
  123. data/ext/sources/ggml/src/ggml-cuda/common.cuh +246 -28
  124. data/ext/sources/ggml/src/ggml-cuda/concat.cu +134 -116
  125. data/ext/sources/ggml/src/ggml-cuda/conv-transpose-1d.cu +14 -12
  126. data/ext/sources/ggml/src/ggml-cuda/conv2d-transpose.cu +45 -21
  127. data/ext/sources/ggml/src/ggml-cuda/conv2d-transpose.cuh +1 -0
  128. data/ext/sources/ggml/src/ggml-cuda/convert.cu +139 -34
  129. data/ext/sources/ggml/src/ggml-cuda/convert.cuh +10 -0
  130. data/ext/sources/ggml/src/ggml-cuda/cpy.cu +88 -29
  131. data/ext/sources/ggml/src/ggml-cuda/dequantize.cuh +22 -0
  132. data/ext/sources/ggml/src/ggml-cuda/fattn-common.cuh +287 -49
  133. data/ext/sources/ggml/src/ggml-cuda/fattn-mma-f16.cuh +335 -130
  134. data/ext/sources/ggml/src/ggml-cuda/fattn-tile.cu +12 -0
  135. data/ext/sources/ggml/src/ggml-cuda/fattn-tile.cuh +127 -24
  136. data/ext/sources/ggml/src/ggml-cuda/fattn-vec.cuh +40 -15
  137. data/ext/sources/ggml/src/ggml-cuda/fattn-wmma-f16.cu +18 -9
  138. data/ext/sources/ggml/src/ggml-cuda/fattn.cu +169 -60
  139. data/ext/sources/ggml/src/ggml-cuda/fattn.cuh +2 -0
  140. data/ext/sources/ggml/src/ggml-cuda/fwht.cu +101 -0
  141. data/ext/sources/ggml/src/ggml-cuda/fwht.cuh +4 -0
  142. data/ext/sources/ggml/src/ggml-cuda/gated_delta_net.cu +109 -45
  143. data/ext/sources/ggml/src/ggml-cuda/gated_delta_net.cuh +10 -0
  144. data/ext/sources/ggml/src/ggml-cuda/getrows.cu +48 -23
  145. data/ext/sources/ggml/src/ggml-cuda/ggml-cuda.cu +2034 -2104
  146. data/ext/sources/ggml/src/ggml-cuda/im2col.cu +32 -29
  147. data/ext/sources/ggml/src/ggml-cuda/mean.cu +4 -2
  148. data/ext/sources/ggml/src/ggml-cuda/mma.cuh +242 -195
  149. data/ext/sources/ggml/src/ggml-cuda/mmf.cuh +3 -3
  150. data/ext/sources/ggml/src/ggml-cuda/mmq.cu +25 -12
  151. data/ext/sources/ggml/src/ggml-cuda/mmq.cuh +502 -423
  152. data/ext/sources/ggml/src/ggml-cuda/mmvf.cu +19 -12
  153. data/ext/sources/ggml/src/ggml-cuda/mmvq.cu +562 -97
  154. data/ext/sources/ggml/src/ggml-cuda/mmvq.cuh +6 -1
  155. data/ext/sources/ggml/src/ggml-cuda/norm.cu +36 -10
  156. data/ext/sources/ggml/src/ggml-cuda/out-prod.cu +66 -7
  157. data/ext/sources/ggml/src/ggml-cuda/quantize.cu +133 -26
  158. data/ext/sources/ggml/src/ggml-cuda/quantize.cuh +1 -1
  159. data/ext/sources/ggml/src/ggml-cuda/reduce_rows.cuh +5 -1
  160. data/ext/sources/ggml/src/ggml-cuda/rope.cu +11 -4
  161. data/ext/sources/ggml/src/ggml-cuda/scale.cu +4 -1
  162. data/ext/sources/ggml/src/ggml-cuda/set-rows.cu +78 -10
  163. data/ext/sources/ggml/src/ggml-cuda/snake.cu +72 -0
  164. data/ext/sources/ggml/src/ggml-cuda/snake.cuh +8 -0
  165. data/ext/sources/ggml/src/ggml-cuda/softcap.cu +4 -1
  166. data/ext/sources/ggml/src/ggml-cuda/ssm-conv.cu +45 -13
  167. data/ext/sources/ggml/src/ggml-cuda/ssm-conv.cuh +1 -1
  168. data/ext/sources/ggml/src/ggml-cuda/ssm-scan.cu +40 -18
  169. data/ext/sources/ggml/src/ggml-cuda/sumrows.cu +8 -4
  170. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_1-ncols2_16.cu +1 -0
  171. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_1-ncols2_32.cu +1 -0
  172. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_1-ncols2_8.cu +2 -0
  173. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_16-ncols2_2.cu +1 -0
  174. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_16-ncols2_4.cu +1 -0
  175. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_2-ncols2_16.cu +1 -0
  176. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_2-ncols2_32.cu +1 -0
  177. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_2-ncols2_4.cu +1 -0
  178. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_2-ncols2_8.cu +2 -0
  179. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_32-ncols2_2.cu +1 -0
  180. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_16.cu +1 -0
  181. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_2.cu +1 -0
  182. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_4.cu +1 -0
  183. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_8.cu +2 -0
  184. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_8-ncols2_2.cu +1 -0
  185. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_8-ncols2_4.cu +1 -0
  186. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_8-ncols2_8.cu +2 -0
  187. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-tile-instance-dkq192-dv128.cu +5 -0
  188. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-tile-instance-dkq320-dv256.cu +5 -0
  189. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-tile-instance-dkq512-dv512.cu +5 -0
  190. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-bf16.cu +7 -0
  191. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-f16.cu +7 -0
  192. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q4_0.cu +7 -0
  193. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q4_1.cu +7 -0
  194. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q5_0.cu +7 -0
  195. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q5_1.cu +7 -0
  196. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q8_0.cu +7 -0
  197. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-f16-bf16.cu +7 -0
  198. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q4_0-bf16.cu +7 -0
  199. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q4_1-bf16.cu +7 -0
  200. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q5_0-bf16.cu +7 -0
  201. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q5_1-bf16.cu +7 -0
  202. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q8_0-bf16.cu +7 -0
  203. data/ext/sources/ggml/src/ggml-cuda/template-instances/mmq-instance-nvfp4.cu +5 -0
  204. data/ext/sources/ggml/src/ggml-cuda/template-instances/mmq-instance-q1_0.cu +5 -0
  205. data/ext/sources/ggml/src/ggml-cuda/top-k.cu +5 -4
  206. data/ext/sources/ggml/src/ggml-cuda/topk-moe.cu +33 -24
  207. data/ext/sources/ggml/src/ggml-cuda/unary.cu +31 -2
  208. data/ext/sources/ggml/src/ggml-cuda/unary.cuh +2 -0
  209. data/ext/sources/ggml/src/ggml-cuda/vecdotq.cuh +80 -0
  210. data/ext/sources/ggml/src/ggml-cuda/vendors/cuda.h +7 -2
  211. data/ext/sources/ggml/src/ggml-cuda/vendors/hip.h +23 -4
  212. data/ext/sources/ggml/src/ggml-cuda/vendors/musa.h +4 -0
  213. data/ext/sources/ggml/src/ggml-hexagon/CMakeLists.txt +1 -5
  214. data/ext/sources/ggml/src/ggml-hexagon/ggml-hexagon.cpp +2788 -1762
  215. data/ext/sources/ggml/src/ggml-hexagon/htp/CMakeLists.txt +13 -4
  216. data/ext/sources/ggml/src/ggml-hexagon/htp/act-ops.c +53 -84
  217. data/ext/sources/ggml/src/ggml-hexagon/htp/argsort-ops.c +25 -12
  218. data/ext/sources/ggml/src/ggml-hexagon/htp/binary-ops.c +165 -184
  219. data/ext/sources/ggml/src/ggml-hexagon/htp/cmake-toolchain.cmake +17 -19
  220. data/ext/sources/ggml/src/ggml-hexagon/htp/concat-ops.c +277 -0
  221. data/ext/sources/ggml/src/ggml-hexagon/htp/cpy-ops.c +170 -127
  222. data/ext/sources/ggml/src/ggml-hexagon/htp/cumsum-ops.c +270 -0
  223. data/ext/sources/ggml/src/ggml-hexagon/htp/diag-ops.c +216 -0
  224. data/ext/sources/ggml/src/ggml-hexagon/htp/fill-ops.c +123 -0
  225. data/ext/sources/ggml/src/ggml-hexagon/htp/flash-attn-ops.c +1774 -396
  226. data/ext/sources/ggml/src/ggml-hexagon/htp/flash-attn-ops.h +303 -0
  227. data/ext/sources/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.c +1148 -0
  228. data/ext/sources/ggml/src/ggml-hexagon/htp/get-rows-ops.c +148 -42
  229. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-common.h +80 -0
  230. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-dma.c +2 -2
  231. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-dma.h +255 -62
  232. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-dump.h +9 -0
  233. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-profile.h +64 -0
  234. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-utils.h +25 -21
  235. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-fa-kernels.h +555 -0
  236. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-mm-kernels-tiled.h +1303 -0
  237. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-queue.c +167 -0
  238. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-queue.h +157 -0
  239. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-utils.h +222 -0
  240. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-ctx.h +104 -13
  241. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-ops.h +222 -57
  242. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-vtcm.h +19 -0
  243. data/ext/sources/ggml/src/ggml-hexagon/htp/htp_iface.idl +10 -3
  244. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-base.h +78 -26
  245. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-copy.h +27 -10
  246. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-div.h +63 -23
  247. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-exp.h +48 -8
  248. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-fa-kernels.h +232 -0
  249. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-flash-attn.h +47 -0
  250. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-log.h +65 -0
  251. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-flat.h +1511 -0
  252. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-tiled.h +1200 -0
  253. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-pow.h +42 -0
  254. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-repl.h +74 -0
  255. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-sigmoid.h +40 -0
  256. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-sin-cos.h +90 -0
  257. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-utils.h +5 -8
  258. data/ext/sources/ggml/src/ggml-hexagon/htp/main.c +625 -816
  259. data/ext/sources/ggml/src/ggml-hexagon/htp/matmul-ops.c +3052 -2166
  260. data/ext/sources/ggml/src/ggml-hexagon/htp/matmul-ops.h +650 -0
  261. data/ext/sources/ggml/src/ggml-hexagon/htp/pad-ops.c +547 -0
  262. data/ext/sources/ggml/src/ggml-hexagon/htp/repeat-ops.c +148 -0
  263. data/ext/sources/ggml/src/ggml-hexagon/htp/rope-ops.c +337 -106
  264. data/ext/sources/ggml/src/ggml-hexagon/htp/set-rows-ops.c +59 -37
  265. data/ext/sources/ggml/src/ggml-hexagon/htp/softmax-ops.c +121 -133
  266. data/ext/sources/ggml/src/ggml-hexagon/htp/solve-tri-ops.c +267 -0
  267. data/ext/sources/ggml/src/ggml-hexagon/htp/ssm-conv.c +245 -151
  268. data/ext/sources/ggml/src/ggml-hexagon/htp/sum-rows-ops.c +6 -6
  269. data/ext/sources/ggml/src/ggml-hexagon/htp/unary-ops.c +719 -45
  270. data/ext/sources/ggml/src/ggml-hexagon/htp/worker-pool.c +15 -3
  271. data/ext/sources/ggml/src/ggml-hexagon/htp/worker-pool.h +8 -0
  272. data/ext/sources/ggml/src/ggml-hexagon/htp-opnode.h +390 -0
  273. data/ext/sources/ggml/src/ggml-hexagon/libggml-htp.inf +3 -5
  274. data/ext/sources/ggml/src/ggml-hip/CMakeLists.txt +27 -9
  275. data/ext/sources/ggml/src/ggml-impl.h +6 -1
  276. data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.cpp +207 -18
  277. data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.h +36 -2
  278. data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.m +186 -29
  279. data/ext/sources/ggml/src/ggml-metal/ggml-metal-impl.h +118 -0
  280. data/ext/sources/ggml/src/ggml-metal/ggml-metal-ops.cpp +322 -21
  281. data/ext/sources/ggml/src/ggml-metal/ggml-metal-ops.h +4 -0
  282. data/ext/sources/ggml/src/ggml-metal/ggml-metal.cpp +39 -26
  283. data/ext/sources/ggml/src/ggml-metal/ggml-metal.metal +1226 -467
  284. data/ext/sources/ggml/src/ggml-musa/CMakeLists.txt +5 -6
  285. data/ext/sources/ggml/src/ggml-opencl/CMakeLists.txt +67 -5
  286. data/ext/sources/ggml/src/ggml-opencl/fa_tune.h +92 -0
  287. data/ext/sources/ggml/src/ggml-opencl/ggml-opencl.cpp +16290 -6246
  288. data/ext/sources/ggml/src/ggml-opencl/kernels/concat.cl +67 -0
  289. data/ext/sources/ggml/src/ggml-opencl/kernels/cpy.cl +59 -0
  290. data/ext/sources/ggml/src/ggml-opencl/kernels/cvt.cl +1997 -92
  291. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f16.cl +81 -41
  292. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32.cl +88 -39
  293. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_f16.cl +1995 -96
  294. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_q4_0.cl +1615 -0
  295. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_q8_0.cl +1486 -0
  296. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_pre_f16.cl +156 -0
  297. data/ext/sources/ggml/src/ggml-opencl/kernels/gated_delta_net.cl +249 -0
  298. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_mxfp4_f32_ns.cl +374 -0
  299. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_0_f32_ns.cl +324 -0
  300. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_1_f32_ns.cl +326 -0
  301. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_k_f32_ns.cl +348 -0
  302. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_0_f32_ns.cl +328 -0
  303. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_1_f32_ns.cl +330 -0
  304. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_k_f32_ns.cl +356 -0
  305. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q6_k_f32_ns.cl +335 -0
  306. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_iq4_nl_f32.cl +150 -0
  307. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q1_0_f32.cl +94 -0
  308. data/ext/sources/ggml/src/ggml-opencl/kernels/{mul_mat_Ab_Bi_8x4.cl → gemm_noshuffle_q4_0_f32.cl} +1 -1
  309. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q4_k_f32.cl +172 -0
  310. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q5_0_f32.cl +131 -0
  311. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q5_1_f32.cl +134 -0
  312. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q5_k_f32.cl +176 -0
  313. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q6_k_f32.cl +140 -0
  314. data/ext/sources/ggml/src/ggml-opencl/kernels/{mul_mm_q8_0_f32_8x4.cl → gemm_noshuffle_q8_0_f32.cl} +1 -1
  315. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_xmem_f16_f32_os8.cl +233 -0
  316. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_mxfp4_f32_ns.cl +165 -0
  317. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q4_0_f32_ns.cl +120 -0
  318. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q4_1_f32_ns.cl +123 -0
  319. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q4_k_f32_ns.cl +155 -0
  320. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q5_0_f32_ns.cl +123 -0
  321. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q5_1_f32_ns.cl +125 -0
  322. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q5_k_f32_ns.cl +160 -0
  323. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q6_k_f32_ns.cl +141 -0
  324. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_iq4_nl_f32.cl +302 -0
  325. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q1_0_f32.cl +121 -0
  326. data/ext/sources/ggml/src/ggml-opencl/kernels/{gemv_noshuffle_general.cl → gemv_noshuffle_q4_0_f32.cl} +5 -5
  327. data/ext/sources/ggml/src/ggml-opencl/kernels/{gemv_noshuffle.cl → gemv_noshuffle_q4_0_f32_spec.cl} +5 -5
  328. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_k_f32.cl +318 -0
  329. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q5_0_f32.cl +291 -0
  330. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q5_1_f32.cl +294 -0
  331. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q5_k_f32.cl +326 -0
  332. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q6_k_f32.cl +293 -0
  333. data/ext/sources/ggml/src/ggml-opencl/kernels/{gemv_noshuffle_general_q8_0_f32.cl → gemv_noshuffle_q8_0_f32.cl} +1 -1
  334. data/ext/sources/ggml/src/ggml-opencl/kernels/get_rows.cl +15 -9
  335. data/ext/sources/ggml/src/ggml-opencl/kernels/moe_reorder_b.cl +30 -0
  336. data/ext/sources/ggml/src/ggml-opencl/kernels/moe_sort_by_expert.cl +82 -0
  337. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_iq4_nl_f32_l4_lm.cl +171 -0
  338. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q1_0_f32_l4_lm.cl +156 -0
  339. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q4_k_f32_l4_lm.cl +179 -0
  340. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q5_0_f32_l4_lm.cl +173 -0
  341. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q5_1_f32_l4_lm.cl +175 -0
  342. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q5_k_f32_l4_lm.cl +192 -0
  343. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_f16_f32_l4.cl +1149 -0
  344. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_iq4_nl_f32.cl +164 -0
  345. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_iq4_nl_f32_flat.cl +202 -0
  346. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q1_0_f32.cl +141 -0
  347. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q1_0_f32_flat.cl +190 -0
  348. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q4_k_f32_flat.cl +196 -0
  349. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_0_f32.cl +241 -0
  350. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_0_f32_flat.cl +243 -0
  351. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_1_f32.cl +243 -0
  352. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_1_f32_flat.cl +247 -0
  353. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_k_f32.cl +187 -0
  354. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_k_f32_flat.cl +203 -0
  355. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q6_k_f32_flat.cl +48 -64
  356. data/ext/sources/ggml/src/ggml-opencl/kernels/norm.cl +5 -2
  357. data/ext/sources/ggml/src/ggml-opencl/kernels/set_rows.cl +500 -0
  358. data/ext/sources/ggml/src/ggml-opencl/libdl.h +79 -0
  359. data/ext/sources/ggml/src/ggml-openvino/.clang-format +0 -5
  360. data/ext/sources/ggml/src/ggml-openvino/CMakeLists.txt +2 -4
  361. data/ext/sources/ggml/src/ggml-openvino/ggml-decoder.cpp +740 -127
  362. data/ext/sources/ggml/src/ggml-openvino/ggml-decoder.h +76 -23
  363. data/ext/sources/ggml/src/ggml-openvino/ggml-openvino-extra.cpp +75 -14
  364. data/ext/sources/ggml/src/ggml-openvino/ggml-openvino-extra.h +29 -8
  365. data/ext/sources/ggml/src/ggml-openvino/ggml-openvino.cpp +339 -69
  366. data/ext/sources/ggml/src/ggml-openvino/ggml-quants.cpp +330 -192
  367. data/ext/sources/ggml/src/ggml-openvino/ggml-quants.h +10 -4
  368. data/ext/sources/ggml/src/ggml-openvino/openvino/decoder.h +56 -16
  369. data/ext/sources/ggml/src/ggml-openvino/openvino/frontend.h +1 -1
  370. data/ext/sources/ggml/src/ggml-openvino/openvino/input_model.h +4 -4
  371. data/ext/sources/ggml/src/ggml-openvino/openvino/node_context.h +94 -37
  372. data/ext/sources/ggml/src/ggml-openvino/openvino/op/add_id.cpp +76 -0
  373. data/ext/sources/ggml/src/ggml-openvino/openvino/op/argsort.cpp +47 -0
  374. data/ext/sources/ggml/src/ggml-openvino/openvino/op/clamp.cpp +33 -0
  375. data/ext/sources/ggml/src/ggml-openvino/openvino/op/concat.cpp +48 -0
  376. data/ext/sources/ggml/src/ggml-openvino/openvino/op/cont.cpp +8 -16
  377. data/ext/sources/ggml/src/ggml-openvino/openvino/op/cpy.cpp +14 -1
  378. data/ext/sources/ggml/src/ggml-openvino/openvino/op/div.cpp +146 -0
  379. data/ext/sources/ggml/src/ggml-openvino/openvino/op/flash_attn_ext.cpp +108 -21
  380. data/ext/sources/ggml/src/ggml-openvino/openvino/op/gated_delta_net.cpp +282 -0
  381. data/ext/sources/ggml/src/ggml-openvino/openvino/op/gated_delta_net.hpp +65 -0
  382. data/ext/sources/ggml/src/ggml-openvino/openvino/op/get_rows.cpp +2 -9
  383. data/ext/sources/ggml/src/ggml-openvino/openvino/op/glu_geglu.cpp +21 -7
  384. data/ext/sources/ggml/src/ggml-openvino/openvino/op/glu_swiglu.cpp +41 -8
  385. data/ext/sources/ggml/src/ggml-openvino/openvino/op/im2col.cpp +120 -0
  386. data/ext/sources/ggml/src/ggml-openvino/openvino/op/l2_norm.cpp +44 -0
  387. data/ext/sources/ggml/src/ggml-openvino/openvino/op/mul_mat_id.cpp +226 -0
  388. data/ext/sources/ggml/src/ggml-openvino/openvino/op/mulmat.cpp +19 -9
  389. data/ext/sources/ggml/src/ggml-openvino/openvino/op/norm.cpp +58 -0
  390. data/ext/sources/ggml/src/ggml-openvino/openvino/op/pad.cpp +95 -0
  391. data/ext/sources/ggml/src/ggml-openvino/openvino/op/permute.cpp +58 -13
  392. data/ext/sources/ggml/src/ggml-openvino/openvino/op/repeat.cpp +74 -0
  393. data/ext/sources/ggml/src/ggml-openvino/openvino/op/reshape.cpp +13 -6
  394. data/ext/sources/ggml/src/ggml-openvino/openvino/op/rms_norm.cpp +1 -1
  395. data/ext/sources/ggml/src/ggml-openvino/openvino/op/rope.cpp +161 -39
  396. data/ext/sources/ggml/src/ggml-openvino/openvino/op/set_rows.cpp +3 -3
  397. data/ext/sources/ggml/src/ggml-openvino/openvino/op/softmax.cpp +126 -49
  398. data/ext/sources/ggml/src/ggml-openvino/openvino/op/ssm_conv.cpp +59 -0
  399. data/ext/sources/ggml/src/ggml-openvino/openvino/op/sum_rows.cpp +27 -0
  400. data/ext/sources/ggml/src/ggml-openvino/openvino/op/transpose.cpp +32 -1
  401. data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_silu.cpp +1 -1
  402. data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_softplus.cpp +38 -0
  403. data/ext/sources/ggml/src/ggml-openvino/openvino/op/view.cpp +90 -25
  404. data/ext/sources/ggml/src/ggml-openvino/openvino/op_table.cpp +41 -22
  405. data/ext/sources/ggml/src/ggml-openvino/openvino/op_table.h +18 -4
  406. data/ext/sources/ggml/src/ggml-openvino/openvino/pass/mark_decompression_convert_constant_folding.h +1 -1
  407. data/ext/sources/ggml/src/ggml-openvino/openvino/rt_info/weightless_caching_attributes.hpp +41 -0
  408. data/ext/sources/ggml/src/ggml-openvino/openvino/translate_session.cpp +70 -43
  409. data/ext/sources/ggml/src/ggml-openvino/openvino/translate_session.h +5 -4
  410. data/ext/sources/ggml/src/ggml-openvino/openvino/utils.cpp +612 -36
  411. data/ext/sources/ggml/src/ggml-openvino/openvino/utils.h +29 -26
  412. data/ext/sources/ggml/src/ggml-openvino/utils.cpp +460 -114
  413. data/ext/sources/ggml/src/ggml-openvino/utils.h +32 -9
  414. data/ext/sources/ggml/src/ggml-opt.cpp +1 -0
  415. data/ext/sources/ggml/src/ggml-quants.c +365 -114
  416. data/ext/sources/ggml/src/ggml-quants.h +6 -0
  417. data/ext/sources/ggml/src/ggml-rpc/CMakeLists.txt +24 -0
  418. data/ext/sources/ggml/src/ggml-rpc/ggml-rpc.cpp +167 -311
  419. data/ext/sources/ggml/src/ggml-rpc/transport.cpp +683 -0
  420. data/ext/sources/ggml/src/ggml-rpc/transport.h +34 -0
  421. data/ext/sources/ggml/src/ggml-sycl/CMakeLists.txt +50 -4
  422. data/ext/sources/ggml/src/ggml-sycl/add-id.cpp +1 -1
  423. data/ext/sources/ggml/src/ggml-sycl/backend.hpp +5 -1
  424. data/ext/sources/ggml/src/ggml-sycl/binbcast.cpp +12 -0
  425. data/ext/sources/ggml/src/ggml-sycl/col2im-1d.cpp +102 -0
  426. data/ext/sources/ggml/src/ggml-sycl/col2im-1d.hpp +8 -0
  427. data/ext/sources/ggml/src/ggml-sycl/common.cpp +72 -2
  428. data/ext/sources/ggml/src/ggml-sycl/common.hpp +59 -2
  429. data/ext/sources/ggml/src/ggml-sycl/concat.cpp +21 -1
  430. data/ext/sources/ggml/src/ggml-sycl/conv2d-dw.cpp +158 -0
  431. data/ext/sources/ggml/src/ggml-sycl/conv2d-dw.hpp +10 -0
  432. data/ext/sources/ggml/src/ggml-sycl/conv2d-transpose.cpp +125 -0
  433. data/ext/sources/ggml/src/ggml-sycl/conv2d-transpose.hpp +10 -0
  434. data/ext/sources/ggml/src/ggml-sycl/conv2d.cpp +150 -0
  435. data/ext/sources/ggml/src/ggml-sycl/conv2d.hpp +10 -0
  436. data/ext/sources/ggml/src/ggml-sycl/conv3d.cpp +224 -0
  437. data/ext/sources/ggml/src/ggml-sycl/conv3d.hpp +8 -0
  438. data/ext/sources/ggml/src/ggml-sycl/convert.cpp +121 -13
  439. data/ext/sources/ggml/src/ggml-sycl/convert.hpp +9 -0
  440. data/ext/sources/ggml/src/ggml-sycl/cpy.cpp +706 -0
  441. data/ext/sources/ggml/src/ggml-sycl/cpy.hpp +281 -0
  442. data/ext/sources/ggml/src/ggml-sycl/cross_entropy_loss.cpp +255 -0
  443. data/ext/sources/ggml/src/ggml-sycl/cross_entropy_loss.hpp +7 -0
  444. data/ext/sources/ggml/src/ggml-sycl/cumsum.cpp +148 -0
  445. data/ext/sources/ggml/src/ggml-sycl/cumsum.hpp +5 -0
  446. data/ext/sources/ggml/src/ggml-sycl/dequantize.hpp +678 -0
  447. data/ext/sources/ggml/src/ggml-sycl/diag.cpp +67 -0
  448. data/ext/sources/ggml/src/ggml-sycl/diag.hpp +5 -0
  449. data/ext/sources/ggml/src/ggml-sycl/dmmv.cpp +997 -244
  450. data/ext/sources/ggml/src/ggml-sycl/dpct/helper.hpp +15 -7
  451. data/ext/sources/ggml/src/ggml-sycl/element_wise.cpp +215 -204
  452. data/ext/sources/ggml/src/ggml-sycl/element_wise.hpp +2 -2
  453. data/ext/sources/ggml/src/ggml-sycl/fattn-buffers.cpp +56 -0
  454. data/ext/sources/ggml/src/ggml-sycl/fattn-buffers.hpp +63 -0
  455. data/ext/sources/ggml/src/ggml-sycl/fattn-common.hpp +7 -5
  456. data/ext/sources/ggml/src/ggml-sycl/fattn-tile.cpp +4 -0
  457. data/ext/sources/ggml/src/ggml-sycl/fattn-tile.hpp +76 -168
  458. data/ext/sources/ggml/src/ggml-sycl/fattn-vec.hpp +7 -0
  459. data/ext/sources/ggml/src/ggml-sycl/fattn.cpp +3 -1
  460. data/ext/sources/ggml/src/ggml-sycl/fill.cpp +55 -0
  461. data/ext/sources/ggml/src/ggml-sycl/fill.hpp +5 -0
  462. data/ext/sources/ggml/src/ggml-sycl/gated_delta_net.cpp +69 -31
  463. data/ext/sources/ggml/src/ggml-sycl/gated_delta_net.hpp +1 -0
  464. data/ext/sources/ggml/src/ggml-sycl/gemm.hpp +3 -0
  465. data/ext/sources/ggml/src/ggml-sycl/getrows.cpp +79 -3
  466. data/ext/sources/ggml/src/ggml-sycl/ggml-sycl.cpp +1758 -455
  467. data/ext/sources/ggml/src/ggml-sycl/im2col.cpp +353 -89
  468. data/ext/sources/ggml/src/ggml-sycl/im2col.hpp +5 -3
  469. data/ext/sources/ggml/src/ggml-sycl/mmvq.cpp +1542 -39
  470. data/ext/sources/ggml/src/ggml-sycl/mmvq.hpp +33 -0
  471. data/ext/sources/ggml/src/ggml-sycl/norm.cpp +103 -49
  472. data/ext/sources/ggml/src/ggml-sycl/outprod.cpp +45 -9
  473. data/ext/sources/ggml/src/ggml-sycl/pad.cpp +27 -27
  474. data/ext/sources/ggml/src/ggml-sycl/pool.cpp +185 -0
  475. data/ext/sources/ggml/src/ggml-sycl/pool.hpp +22 -0
  476. data/ext/sources/ggml/src/ggml-sycl/presets.hpp +3 -1
  477. data/ext/sources/ggml/src/ggml-sycl/quants.hpp +71 -0
  478. data/ext/sources/ggml/src/ggml-sycl/set_rows.cpp +17 -3
  479. data/ext/sources/ggml/src/ggml-sycl/softmax.cpp +9 -10
  480. data/ext/sources/ggml/src/ggml-sycl/solve_tri.cpp +172 -0
  481. data/ext/sources/ggml/src/ggml-sycl/solve_tri.hpp +8 -0
  482. data/ext/sources/ggml/src/ggml-sycl/ssm_conv.cpp +6 -1
  483. data/ext/sources/ggml/src/ggml-sycl/ssm_scan.cpp +156 -0
  484. data/ext/sources/ggml/src/ggml-sycl/ssm_scan.hpp +5 -0
  485. data/ext/sources/ggml/src/ggml-sycl/sycl_hw.cpp +62 -10
  486. data/ext/sources/ggml/src/ggml-sycl/sycl_hw.hpp +18 -6
  487. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-tile-instance-dkq512-dv512.cpp +6 -0
  488. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-f16.cpp +1 -0
  489. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q4_0.cpp +1 -0
  490. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q4_1.cpp +1 -0
  491. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q5_0.cpp +1 -0
  492. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q5_1.cpp +1 -0
  493. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q8_0.cpp +1 -0
  494. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-f16.cpp +1 -0
  495. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q4_0.cpp +1 -0
  496. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q4_1.cpp +1 -0
  497. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q5_0.cpp +1 -0
  498. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q5_1.cpp +1 -0
  499. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q8_0.cpp +1 -0
  500. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-f16.cpp +1 -0
  501. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q4_0.cpp +1 -0
  502. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q4_1.cpp +1 -0
  503. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q5_0.cpp +1 -0
  504. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q5_1.cpp +1 -0
  505. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q8_0.cpp +1 -0
  506. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-f16.cpp +1 -0
  507. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q4_0.cpp +1 -0
  508. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q4_1.cpp +1 -0
  509. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q5_0.cpp +1 -0
  510. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q5_1.cpp +1 -0
  511. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q8_0.cpp +1 -0
  512. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-f16.cpp +1 -0
  513. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q4_0.cpp +1 -0
  514. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q4_1.cpp +1 -0
  515. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q5_0.cpp +1 -0
  516. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q5_1.cpp +1 -0
  517. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q8_0.cpp +1 -0
  518. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-f16.cpp +1 -0
  519. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q4_0.cpp +1 -0
  520. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q4_1.cpp +1 -0
  521. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q5_0.cpp +1 -0
  522. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q5_1.cpp +1 -0
  523. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q8_0.cpp +1 -0
  524. data/ext/sources/ggml/src/ggml-sycl/type.hpp +112 -0
  525. data/ext/sources/ggml/src/ggml-sycl/upscale.cpp +410 -0
  526. data/ext/sources/ggml/src/ggml-sycl/upscale.hpp +9 -0
  527. data/ext/sources/ggml/src/ggml-sycl/vecdotq.hpp +242 -45
  528. data/ext/sources/ggml/src/ggml-virtgpu/ggml-backend-buffer.cpp +4 -0
  529. data/ext/sources/ggml/src/ggml-virtgpu/ggml-backend-device.cpp +2 -0
  530. data/ext/sources/ggml/src/ggml-virtgpu/ggml-backend.cpp +2 -0
  531. data/ext/sources/ggml/src/ggml-virtgpu/virtgpu-shm.cpp +1 -0
  532. data/ext/sources/ggml/src/ggml-virtgpu/virtgpu.cpp +1 -0
  533. data/ext/sources/ggml/src/ggml-virtgpu/virtgpu.h +0 -2
  534. data/ext/sources/ggml/src/ggml-vulkan/CMakeLists.txt +16 -0
  535. data/ext/sources/ggml/src/ggml-vulkan/ggml-vulkan.cpp +2843 -700
  536. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/CMakeLists.txt +4 -0
  537. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/col2im_1d.comp +61 -0
  538. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/contig_copy.comp +6 -2
  539. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/conv2d_mm.comp +146 -13
  540. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/conv3d_mm.comp +431 -0
  541. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/copy.comp +3 -1
  542. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/copy_from_quant.comp +1 -1
  543. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/copy_to_quant.comp +25 -1
  544. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs.glsl +88 -0
  545. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs_cm2.glsl +643 -1
  546. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dequant_nvfp4.comp +32 -0
  547. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q1_0.comp +29 -0
  548. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/diag.comp +3 -4
  549. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dot_product_funcs.glsl +27 -0
  550. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/feature-tests/coopmat2_decode_vector.comp +7 -0
  551. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp +198 -48
  552. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl +60 -59
  553. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp +116 -113
  554. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm2.comp +122 -31
  555. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_dequant.glsl +131 -0
  556. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_mmq_funcs.glsl +203 -0
  557. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/fwht.comp +115 -0
  558. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gated_delta_net.comp +125 -64
  559. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/generic_binary_head.glsl +0 -1
  560. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/generic_unary_head.glsl +21 -19
  561. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/get_rows_back.comp +25 -0
  562. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/glu_head.glsl +29 -1
  563. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/glu_main.glsl +17 -11
  564. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/im2col.comp +76 -54
  565. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/im2col_3d.comp +0 -1
  566. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/l2_norm.comp +4 -7
  567. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/log.comp +0 -1
  568. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec.comp +122 -27
  569. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_iface.glsl +6 -6
  570. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_q2_k.comp +1 -1
  571. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_q4_k.comp +1 -1
  572. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_q5_k.comp +1 -1
  573. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq.comp +22 -24
  574. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq_funcs.glsl +88 -55
  575. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp +42 -40
  576. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp +49 -15
  577. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl +222 -171
  578. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_funcs.glsl +8 -8
  579. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_shmem_types.glsl +24 -9
  580. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/multi_add.comp +0 -1
  581. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/norm.comp +10 -10
  582. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/repeat_back.comp +3 -3
  583. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/roll.comp +3 -3
  584. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/rope_funcs.glsl +5 -2
  585. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/rope_head.glsl +0 -1
  586. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/rope_params.glsl +3 -2
  587. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/snake.comp +49 -0
  588. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/ssm_conv.comp +11 -1
  589. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/tri.comp +3 -4
  590. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/types.glsl +79 -2
  591. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/unary.comp +168 -0
  592. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +282 -211
  593. data/ext/sources/ggml/src/ggml-webgpu/CMakeLists.txt +5 -2
  594. data/ext/sources/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp +2209 -283
  595. data/ext/sources/ggml/src/ggml-webgpu/ggml-webgpu.cpp +2618 -1416
  596. data/ext/sources/ggml/src/ggml-webgpu/pre_wgsl.hpp +37 -7
  597. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/add_id.wgsl +64 -0
  598. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/binary.wgsl +8 -7
  599. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/common_decls.tmpl +90 -95
  600. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/concat.wgsl +19 -1
  601. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/conv2d.wgsl +165 -0
  602. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/{cpy.tmpl.wgsl → cpy.wgsl} +25 -50
  603. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn.wgsl +107 -184
  604. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_quant_staging.tmpl +124 -0
  605. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_tile.wgsl +397 -0
  606. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_blk.wgsl +101 -0
  607. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_reduce.wgsl +84 -0
  608. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_split.wgsl +619 -0
  609. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/gated_delta_net.wgsl +149 -0
  610. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/get_rows.wgsl +204 -78
  611. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/glu.wgsl +155 -0
  612. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/im2col.wgsl +101 -0
  613. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_decls.tmpl +805 -526
  614. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_id.wgsl +195 -0
  615. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_id_gather.wgsl +52 -0
  616. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_id_vec.wgsl +154 -0
  617. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_reg_tile.wgsl +8 -6
  618. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_subgroup_matrix.wgsl +5 -1
  619. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec.wgsl +90 -413
  620. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_acc.tmpl +1553 -0
  621. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_q_acc.tmpl +297 -0
  622. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/quant_inner_loops.tmpl +21 -0
  623. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/quantize_q8.wgsl +178 -0
  624. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/rms_norm_mul.wgsl +152 -0
  625. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/{rope.tmpl.wgsl → rope.wgsl} +71 -142
  626. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/row_norm.wgsl +153 -0
  627. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/scale.wgsl +6 -4
  628. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/set.wgsl +109 -0
  629. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/set_rows.wgsl +2 -3
  630. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/set_rows_quant.wgsl +224 -0
  631. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/{soft_max.tmpl.wgsl → soft_max.wgsl} +106 -206
  632. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/solve_tri.wgsl +121 -0
  633. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/ssm_conv.wgsl +65 -0
  634. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/ssm_scan.wgsl +193 -0
  635. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/unary.wgsl +68 -48
  636. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/upscale.wgsl +240 -0
  637. data/ext/sources/ggml/src/ggml-zdnn/ggml-zdnn.cpp +18 -14
  638. data/ext/sources/ggml/src/ggml-zendnn/CMakeLists.txt +1 -1
  639. data/ext/sources/ggml/src/ggml-zendnn/ggml-zendnn.cpp +244 -10
  640. data/ext/sources/ggml/src/ggml.c +146 -42
  641. data/ext/sources/ggml/src/gguf.cpp +173 -28
  642. data/ext/sources/include/parakeet.h +342 -0
  643. data/ext/sources/include/whisper.h +31 -0
  644. data/ext/sources/media/matmul.png +0 -0
  645. data/ext/sources/src/CMakeLists.txt +23 -0
  646. data/ext/sources/src/parakeet-arch.h +188 -0
  647. data/ext/sources/src/parakeet.cpp +3838 -0
  648. data/ext/sources/src/whisper.cpp +220 -26
  649. data/extsources.rb +26 -10
  650. data/lib/whisper/log_settable.rb +33 -0
  651. data/lib/whisper/model/uri.rb +13 -8
  652. data/lib/whisper/output.rb +74 -0
  653. data/sig/whisper.rbs +417 -62
  654. data/test/helper.rb +2 -0
  655. data/test/jfk_reader/jfk_reader.c +50 -7
  656. data/test/test_callback.rb +1 -0
  657. data/test/test_package.rb +6 -5
  658. data/test/test_parakeet.rb +28 -0
  659. data/test/test_parakeet_callback.rb +107 -0
  660. data/test/test_parakeet_context.rb +116 -0
  661. data/test/test_parakeet_context_params.rb +24 -0
  662. data/test/test_parakeet_model.rb +21 -0
  663. data/test/test_parakeet_params.rb +78 -0
  664. data/test/test_parakeet_segment.rb +42 -0
  665. data/test/test_parakeet_token.rb +73 -0
  666. data/test/test_params.rb +2 -0
  667. data/test/test_vad.rb +9 -0
  668. data/test/test_vad_context.rb +2 -2
  669. data/test/test_vad_segment.rb +1 -1
  670. data/test/test_whisper.rb +24 -6
  671. data/whispercpp.gemspec +2 -2
  672. metadata +263 -304
  673. data/ext/sources/bindings/javascript/CMakeLists.txt +0 -41
  674. data/ext/sources/bindings/javascript/emscripten.cpp +0 -93
  675. data/ext/sources/bindings/javascript/libwhisper.worker.js +0 -1
  676. data/ext/sources/bindings/javascript/package.json +0 -26
  677. data/ext/sources/bindings/javascript/whisper.js +0 -19
  678. data/ext/sources/examples/addon.node/CMakeLists.txt +0 -31
  679. data/ext/sources/examples/addon.node/__test__/whisper.spec.js +0 -133
  680. data/ext/sources/examples/addon.node/addon.cpp +0 -557
  681. data/ext/sources/examples/addon.node/index.js +0 -59
  682. data/ext/sources/examples/addon.node/package.json +0 -16
  683. data/ext/sources/examples/addon.node/vad-example.js +0 -132
  684. data/ext/sources/examples/bench.wasm/CMakeLists.txt +0 -49
  685. data/ext/sources/examples/bench.wasm/emscripten.cpp +0 -87
  686. data/ext/sources/examples/bench.wasm/index-tmpl.html +0 -285
  687. data/ext/sources/examples/coi-serviceworker.js +0 -146
  688. data/ext/sources/examples/command/CMakeLists.txt +0 -10
  689. data/ext/sources/examples/command/command.cpp +0 -802
  690. data/ext/sources/examples/command/commands.txt +0 -9
  691. data/ext/sources/examples/command.wasm/CMakeLists.txt +0 -50
  692. data/ext/sources/examples/command.wasm/emscripten.cpp +0 -327
  693. data/ext/sources/examples/command.wasm/index-tmpl.html +0 -415
  694. data/ext/sources/examples/generate-karaoke.sh +0 -57
  695. data/ext/sources/examples/helpers.js +0 -191
  696. data/ext/sources/examples/livestream.sh +0 -112
  697. data/ext/sources/examples/lsp/CMakeLists.txt +0 -10
  698. data/ext/sources/examples/lsp/lsp.cpp +0 -471
  699. data/ext/sources/examples/lsp/whisper.vim +0 -362
  700. data/ext/sources/examples/python/test_whisper_processor.py +0 -7
  701. data/ext/sources/examples/python/whisper_processor.py +0 -54
  702. data/ext/sources/examples/server/bench.js +0 -29
  703. data/ext/sources/examples/server.py +0 -120
  704. data/ext/sources/examples/stream/CMakeLists.txt +0 -10
  705. data/ext/sources/examples/stream/stream.cpp +0 -437
  706. data/ext/sources/examples/stream.wasm/CMakeLists.txt +0 -49
  707. data/ext/sources/examples/stream.wasm/emscripten.cpp +0 -216
  708. data/ext/sources/examples/stream.wasm/index-tmpl.html +0 -491
  709. data/ext/sources/examples/sycl/CMakeLists.txt +0 -9
  710. data/ext/sources/examples/sycl/build.sh +0 -22
  711. data/ext/sources/examples/sycl/ls-sycl-device.cpp +0 -11
  712. data/ext/sources/examples/sycl/run-whisper.sh +0 -17
  713. data/ext/sources/examples/talk-llama/CMakeLists.txt +0 -48
  714. data/ext/sources/examples/talk-llama/eleven-labs.py +0 -80
  715. data/ext/sources/examples/talk-llama/llama-adapter.cpp +0 -488
  716. data/ext/sources/examples/talk-llama/llama-adapter.h +0 -89
  717. data/ext/sources/examples/talk-llama/llama-arch.cpp +0 -2877
  718. data/ext/sources/examples/talk-llama/llama-arch.h +0 -628
  719. data/ext/sources/examples/talk-llama/llama-batch.cpp +0 -919
  720. data/ext/sources/examples/talk-llama/llama-batch.h +0 -173
  721. data/ext/sources/examples/talk-llama/llama-chat.cpp +0 -896
  722. data/ext/sources/examples/talk-llama/llama-chat.h +0 -71
  723. data/ext/sources/examples/talk-llama/llama-context.cpp +0 -3633
  724. data/ext/sources/examples/talk-llama/llama-context.h +0 -359
  725. data/ext/sources/examples/talk-llama/llama-cparams.cpp +0 -5
  726. data/ext/sources/examples/talk-llama/llama-cparams.h +0 -47
  727. data/ext/sources/examples/talk-llama/llama-ext.h +0 -12
  728. data/ext/sources/examples/talk-llama/llama-grammar.cpp +0 -1464
  729. data/ext/sources/examples/talk-llama/llama-grammar.h +0 -194
  730. data/ext/sources/examples/talk-llama/llama-graph.cpp +0 -2735
  731. data/ext/sources/examples/talk-llama/llama-graph.h +0 -1031
  732. data/ext/sources/examples/talk-llama/llama-hparams.cpp +0 -258
  733. data/ext/sources/examples/talk-llama/llama-hparams.h +0 -353
  734. data/ext/sources/examples/talk-llama/llama-impl.cpp +0 -171
  735. data/ext/sources/examples/talk-llama/llama-impl.h +0 -75
  736. data/ext/sources/examples/talk-llama/llama-io.cpp +0 -15
  737. data/ext/sources/examples/talk-llama/llama-io.h +0 -35
  738. data/ext/sources/examples/talk-llama/llama-kv-cache-iswa.cpp +0 -330
  739. data/ext/sources/examples/talk-llama/llama-kv-cache-iswa.h +0 -137
  740. data/ext/sources/examples/talk-llama/llama-kv-cache.cpp +0 -2285
  741. data/ext/sources/examples/talk-llama/llama-kv-cache.h +0 -389
  742. data/ext/sources/examples/talk-llama/llama-kv-cells.h +0 -533
  743. data/ext/sources/examples/talk-llama/llama-memory-hybrid-iswa.cpp +0 -275
  744. data/ext/sources/examples/talk-llama/llama-memory-hybrid-iswa.h +0 -140
  745. data/ext/sources/examples/talk-llama/llama-memory-hybrid.cpp +0 -268
  746. data/ext/sources/examples/talk-llama/llama-memory-hybrid.h +0 -139
  747. data/ext/sources/examples/talk-llama/llama-memory-recurrent.cpp +0 -1165
  748. data/ext/sources/examples/talk-llama/llama-memory-recurrent.h +0 -182
  749. data/ext/sources/examples/talk-llama/llama-memory.cpp +0 -59
  750. data/ext/sources/examples/talk-llama/llama-memory.h +0 -122
  751. data/ext/sources/examples/talk-llama/llama-mmap.cpp +0 -752
  752. data/ext/sources/examples/talk-llama/llama-mmap.h +0 -73
  753. data/ext/sources/examples/talk-llama/llama-model-loader.cpp +0 -1655
  754. data/ext/sources/examples/talk-llama/llama-model-loader.h +0 -206
  755. data/ext/sources/examples/talk-llama/llama-model-saver.cpp +0 -299
  756. data/ext/sources/examples/talk-llama/llama-model-saver.h +0 -40
  757. data/ext/sources/examples/talk-llama/llama-model.cpp +0 -9056
  758. data/ext/sources/examples/talk-llama/llama-model.h +0 -597
  759. data/ext/sources/examples/talk-llama/llama-quant.cpp +0 -1304
  760. data/ext/sources/examples/talk-llama/llama-quant.h +0 -1
  761. data/ext/sources/examples/talk-llama/llama-sampler.cpp +0 -3885
  762. data/ext/sources/examples/talk-llama/llama-sampler.h +0 -42
  763. data/ext/sources/examples/talk-llama/llama-vocab.cpp +0 -3970
  764. data/ext/sources/examples/talk-llama/llama-vocab.h +0 -187
  765. data/ext/sources/examples/talk-llama/llama.cpp +0 -1194
  766. data/ext/sources/examples/talk-llama/llama.h +0 -1573
  767. data/ext/sources/examples/talk-llama/models/afmoe.cpp +0 -190
  768. data/ext/sources/examples/talk-llama/models/apertus.cpp +0 -125
  769. data/ext/sources/examples/talk-llama/models/arcee.cpp +0 -135
  770. data/ext/sources/examples/talk-llama/models/arctic.cpp +0 -137
  771. data/ext/sources/examples/talk-llama/models/arwkv7.cpp +0 -86
  772. data/ext/sources/examples/talk-llama/models/baichuan.cpp +0 -123
  773. data/ext/sources/examples/talk-llama/models/bailingmoe.cpp +0 -143
  774. data/ext/sources/examples/talk-llama/models/bailingmoe2.cpp +0 -133
  775. data/ext/sources/examples/talk-llama/models/bert.cpp +0 -184
  776. data/ext/sources/examples/talk-llama/models/bitnet.cpp +0 -145
  777. data/ext/sources/examples/talk-llama/models/bloom.cpp +0 -101
  778. data/ext/sources/examples/talk-llama/models/chameleon.cpp +0 -178
  779. data/ext/sources/examples/talk-llama/models/chatglm.cpp +0 -132
  780. data/ext/sources/examples/talk-llama/models/codeshell.cpp +0 -111
  781. data/ext/sources/examples/talk-llama/models/cogvlm.cpp +0 -102
  782. data/ext/sources/examples/talk-llama/models/cohere2-iswa.cpp +0 -134
  783. data/ext/sources/examples/talk-llama/models/command-r.cpp +0 -122
  784. data/ext/sources/examples/talk-llama/models/dbrx.cpp +0 -122
  785. data/ext/sources/examples/talk-llama/models/deci.cpp +0 -135
  786. data/ext/sources/examples/talk-llama/models/deepseek.cpp +0 -142
  787. data/ext/sources/examples/talk-llama/models/deepseek2.cpp +0 -262
  788. data/ext/sources/examples/talk-llama/models/delta-net-base.cpp +0 -445
  789. data/ext/sources/examples/talk-llama/models/dots1.cpp +0 -132
  790. data/ext/sources/examples/talk-llama/models/dream.cpp +0 -105
  791. data/ext/sources/examples/talk-llama/models/ernie4-5-moe.cpp +0 -148
  792. data/ext/sources/examples/talk-llama/models/ernie4-5.cpp +0 -110
  793. data/ext/sources/examples/talk-llama/models/eurobert.cpp +0 -97
  794. data/ext/sources/examples/talk-llama/models/exaone-moe.cpp +0 -145
  795. data/ext/sources/examples/talk-llama/models/exaone.cpp +0 -114
  796. data/ext/sources/examples/talk-llama/models/exaone4.cpp +0 -123
  797. data/ext/sources/examples/talk-llama/models/falcon-h1.cpp +0 -111
  798. data/ext/sources/examples/talk-llama/models/falcon.cpp +0 -120
  799. data/ext/sources/examples/talk-llama/models/gemma-embedding.cpp +0 -116
  800. data/ext/sources/examples/talk-llama/models/gemma.cpp +0 -112
  801. data/ext/sources/examples/talk-llama/models/gemma2-iswa.cpp +0 -128
  802. data/ext/sources/examples/talk-llama/models/gemma3.cpp +0 -155
  803. data/ext/sources/examples/talk-llama/models/gemma3n-iswa.cpp +0 -384
  804. data/ext/sources/examples/talk-llama/models/glm4-moe.cpp +0 -170
  805. data/ext/sources/examples/talk-llama/models/glm4.cpp +0 -157
  806. data/ext/sources/examples/talk-llama/models/gpt2.cpp +0 -105
  807. data/ext/sources/examples/talk-llama/models/gptneox.cpp +0 -144
  808. data/ext/sources/examples/talk-llama/models/granite-hybrid.cpp +0 -195
  809. data/ext/sources/examples/talk-llama/models/granite.cpp +0 -210
  810. data/ext/sources/examples/talk-llama/models/grok.cpp +0 -159
  811. data/ext/sources/examples/talk-llama/models/grovemoe.cpp +0 -139
  812. data/ext/sources/examples/talk-llama/models/hunyuan-dense.cpp +0 -132
  813. data/ext/sources/examples/talk-llama/models/hunyuan-moe.cpp +0 -153
  814. data/ext/sources/examples/talk-llama/models/internlm2.cpp +0 -120
  815. data/ext/sources/examples/talk-llama/models/jais.cpp +0 -86
  816. data/ext/sources/examples/talk-llama/models/jais2.cpp +0 -123
  817. data/ext/sources/examples/talk-llama/models/jamba.cpp +0 -106
  818. data/ext/sources/examples/talk-llama/models/kimi-linear.cpp +0 -381
  819. data/ext/sources/examples/talk-llama/models/lfm2.cpp +0 -196
  820. data/ext/sources/examples/talk-llama/models/llada-moe.cpp +0 -122
  821. data/ext/sources/examples/talk-llama/models/llada.cpp +0 -99
  822. data/ext/sources/examples/talk-llama/models/llama-iswa.cpp +0 -178
  823. data/ext/sources/examples/talk-llama/models/llama.cpp +0 -175
  824. data/ext/sources/examples/talk-llama/models/maincoder.cpp +0 -117
  825. data/ext/sources/examples/talk-llama/models/mamba-base.cpp +0 -289
  826. data/ext/sources/examples/talk-llama/models/mamba.cpp +0 -54
  827. data/ext/sources/examples/talk-llama/models/mimo2-iswa.cpp +0 -129
  828. data/ext/sources/examples/talk-llama/models/minicpm3.cpp +0 -200
  829. data/ext/sources/examples/talk-llama/models/minimax-m2.cpp +0 -123
  830. data/ext/sources/examples/talk-llama/models/mistral3.cpp +0 -160
  831. data/ext/sources/examples/talk-llama/models/models.h +0 -704
  832. data/ext/sources/examples/talk-llama/models/modern-bert.cpp +0 -109
  833. data/ext/sources/examples/talk-llama/models/mpt.cpp +0 -126
  834. data/ext/sources/examples/talk-llama/models/nemotron-h.cpp +0 -162
  835. data/ext/sources/examples/talk-llama/models/nemotron.cpp +0 -122
  836. data/ext/sources/examples/talk-llama/models/neo-bert.cpp +0 -104
  837. data/ext/sources/examples/talk-llama/models/olmo.cpp +0 -121
  838. data/ext/sources/examples/talk-llama/models/olmo2.cpp +0 -150
  839. data/ext/sources/examples/talk-llama/models/olmoe.cpp +0 -124
  840. data/ext/sources/examples/talk-llama/models/openai-moe-iswa.cpp +0 -127
  841. data/ext/sources/examples/talk-llama/models/openelm.cpp +0 -124
  842. data/ext/sources/examples/talk-llama/models/orion.cpp +0 -123
  843. data/ext/sources/examples/talk-llama/models/paddleocr.cpp +0 -122
  844. data/ext/sources/examples/talk-llama/models/pangu-embedded.cpp +0 -121
  845. data/ext/sources/examples/talk-llama/models/phi2.cpp +0 -121
  846. data/ext/sources/examples/talk-llama/models/phi3.cpp +0 -152
  847. data/ext/sources/examples/talk-llama/models/plamo.cpp +0 -110
  848. data/ext/sources/examples/talk-llama/models/plamo2.cpp +0 -320
  849. data/ext/sources/examples/talk-llama/models/plamo3.cpp +0 -128
  850. data/ext/sources/examples/talk-llama/models/plm.cpp +0 -169
  851. data/ext/sources/examples/talk-llama/models/qwen.cpp +0 -108
  852. data/ext/sources/examples/talk-llama/models/qwen2.cpp +0 -126
  853. data/ext/sources/examples/talk-llama/models/qwen2moe.cpp +0 -151
  854. data/ext/sources/examples/talk-llama/models/qwen2vl.cpp +0 -117
  855. data/ext/sources/examples/talk-llama/models/qwen3.cpp +0 -120
  856. data/ext/sources/examples/talk-llama/models/qwen35.cpp +0 -381
  857. data/ext/sources/examples/talk-llama/models/qwen35moe.cpp +0 -422
  858. data/ext/sources/examples/talk-llama/models/qwen3moe.cpp +0 -131
  859. data/ext/sources/examples/talk-llama/models/qwen3next.cpp +0 -525
  860. data/ext/sources/examples/talk-llama/models/qwen3vl-moe.cpp +0 -140
  861. data/ext/sources/examples/talk-llama/models/qwen3vl.cpp +0 -132
  862. data/ext/sources/examples/talk-llama/models/refact.cpp +0 -94
  863. data/ext/sources/examples/talk-llama/models/rnd1.cpp +0 -126
  864. data/ext/sources/examples/talk-llama/models/rwkv6-base.cpp +0 -164
  865. data/ext/sources/examples/talk-llama/models/rwkv6.cpp +0 -94
  866. data/ext/sources/examples/talk-llama/models/rwkv6qwen2.cpp +0 -86
  867. data/ext/sources/examples/talk-llama/models/rwkv7-base.cpp +0 -137
  868. data/ext/sources/examples/talk-llama/models/rwkv7.cpp +0 -90
  869. data/ext/sources/examples/talk-llama/models/seed-oss.cpp +0 -124
  870. data/ext/sources/examples/talk-llama/models/smallthinker.cpp +0 -126
  871. data/ext/sources/examples/talk-llama/models/smollm3.cpp +0 -128
  872. data/ext/sources/examples/talk-llama/models/stablelm.cpp +0 -146
  873. data/ext/sources/examples/talk-llama/models/starcoder.cpp +0 -100
  874. data/ext/sources/examples/talk-llama/models/starcoder2.cpp +0 -121
  875. data/ext/sources/examples/talk-llama/models/step35-iswa.cpp +0 -165
  876. data/ext/sources/examples/talk-llama/models/t5-dec.cpp +0 -166
  877. data/ext/sources/examples/talk-llama/models/t5-enc.cpp +0 -96
  878. data/ext/sources/examples/talk-llama/models/wavtokenizer-dec.cpp +0 -149
  879. data/ext/sources/examples/talk-llama/models/xverse.cpp +0 -108
  880. data/ext/sources/examples/talk-llama/prompts/talk-alpaca.txt +0 -23
  881. data/ext/sources/examples/talk-llama/speak +0 -40
  882. data/ext/sources/examples/talk-llama/speak.bat +0 -1
  883. data/ext/sources/examples/talk-llama/speak.ps1 +0 -14
  884. data/ext/sources/examples/talk-llama/talk-llama.cpp +0 -813
  885. data/ext/sources/examples/talk-llama/unicode-data.cpp +0 -7034
  886. data/ext/sources/examples/talk-llama/unicode-data.h +0 -20
  887. data/ext/sources/examples/talk-llama/unicode.cpp +0 -1103
  888. data/ext/sources/examples/talk-llama/unicode.h +0 -111
  889. data/ext/sources/examples/wchess/CMakeLists.txt +0 -10
  890. data/ext/sources/examples/wchess/libwchess/CMakeLists.txt +0 -19
  891. data/ext/sources/examples/wchess/libwchess/Chessboard.cpp +0 -803
  892. data/ext/sources/examples/wchess/libwchess/Chessboard.h +0 -33
  893. data/ext/sources/examples/wchess/libwchess/WChess.cpp +0 -193
  894. data/ext/sources/examples/wchess/libwchess/WChess.h +0 -63
  895. data/ext/sources/examples/wchess/libwchess/test-chessboard.cpp +0 -117
  896. data/ext/sources/examples/wchess/wchess.cmd/CMakeLists.txt +0 -8
  897. data/ext/sources/examples/wchess/wchess.cmd/wchess.cmd.cpp +0 -253
  898. data/ext/sources/examples/whisper.wasm/CMakeLists.txt +0 -50
  899. data/ext/sources/examples/whisper.wasm/emscripten.cpp +0 -118
  900. data/ext/sources/examples/whisper.wasm/index-tmpl.html +0 -659
  901. data/ext/sources/ggml/src/ggml-cuda/template-instances/generate_cu_files.py +0 -99
  902. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-msg.h +0 -155
  903. data/ext/sources/ggml/src/ggml-hexagon/op-desc.h +0 -153
  904. data/ext/sources/ggml/src/ggml-opencl/kernels/embed_kernel.py +0 -26
  905. data/ext/sources/ggml/src/ggml-openvino/openvino/pass/eliminate_zp.cpp +0 -123
  906. data/ext/sources/ggml/src/ggml-openvino/openvino/pass/eliminate_zp.h +0 -17
  907. data/ext/sources/ggml/src/ggml-virtgpu/regenerate_remoting.py +0 -333
  908. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/abs.comp +0 -21
  909. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/ceil.comp +0 -22
  910. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/clamp.comp +0 -17
  911. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/cos.comp +0 -17
  912. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/elu.comp +0 -27
  913. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/exp.comp +0 -21
  914. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/floor.comp +0 -22
  915. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu.comp +0 -25
  916. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu_erf.comp +0 -39
  917. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu_quick.comp +0 -23
  918. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/hardsigmoid.comp +0 -22
  919. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/hardswish.comp +0 -22
  920. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/leaky_relu.comp +0 -22
  921. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/neg.comp +0 -20
  922. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/relu.comp +0 -21
  923. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/round.comp +0 -29
  924. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/rte.glsl +0 -5
  925. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sgn.comp +0 -21
  926. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sigmoid.comp +0 -20
  927. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/silu.comp +0 -22
  928. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sin.comp +0 -17
  929. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/softplus.comp +0 -23
  930. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sqrt.comp +0 -17
  931. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/square.comp +0 -17
  932. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/step.comp +0 -22
  933. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/tanh.comp +0 -20
  934. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/trunc.comp +0 -22
  935. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/xielu.comp +0 -35
  936. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/embed_wgsl.py +0 -182
  937. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/glu.tmpl.wgsl +0 -323
  938. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat.wgsl +0 -718
  939. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/rms_norm.wgsl +0 -123
  940. data/ext/sources/tests/CMakeLists.txt +0 -112
  941. data/ext/sources/tests/earnings21/eval.mk +0 -58
  942. data/ext/sources/tests/earnings21/eval.py +0 -68
  943. data/ext/sources/tests/earnings21/normalizers/__init__.py +0 -2
  944. data/ext/sources/tests/earnings21/normalizers/basic.py +0 -80
  945. data/ext/sources/tests/earnings21/normalizers/english.json +0 -1741
  946. data/ext/sources/tests/earnings21/normalizers/english.py +0 -550
  947. data/ext/sources/tests/earnings21/requirements.txt +0 -6
  948. data/ext/sources/tests/en-0-ref.txt +0 -1
  949. data/ext/sources/tests/en-1-ref.txt +0 -1
  950. data/ext/sources/tests/en-2-ref.txt +0 -1
  951. data/ext/sources/tests/es-0-ref.txt +0 -1
  952. data/ext/sources/tests/librispeech/eval.mk +0 -39
  953. data/ext/sources/tests/librispeech/eval.py +0 -47
  954. data/ext/sources/tests/librispeech/normalizers/__init__.py +0 -2
  955. data/ext/sources/tests/librispeech/normalizers/basic.py +0 -80
  956. data/ext/sources/tests/librispeech/normalizers/english.json +0 -1741
  957. data/ext/sources/tests/librispeech/normalizers/english.py +0 -550
  958. data/ext/sources/tests/librispeech/requirements.txt +0 -6
  959. data/ext/sources/tests/run-tests.sh +0 -130
  960. data/ext/sources/tests/test-c.c +0 -3
  961. data/ext/sources/tests/test-vad-full.cpp +0 -56
  962. data/ext/sources/tests/test-vad.cpp +0 -83
  963. data/ext/sources/tests/test-whisper.js +0 -58
  964. data/lib/whisper/context.rb +0 -15
  965. data/lib/whisper/segment.rb +0 -58
@@ -1,5 +1,37 @@
1
1
  #pragma OPENCL EXTENSION cl_khr_fp16 : enable
2
2
 
3
+ #ifdef cl_intel_subgroups
4
+ #pragma OPENCL EXTENSION cl_intel_subgroups : enable
5
+ #else
6
+ #pragma OPENCL EXTENSION cl_khr_subgroups : enable
7
+ #endif
8
+
9
+ #ifdef cl_qcom_reqd_sub_group_size
10
+ #pragma OPENCL EXTENSION cl_qcom_reqd_sub_group_size : enable
11
+ #define REQD_SUBGROUP_SIZE_64 __attribute__((qcom_reqd_sub_group_size("half")))
12
+ #else
13
+ #define REQD_SUBGROUP_SIZE_64
14
+ #endif
15
+
16
+ // subgroup size for q1 kernels
17
+ #ifndef FA_SG
18
+ #define FA_SG 64
19
+ #endif
20
+ #ifdef cl_intel_required_subgroup_size
21
+ #pragma OPENCL EXTENSION cl_intel_required_subgroup_size : enable
22
+ #define REQD_FA_SG __attribute__((intel_reqd_sub_group_size(FA_SG)))
23
+ #else
24
+ #define REQD_FA_SG
25
+ #endif
26
+
27
+ #ifdef cl_khr_subgroup_shuffle
28
+ #pragma OPENCL EXTENSION cl_khr_subgroup_shuffle : enable
29
+ #define HAS_SUBGROUP_SHUFFLE 1
30
+ #elif defined(cl_qcom_subgroup_shuffle)
31
+ #pragma OPENCL EXTENSION cl_qcom_subgroup_shuffle : enable
32
+ #define HAS_SUBGROUP_SHUFFLE 1
33
+ #endif
34
+
3
35
  #define ACC_TYPE float
4
36
  #define ACC_TYPE4 float4
5
37
  #define Q_DATA_TYPE4 float4
@@ -12,8 +44,37 @@
12
44
 
13
45
  #define DK_VEC (DK/4)
14
46
  #define DV_VEC (DV/4)
47
+
48
+ #ifndef FA_PARTIAL_FLOATS
49
+ #define FA_PARTIAL_FLOATS (2 + DV)
50
+ #endif
51
+ #define Q1_WG_SIZE FA_SG
52
+
53
+ // The kernels are built with -cl-finite-math-only. On some older Adreno GPUs,
54
+ // infinite operand can cause undefined behavior and miscompilation for exp.
55
+ // Therefore, a large negative value is used instead.
56
+ #define FA_M_INIT (-3.0e38f)
57
+
58
+ // Drop full unroll at DK>=192 — Adreno compiler host-memory budget.
59
+ #if DK >= 192
60
+ #define FA_UNROLL
61
+ #else
62
+ #define FA_UNROLL _Pragma("unroll")
63
+ #endif
64
+
65
+ // N_SPLIT>1 splits DK/DV across threads to cut per-thread register use.
66
+ #ifndef N_SPLIT
67
+ #define N_SPLIT 1
68
+ #endif
69
+
70
+ #define SPLIT_DK_VEC (DK_VEC / N_SPLIT)
71
+ #define SPLIT_DV_VEC (DV_VEC / N_SPLIT)
72
+
73
+ #if N_SPLIT > 1
74
+ #define WG_SIZE (BLOCK_M * N_SPLIT)
75
+ #else
15
76
  #define WG_SIZE (BLOCK_M)
16
- #define Q1_WG_SIZE 64
77
+ #endif
17
78
 
18
79
  inline float get_alibi_slope(
19
80
  const float max_bias, const uint h, const uint n_head_log2, const float m0, const float m1
@@ -26,9 +87,20 @@ inline float get_alibi_slope(
26
87
 
27
88
  return pow(base, exph);
28
89
  }
29
- __kernel void flash_attn_f32_f16(
90
+
91
+ // Adreno compiler crashes when attempting to compile the entire program for DK=512,
92
+ // FA_DECODE_ONLY allows bypass the encoding kernel.
93
+ #if !defined(FA_DECODE_ONLY) && !defined(FA_MQ_ONLY)
94
+ #ifndef FA_TILE_NAME
95
+ #define FA_TILE_NAME flash_attn_f32_f16
96
+ #endif
97
+ __kernel void FA_TILE_NAME(
30
98
  const global void * q_void, ulong q_offset,
99
+ #ifdef FA_K_IMG
100
+ __read_only image1d_buffer_t k_img, ulong k_offset_unused,
101
+ #else
31
102
  const global void * k_void, ulong k_offset,
103
+ #endif
32
104
  const global void * v_void, ulong v_offset,
33
105
  global void * o_void, ulong o_offset,
34
106
  const float scale,
@@ -54,48 +126,83 @@ __kernel void flash_attn_f32_f16(
54
126
  const int mask_ne2,
55
127
  const int mask_ne3,
56
128
  const global void* sinks_void,
57
- const ulong sinks_offset
129
+ const ulong sinks_offset,
130
+ const global void * k_pad_void,
131
+ const global void * v_pad_void,
132
+ const global void * mask_pad_void,
133
+ const global char * blk,
134
+ const int n_kv_blocks,
135
+ const ulong mask_pad_nb1,
136
+ const ulong mask_pad_nb2,
137
+ const ulong mask_pad_nb3
58
138
  ) {
59
139
  const int tid = get_local_id(0);
60
140
  const int block_q_idx = get_group_id(0);
61
141
  const int head_batch_idx = get_global_id(1);
62
142
 
63
- const int my_query_row = block_q_idx * BLOCK_M + tid;
143
+ #if N_SPLIT > 1
144
+ const int q_lane = tid / N_SPLIT;
145
+ const int split_idx = tid % N_SPLIT;
146
+ #else
147
+ const int q_lane = tid;
148
+ const int split_idx = 0;
149
+ #endif
150
+
151
+ const int my_query_row = block_q_idx * BLOCK_M + q_lane;
152
+ const int query_valid = my_query_row < n_q;
64
153
 
65
154
  const int batch_idx = head_batch_idx / n_head;
66
155
  const int head_idx = head_batch_idx % n_head;
67
156
 
68
157
  const int gqa_ratio = n_head / n_head_kv;
69
158
  const int head_kv_idx = head_idx / gqa_ratio;
159
+ const int mask_head_idx = mask_void != NULL ? head_idx % mask_ne2 : 0;
160
+ const int mask_batch_idx = mask_void != NULL ? batch_idx % mask_ne3 : 0;
70
161
 
71
162
  const global char* q_base = (const global char*)q_void + q_offset;
163
+ #ifndef FA_K_IMG
72
164
  const global char* k_base = (const global char*)k_void + k_offset;
165
+ #endif
73
166
  const global char* v_base = (const global char*)v_void + v_offset;
74
167
  global char* o_base = (global char*)o_void + o_offset;
75
168
 
76
169
  const global char* mask_base = NULL;
77
170
  if (mask_void != NULL) {
78
- const int mask_head_idx = head_idx % mask_ne2;
79
- const int mask_batch_idx = batch_idx % mask_ne3;
80
171
  mask_base = (const global char*)mask_void + mask_offset + mask_batch_idx * mask_nb3 + mask_head_idx * mask_nb2;
81
172
  }
173
+ const global char* mask_pad_base = NULL;
174
+ if (mask_pad_void != NULL) {
175
+ mask_pad_base = (const global char*)mask_pad_void + mask_batch_idx * mask_pad_nb3 + mask_head_idx * mask_pad_nb2;
176
+ }
177
+ const global char* blk_base = NULL;
178
+ if (blk != NULL) {
179
+ const int n_q_blocks = (n_q + BLOCK_M - 1) / BLOCK_M;
180
+ blk_base = blk + (((mask_batch_idx * mask_ne2) + mask_head_idx) * n_q_blocks + block_q_idx) * n_kv_blocks;
181
+ }
82
182
 
83
- ACC_TYPE4 q_priv[DK_VEC];
84
- if (my_query_row < n_q) {
183
+ ACC_TYPE4 q_priv[SPLIT_DK_VEC];
184
+ const int dk_off = split_idx * SPLIT_DK_VEC;
185
+ if (query_valid) {
85
186
  const ulong q_row_offset = batch_idx * q_nb3 + head_idx * q_nb2 + my_query_row * q_nb1;
86
187
  const global Q_DATA_TYPE4* q_ptr = (const global Q_DATA_TYPE4*)(q_base + q_row_offset);
87
- #pragma unroll
88
- for (int i = 0; i < DK_VEC; ++i) {
89
- q_priv[i] = CONVERT_Q_ACC4(q_ptr[i]);
188
+ FA_UNROLL
189
+ for (int i = 0; i < SPLIT_DK_VEC; ++i) {
190
+ q_priv[i] = CONVERT_Q_ACC4(q_ptr[dk_off + i]);
191
+ }
192
+ } else {
193
+ FA_UNROLL
194
+ for (int i = 0; i < SPLIT_DK_VEC; ++i) {
195
+ q_priv[i] = (ACC_TYPE4)(0.0f);
90
196
  }
91
197
  }
92
198
 
93
- ACC_TYPE4 o_acc[DV_VEC];
94
- #pragma unroll
95
- for (int i = 0; i < DV_VEC; ++i) {
199
+ ACC_TYPE4 o_acc[SPLIT_DV_VEC];
200
+ FA_UNROLL
201
+ for (int i = 0; i < SPLIT_DV_VEC; ++i) {
96
202
  o_acc[i] = (ACC_TYPE4)(0.0f);
97
203
  }
98
- ACC_TYPE m_i = -INFINITY;
204
+
205
+ ACC_TYPE m_i = FA_M_INIT;
99
206
  ACC_TYPE l_i = 0.0f;
100
207
 
101
208
  float slope = get_alibi_slope(max_bias, head_idx, n_head_log2, m0, m1);
@@ -103,86 +210,388 @@ __kernel void flash_attn_f32_f16(
103
210
  __local KV_DATA_TYPE4 l_k[BLOCK_N][DK_VEC];
104
211
  __local KV_DATA_TYPE4 l_v[BLOCK_N][DV_VEC];
105
212
 
213
+ #if N_SPLIT > 1 && !defined(HAS_SUBGROUP_SHUFFLE)
214
+ __local ACC_TYPE local_partial[BLOCK_N][WG_SIZE];
215
+ __local ACC_TYPE local_p[BLOCK_M][BLOCK_N];
216
+ __local ACC_TYPE local_softmax_scale[BLOCK_M];
217
+ __local ACC_TYPE local_l_inv[BLOCK_M];
218
+ #endif
219
+
106
220
  for (int k_start = 0; k_start < n_kv; k_start += BLOCK_N) {
221
+ char blk_cur = 1;
222
+ if (blk_base != NULL) {
223
+ blk_cur = blk_base[k_start / BLOCK_N];
224
+ if (blk_cur == 0) continue;
225
+ }
226
+
227
+ const int use_kv_pad = k_pad_void != NULL && k_start + BLOCK_N > n_kv;
228
+ const int k_tile_start = use_kv_pad ? 0 : k_start;
229
+ const ulong k_tile_nb2 = use_kv_pad ? (ulong) BLOCK_N * k_nb1 : k_nb2;
230
+ const ulong k_tile_nb3 = use_kv_pad ? (ulong) n_head_kv * k_tile_nb2 : k_nb3;
231
+ const ulong v_tile_nb2 = use_kv_pad ? (ulong) BLOCK_N * v_nb1 : v_nb2;
232
+ const ulong v_tile_nb3 = use_kv_pad ? (ulong) n_head_kv * v_tile_nb2 : v_nb3;
233
+ #ifdef FA_K_IMG
234
+ // K via texture cache for the bulk (aligned) tiles; the ragged last
235
+ // tile (use_kv_pad) still reads the f32-strided pad buffer from global.
236
+ const global char* k_tile_base = use_kv_pad ? (const global char*) k_pad_void : (const global char*) 0;
237
+ const int k_pitch_px_row = (int)(k_nb1 >> 3);
238
+ const int k_pitch_px_head = (int)(k_nb2 >> 3);
239
+ const int k_pitch_px_batch = (int)(k_nb3 >> 3);
240
+ #else
241
+ const global char* k_tile_base = use_kv_pad ? (const global char*) k_pad_void : k_base;
242
+ #endif
243
+ const global char* v_tile_base = use_kv_pad ? (const global char*) v_pad_void : v_base;
244
+
107
245
  for (int i = tid; i < BLOCK_N * DK_VEC; i += WG_SIZE) {
108
246
  const int row = i / DK_VEC;
109
247
  const int col = i % DK_VEC;
110
- const int k_row_idx = k_start + row;
111
- if (k_row_idx < n_kv) {
112
- const ulong k_row_offset = batch_idx * k_nb3 + head_kv_idx * k_nb2 + k_row_idx * k_nb1;
113
- l_k[row][col] = ((__global KV_DATA_TYPE4*)(k_base + k_row_offset))[col];
248
+ const int k_row_idx = k_tile_start + row;
249
+ if (use_kv_pad || k_row_idx < n_kv) {
250
+ #ifdef FA_K_IMG
251
+ if (use_kv_pad) {
252
+ const ulong k_row_offset = batch_idx * k_tile_nb3 + head_kv_idx * k_tile_nb2 + k_row_idx * k_nb1;
253
+ l_k[row][col] = ((__global KV_DATA_TYPE4*)(k_tile_base + k_row_offset))[col];
254
+ } else {
255
+ const int k_row_px = batch_idx * k_pitch_px_batch + head_kv_idx * k_pitch_px_head + k_row_idx * k_pitch_px_row;
256
+ l_k[row][col] = read_imageh(k_img, k_row_px + col);
257
+ }
258
+ #else
259
+ const ulong k_row_offset = batch_idx * k_tile_nb3 + head_kv_idx * k_tile_nb2 + k_row_idx * k_nb1;
260
+ l_k[row][col] = ((__global KV_DATA_TYPE4*)(k_tile_base + k_row_offset))[col];
261
+ #endif
262
+ } else {
263
+ l_k[row][col] = (KV_DATA_TYPE4)(0.0h);
114
264
  }
115
265
  }
116
266
  for (int i = tid; i < BLOCK_N * DV_VEC; i += WG_SIZE) {
117
267
  const int row = i / DV_VEC;
118
268
  const int col = i % DV_VEC;
119
- const int v_row_idx = k_start + row;
120
- if (v_row_idx < n_kv) {
121
- const ulong v_row_offset = batch_idx * v_nb3 + head_kv_idx * v_nb2 + v_row_idx * v_nb1;
122
- l_v[row][col] = ((__global KV_DATA_TYPE4*)(v_base + v_row_offset))[col];
269
+ const int v_row_idx = k_tile_start + row;
270
+ if (use_kv_pad || v_row_idx < n_kv) {
271
+ const ulong v_row_offset = batch_idx * v_tile_nb3 + head_kv_idx * v_tile_nb2 + v_row_idx * v_nb1;
272
+ l_v[row][col] = ((__global KV_DATA_TYPE4*)(v_tile_base + v_row_offset))[col];
273
+ } else {
274
+ l_v[row][col] = (KV_DATA_TYPE4)(0.0h);
123
275
  }
124
276
  }
125
277
  barrier(CLK_LOCAL_MEM_FENCE);
126
278
 
127
- if (my_query_row >= n_q) {
128
- continue;
129
- }
279
+ #if N_SPLIT > 1 && defined(HAS_SUBGROUP_SHUFFLE)
280
+ {
281
+ const int dv_off = split_idx * SPLIT_DV_VEC;
282
+ for (int j = 0; j < BLOCK_N; j += 2) {
283
+ const int k_row0 = k_start + j;
284
+ const int k_row1 = k_start + j + 1;
130
285
 
131
- for (int j = 0; j < BLOCK_N; j += 2) {
132
- const int k_row0 = k_start + j;
133
- const int k_row1 = k_start + j + 1;
286
+ ACC_TYPE partial0 = 0.0f;
287
+ ACC_TYPE partial1 = 0.0f;
288
+ FA_UNROLL
289
+ for (int k = 0; k < SPLIT_DK_VEC; k++) {
290
+ const ACC_TYPE4 qk = q_priv[k];
291
+ ACC_TYPE4 dot0 = qk * CONVERT_KV_ACC4(l_k[j ][dk_off + k]);
292
+ ACC_TYPE4 dot1 = qk * CONVERT_KV_ACC4(l_k[j+1][dk_off + k]);
293
+ partial0 += dot0.s0 + dot0.s1 + dot0.s2 + dot0.s3;
294
+ partial1 += dot1.s0 + dot1.s1 + dot1.s2 + dot1.s3;
295
+ }
134
296
 
135
- ACC_TYPE4 dot_acc0 = (ACC_TYPE4)(0.0f);
136
- ACC_TYPE4 dot_acc1 = (ACC_TYPE4)(0.0f);
137
- #pragma unroll
138
- for (int k = 0; k < DK_VEC; k++) {
139
- dot_acc0 = mad(q_priv[k], CONVERT_KV_ACC4(l_k[j][k]), dot_acc0);
140
- dot_acc1 = mad(q_priv[k], CONVERT_KV_ACC4(l_k[j+1][k]), dot_acc1);
141
- }
142
- ACC_TYPE score0 = (dot_acc0.s0 + dot_acc0.s1 + dot_acc0.s2 + dot_acc0.s3) * scale;
143
- ACC_TYPE score1 = (dot_acc1.s0 + dot_acc1.s1 + dot_acc1.s2 + dot_acc1.s3) * scale;
297
+ FA_UNROLL
298
+ for (int step = 1; step < N_SPLIT; step <<= 1) {
299
+ partial0 += sub_group_shuffle_xor(partial0, step);
300
+ partial1 += sub_group_shuffle_xor(partial1, step);
301
+ }
302
+
303
+ ACC_TYPE score0 = partial0 * scale;
304
+ ACC_TYPE score1 = partial1 * scale;
305
+
306
+ if (!query_valid) { score0 = FA_M_INIT; score1 = FA_M_INIT; }
307
+ if (is_causal) {
308
+ if (k_row0 > (n_kv - n_q + my_query_row)) score0 = FA_M_INIT;
309
+ if (k_row1 > (n_kv - n_q + my_query_row)) score1 = FA_M_INIT;
310
+ }
311
+ if (k_row0 >= n_kv) score0 = FA_M_INIT;
312
+ if (k_row1 >= n_kv) score1 = FA_M_INIT;
144
313
 
145
- if (is_causal) {
146
- if (k_row0 > (n_kv - n_q + my_query_row)) score0 = -INFINITY;
147
- if (k_row1 > (n_kv - n_q + my_query_row)) score1 = -INFINITY;
314
+ if (query_valid && mask_base != NULL && blk_cur != 2) {
315
+ if (use_kv_pad && mask_pad_base != NULL) {
316
+ const global MASK_DATA_TYPE* mask_ptr =
317
+ (const global MASK_DATA_TYPE*)(mask_pad_base + my_query_row * mask_pad_nb1);
318
+ score0 += slope * (ACC_TYPE)mask_ptr[j];
319
+ score1 += slope * (ACC_TYPE)mask_ptr[j + 1];
320
+ } else {
321
+ const global MASK_DATA_TYPE* mask_ptr =
322
+ (const global MASK_DATA_TYPE*)(mask_base + my_query_row * mask_nb1);
323
+ if (k_row0 < n_kv) score0 += slope * (ACC_TYPE)mask_ptr[k_row0];
324
+ if (k_row1 < n_kv) score1 += slope * (ACC_TYPE)mask_ptr[k_row1];
325
+ }
326
+ }
327
+
328
+ if (logit_softcap > 0.0f) {
329
+ score0 = logit_softcap * tanh(score0 / logit_softcap);
330
+ score1 = logit_softcap * tanh(score1 / logit_softcap);
331
+ }
332
+
333
+ const ACC_TYPE m_new = max(m_i, max(score0, score1));
334
+ // Whole tile masked (m_new == FA_M_INIT): force the exp() args
335
+ // far negative so the tile contributes 0, not exp(0)=1.
336
+ const ACC_TYPE m_exp = (m_new == FA_M_INIT) ? 0.0f : m_new;
337
+ const ACC_TYPE sp = native_exp(m_i - m_exp);
338
+ const ACC_TYPE p0 = native_exp(score0 - m_exp);
339
+ const ACC_TYPE p1 = native_exp(score1 - m_exp);
340
+
341
+ FA_UNROLL
342
+ for (int i = 0; i < SPLIT_DV_VEC; ++i) {
343
+ o_acc[i] = o_acc[i] * sp
344
+ + p0 * CONVERT_KV_ACC4(l_v[j ][dv_off + i])
345
+ + p1 * CONVERT_KV_ACC4(l_v[j+1][dv_off + i]);
346
+ }
347
+ l_i = l_i * sp + p0 + p1;
348
+ m_i = m_new;
349
+ }
350
+ }
351
+ #elif N_SPLIT > 1
352
+ // N_SPLIT>1 fallback (no shuffle): 3-phase local-memory reduction.
353
+ // Phase 1 — partial dots for all BLOCK_N tokens.
354
+ for (int j = 0; j < BLOCK_N; ++j) {
355
+ ACC_TYPE4 dot_acc = (ACC_TYPE4)(0.0f);
356
+ FA_UNROLL
357
+ for (int k = 0; k < SPLIT_DK_VEC; k++) {
358
+ dot_acc = mad(q_priv[k], CONVERT_KV_ACC4(l_k[j][dk_off + k]), dot_acc);
148
359
  }
360
+ local_partial[j][tid] =
361
+ dot_acc.s0 + dot_acc.s1 + dot_acc.s2 + dot_acc.s3;
362
+ }
363
+ barrier(CLK_LOCAL_MEM_FENCE); // 1 barrier: partial dots visible
364
+
365
+ // Phase 2 — split_idx==0 reduces partial sums and computes block softmax.
366
+ if (split_idx == 0) {
367
+ if (query_valid) {
368
+ ACC_TYPE m_new = m_i;
369
+ for (int j = 0; j < BLOCK_N; ++j) {
370
+ const int k_row = k_start + j;
371
+ ACC_TYPE score = 0.0f;
372
+ FA_UNROLL
373
+ for (int s = 0; s < N_SPLIT; s++) {
374
+ score += local_partial[j][q_lane * N_SPLIT + s];
375
+ }
376
+ score *= scale;
377
+
378
+ if (is_causal && k_row > (n_kv - n_q + my_query_row)) score = FA_M_INIT;
379
+ if (k_row >= n_kv) score = FA_M_INIT;
149
380
 
150
- if (k_row0 >= n_kv) score0 = -INFINITY;
151
- if (k_row1 >= n_kv) score1 = -INFINITY;
381
+ if (mask_base != NULL && blk_cur != 2) {
382
+ if (use_kv_pad && mask_pad_base != NULL) {
383
+ const global MASK_DATA_TYPE* mask_ptr =
384
+ (const global MASK_DATA_TYPE*)(mask_pad_base + my_query_row * mask_pad_nb1);
385
+ score += slope * (ACC_TYPE)mask_ptr[j];
386
+ } else {
387
+ const global MASK_DATA_TYPE* mask_ptr =
388
+ (const global MASK_DATA_TYPE*)(mask_base + my_query_row * mask_nb1);
389
+ if (k_row < n_kv) score += slope * (ACC_TYPE)mask_ptr[k_row];
390
+ }
391
+ }
152
392
 
153
- if (mask_base != NULL) {
154
- const global MASK_DATA_TYPE* mask_ptr = (const global MASK_DATA_TYPE*)(mask_base + my_query_row * mask_nb1);
155
- if (k_row0 < n_kv) score0 += slope * (ACC_TYPE)mask_ptr[k_row0];
156
- if (k_row1 < n_kv) score1 += slope * (ACC_TYPE)mask_ptr[k_row1];
393
+ if (logit_softcap > 0.0f) {
394
+ score = logit_softcap * tanh(score / logit_softcap);
395
+ }
396
+
397
+ m_new = max(m_new, score);
398
+ local_p[q_lane][j] = score;
399
+ }
400
+
401
+ const ACC_TYPE m_exp = (m_new == FA_M_INIT) ? 0.0f : m_new;
402
+ const ACC_TYPE sp = native_exp(m_i - m_exp);
403
+ ACC_TYPE l_new = l_i * sp;
404
+ for (int j = 0; j < BLOCK_N; ++j) {
405
+ const ACC_TYPE p = native_exp(local_p[q_lane][j] - m_exp);
406
+ local_p[q_lane][j] = p;
407
+ l_new += p;
408
+ }
409
+ local_softmax_scale[q_lane] = sp;
410
+ l_i = l_new;
411
+ m_i = m_new;
412
+ } else {
413
+ local_softmax_scale[q_lane] = 1.0f;
414
+ for (int j = 0; j < BLOCK_N; ++j) local_p[q_lane][j] = 0.0f;
157
415
  }
416
+ }
417
+ barrier(CLK_LOCAL_MEM_FENCE);
158
418
 
159
- if (logit_softcap > 0.0f) {
160
- score0 = logit_softcap * tanh(score0 / logit_softcap);
161
- score1 = logit_softcap * tanh(score1 / logit_softcap);
419
+ // Phase 3 — V accumulate using broadcast probabilities.
420
+ {
421
+ const ACC_TYPE sp_block = local_softmax_scale[q_lane];
422
+ const int dv_off = split_idx * SPLIT_DV_VEC;
423
+ FA_UNROLL
424
+ for (int i = 0; i < SPLIT_DV_VEC; ++i) {
425
+ o_acc[i] *= sp_block;
426
+ }
427
+ for (int j = 0; j < BLOCK_N; ++j) {
428
+ const ACC_TYPE p = local_p[q_lane][j];
429
+ FA_UNROLL
430
+ for (int i = 0; i < SPLIT_DV_VEC; ++i) {
431
+ o_acc[i] = mad(p, CONVERT_KV_ACC4(l_v[j][dv_off + i]), o_acc[i]);
432
+ }
162
433
  }
434
+ }
435
+ #else
436
+ // N_SPLIT==1: j+=4 unroll. Requires BLOCK_N % 4 == 0.
437
+ if (query_valid) {
438
+ for (int j = 0; j < BLOCK_N; j += 4) {
439
+ const int k_row0 = k_start + j;
440
+ const int k_row1 = k_start + j + 1;
441
+ const int k_row2 = k_start + j + 2;
442
+ const int k_row3 = k_start + j + 3;
163
443
 
164
- const ACC_TYPE m_new = max(m_i, max(score0, score1));
165
- const ACC_TYPE p0 = exp(score0 - m_new);
166
- const ACC_TYPE p1 = exp(score1 - m_new);
167
- const ACC_TYPE scale_prev = exp(m_i - m_new);
444
+ ACC_TYPE4 dot_acc0 = (ACC_TYPE4)(0.0f);
445
+ ACC_TYPE4 dot_acc1 = (ACC_TYPE4)(0.0f);
446
+ ACC_TYPE4 dot_acc2 = (ACC_TYPE4)(0.0f);
447
+ ACC_TYPE4 dot_acc3 = (ACC_TYPE4)(0.0f);
448
+ FA_UNROLL
449
+ for (int k = 0; k < DK_VEC; k++) {
450
+ const ACC_TYPE4 qk = q_priv[k];
451
+ dot_acc0 = mad(qk, CONVERT_KV_ACC4(l_k[j][k]), dot_acc0);
452
+ dot_acc1 = mad(qk, CONVERT_KV_ACC4(l_k[j+1][k]), dot_acc1);
453
+ dot_acc2 = mad(qk, CONVERT_KV_ACC4(l_k[j+2][k]), dot_acc2);
454
+ dot_acc3 = mad(qk, CONVERT_KV_ACC4(l_k[j+3][k]), dot_acc3);
455
+ }
456
+ ACC_TYPE s0 = (dot_acc0.s0 + dot_acc0.s1 + dot_acc0.s2 + dot_acc0.s3) * scale;
457
+ ACC_TYPE s1 = (dot_acc1.s0 + dot_acc1.s1 + dot_acc1.s2 + dot_acc1.s3) * scale;
458
+ ACC_TYPE s2 = (dot_acc2.s0 + dot_acc2.s1 + dot_acc2.s2 + dot_acc2.s3) * scale;
459
+ ACC_TYPE s3 = (dot_acc3.s0 + dot_acc3.s1 + dot_acc3.s2 + dot_acc3.s3) * scale;
168
460
 
169
- #pragma unroll
170
- for (int i = 0; i < DV_VEC; ++i) {
171
- o_acc[i] = o_acc[i] * scale_prev + p0 * CONVERT_KV_ACC4(l_v[j][i]) + p1 * CONVERT_KV_ACC4(l_v[j+1][i]);
461
+ if (is_causal) {
462
+ const int causal_limit = n_kv - n_q + my_query_row;
463
+ if (k_row0 > causal_limit) s0 = FA_M_INIT;
464
+ if (k_row1 > causal_limit) s1 = FA_M_INIT;
465
+ if (k_row2 > causal_limit) s2 = FA_M_INIT;
466
+ if (k_row3 > causal_limit) s3 = FA_M_INIT;
467
+ }
468
+ if (k_row0 >= n_kv) s0 = FA_M_INIT;
469
+ if (k_row1 >= n_kv) s1 = FA_M_INIT;
470
+ if (k_row2 >= n_kv) s2 = FA_M_INIT;
471
+ if (k_row3 >= n_kv) s3 = FA_M_INIT;
472
+
473
+ if (mask_base != NULL && blk_cur != 2) {
474
+ if (use_kv_pad && mask_pad_base != NULL) {
475
+ const global MASK_DATA_TYPE* mask_ptr = (const global MASK_DATA_TYPE*)(mask_pad_base + my_query_row * mask_pad_nb1);
476
+ s0 += slope * (ACC_TYPE)mask_ptr[j];
477
+ s1 += slope * (ACC_TYPE)mask_ptr[j + 1];
478
+ s2 += slope * (ACC_TYPE)mask_ptr[j + 2];
479
+ s3 += slope * (ACC_TYPE)mask_ptr[j + 3];
480
+ } else {
481
+ const global MASK_DATA_TYPE* mask_ptr = (const global MASK_DATA_TYPE*)(mask_base + my_query_row * mask_nb1);
482
+ if (k_row0 < n_kv) s0 += slope * (ACC_TYPE)mask_ptr[k_row0];
483
+ if (k_row1 < n_kv) s1 += slope * (ACC_TYPE)mask_ptr[k_row1];
484
+ if (k_row2 < n_kv) s2 += slope * (ACC_TYPE)mask_ptr[k_row2];
485
+ if (k_row3 < n_kv) s3 += slope * (ACC_TYPE)mask_ptr[k_row3];
486
+ }
487
+ }
488
+
489
+ if (logit_softcap > 0.0f) {
490
+ s0 = logit_softcap * tanh(s0 / logit_softcap);
491
+ s1 = logit_softcap * tanh(s1 / logit_softcap);
492
+ s2 = logit_softcap * tanh(s2 / logit_softcap);
493
+ s3 = logit_softcap * tanh(s3 / logit_softcap);
494
+ }
495
+
496
+ const ACC_TYPE m_new = max(m_i, max(max(s0, s1), max(s2, s3)));
497
+ // Whole tile masked (m_new == FA_M_INIT): force the exp() args
498
+ // far negative so the tile contributes 0, not exp(0)=1.
499
+ const ACC_TYPE m_exp = (m_new == FA_M_INIT) ? 0.0f : m_new;
500
+ const ACC_TYPE scale_prev = native_exp(m_i - m_exp);
501
+ const ACC_TYPE p0 = native_exp(s0 - m_exp);
502
+ const ACC_TYPE p1 = native_exp(s1 - m_exp);
503
+ const ACC_TYPE p2 = native_exp(s2 - m_exp);
504
+ const ACC_TYPE p3 = native_exp(s3 - m_exp);
505
+
506
+ FA_UNROLL
507
+ for (int i = 0; i < DV_VEC; ++i) {
508
+ o_acc[i] = mad(p3, CONVERT_KV_ACC4(l_v[j+3][i]),
509
+ mad(p2, CONVERT_KV_ACC4(l_v[j+2][i]),
510
+ mad(p1, CONVERT_KV_ACC4(l_v[j+1][i]),
511
+ mad(p0, CONVERT_KV_ACC4(l_v[j][i]),
512
+ o_acc[i] * scale_prev))));
513
+ }
514
+ l_i = l_i * scale_prev + p0 + p1 + p2 + p3;
515
+ m_i = m_new;
516
+ }
517
+ }
518
+ #endif
519
+ // End of tile: every thread must finish reading l_k/l_v before the
520
+ // next iteration's load overwrites them (WAR hazard on local memory).
521
+ barrier(CLK_LOCAL_MEM_FENCE);
522
+ }
523
+
524
+ // Write output.
525
+ #if N_SPLIT > 1 && defined(HAS_SUBGROUP_SHUFFLE)
526
+ if (query_valid) {
527
+ ACC_TYPE sinks_sp = 1.0f;
528
+ if (sinks_void != NULL) {
529
+ const global ACC_TYPE* sinks_ptr = (const global ACC_TYPE*)((const global char*)sinks_void + sinks_offset);
530
+ const ACC_TYPE m_sink = sinks_ptr[head_idx];
531
+ const ACC_TYPE m_final = max(m_i, m_sink);
532
+ sinks_sp = exp(m_i - m_final);
533
+ l_i = l_i * sinks_sp + exp(m_sink - m_final);
534
+ m_i = m_final;
535
+ }
536
+ const ACC_TYPE l_inv = (l_i > 0.0f) ? (1.0f / l_i) : 0.0f;
537
+ const int dv_off = split_idx * SPLIT_DV_VEC;
538
+ const ulong o_row_offset = batch_idx * o_nb3 + my_query_row * o_nb2 + head_idx * o_nb1;
539
+ global O_DATA_TYPE4 *o_row = (global O_DATA_TYPE4 *)(o_base + o_row_offset);
540
+ if (l_inv > 0.0f) {
541
+ FA_UNROLL
542
+ for (int i = 0; i < SPLIT_DV_VEC; ++i) {
543
+ o_row[dv_off + i] = CONVERT_O_DATA4(o_acc[i] * sinks_sp * l_inv);
172
544
  }
173
- l_i = l_i * scale_prev + p0 + p1;
174
- m_i = m_new;
545
+ } else {
546
+ FA_UNROLL
547
+ for (int i = 0; i < SPLIT_DV_VEC; ++i) {
548
+ o_row[dv_off + i] = (O_DATA_TYPE4)(0.0f);
549
+ }
550
+ }
551
+ }
552
+ #elif N_SPLIT > 1
553
+ if (split_idx == 0) {
554
+ ACC_TYPE sinks_sp = 1.0f;
555
+ if (query_valid && sinks_void != NULL) {
556
+ const global ACC_TYPE* sinks_ptr = (const global ACC_TYPE*)((const global char*)sinks_void + sinks_offset);
557
+ const ACC_TYPE m_sink = sinks_ptr[head_idx];
558
+ const ACC_TYPE m_final = max(m_i, m_sink);
559
+ sinks_sp = exp(m_i - m_final);
560
+ l_i = l_i * sinks_sp + exp(m_sink - m_final);
561
+ m_i = m_final;
175
562
  }
563
+ local_softmax_scale[q_lane] = sinks_sp;
564
+ local_l_inv[q_lane] = (query_valid && l_i > 0.0f) ? (1.0f / l_i) : 0.0f;
176
565
  }
566
+ barrier(CLK_LOCAL_MEM_FENCE);
177
567
 
178
- if (my_query_row < n_q) {
568
+ if (query_valid) {
569
+ const ACC_TYPE sinks_sp = local_softmax_scale[q_lane];
570
+ const ACC_TYPE l_inv = local_l_inv[q_lane];
571
+ const int dv_off = split_idx * SPLIT_DV_VEC;
572
+ const ulong o_row_offset = batch_idx * o_nb3 + my_query_row * o_nb2 + head_idx * o_nb1;
573
+ global O_DATA_TYPE4 *o_row = (global O_DATA_TYPE4 *)(o_base + o_row_offset);
574
+ if (l_inv > 0.0f) {
575
+ FA_UNROLL
576
+ for (int i = 0; i < SPLIT_DV_VEC; ++i) {
577
+ o_row[dv_off + i] = CONVERT_O_DATA4(o_acc[i] * sinks_sp * l_inv);
578
+ }
579
+ } else {
580
+ FA_UNROLL
581
+ for (int i = 0; i < SPLIT_DV_VEC; ++i) {
582
+ o_row[dv_off + i] = (O_DATA_TYPE4)(0.0f);
583
+ }
584
+ }
585
+ }
586
+ #else
587
+ if (query_valid) {
179
588
  if (sinks_void != NULL) {
180
589
  const global ACC_TYPE* sinks_ptr = (const global ACC_TYPE*)((const global char*)sinks_void + sinks_offset);
181
590
  const ACC_TYPE m_sink = sinks_ptr[head_idx];
182
591
  const ACC_TYPE m_final = max(m_i, m_sink);
183
592
 
184
593
  const ACC_TYPE scale_o = exp(m_i - m_final);
185
- #pragma unroll
594
+ FA_UNROLL
186
595
  for (int i = 0; i < DV_VEC; ++i) {
187
596
  o_acc[i] *= scale_o;
188
597
  }
@@ -194,19 +603,25 @@ __kernel void flash_attn_f32_f16(
194
603
  global O_DATA_TYPE4 *o_row = (global O_DATA_TYPE4 *)(o_base + o_row_offset);
195
604
  if (l_i > 0.0f) {
196
605
  const ACC_TYPE l_inv = 1.0f / l_i;
197
- #pragma unroll
606
+ FA_UNROLL
198
607
  for (int i = 0; i < DV_VEC; ++i) {
199
608
  o_row[i] = CONVERT_O_DATA4(o_acc[i] * l_inv);
200
609
  }
201
610
  } else {
202
- #pragma unroll
611
+ FA_UNROLL
203
612
  for (int i = 0; i < DV_VEC; ++i) {
204
613
  o_row[i] = (O_DATA_TYPE4)(0.0f);
205
614
  }
206
615
  }
207
616
  }
617
+ #endif
208
618
  }
619
+ #endif // !FA_DECODE_ONLY
209
620
 
621
+ // allow bypassing decode kernels to avoid compiler crash for DK=512 on Adreno GPUs
622
+ #ifndef FA_PREFILL_ONLY
623
+ #ifndef FA_MQ_ONLY // q1 excluded from the MQ-only (g8) program
624
+ REQD_FA_SG
210
625
  __kernel void flash_attn_f32_f16_q1(
211
626
  const global void * q_void, ulong q_offset,
212
627
  const global void * k_void, ulong k_offset,
@@ -247,7 +662,9 @@ __kernel void flash_attn_f32_f16_q1(
247
662
  const int head_kv_idx = head_idx / gqa_ratio;
248
663
 
249
664
  const global char* q_base = (const global char*)q_void + q_offset;
665
+ #ifndef FA_K_IMG
250
666
  const global char* k_base = (const global char*)k_void + k_offset;
667
+ #endif
251
668
  const global char* v_base = (const global char*)v_void + v_offset;
252
669
  global char* o_base = (global char*)o_void + o_offset;
253
670
 
@@ -258,13 +675,16 @@ __kernel void flash_attn_f32_f16_q1(
258
675
  mask_base = (const global char*)mask_void + mask_offset + mask_batch_idx * mask_nb3 + mask_head_idx * mask_nb2;
259
676
  }
260
677
 
261
- ACC_TYPE4 q_priv[DK_VEC];
678
+ // Q is uniform across WG threads (n_q=1). Share via local memory to
679
+ // avoid per-thread q_priv[DK_VEC] dynamic-indexed private array that
680
+ // spills to DDR on Adreno.
681
+ __local ACC_TYPE4 q_shared[DK_VEC];
262
682
  const ulong q_row_offset = batch_idx * q_nb3 + head_idx * q_nb2;
263
683
  const global Q_DATA_TYPE4* q_ptr = (const global Q_DATA_TYPE4*)(q_base + q_row_offset);
264
- #pragma unroll
265
- for (int i = 0; i < DK_VEC; ++i) {
266
- q_priv[i] = CONVERT_Q_ACC4(q_ptr[i]);
684
+ for (int i = tid; i < DK_VEC; i += Q1_WG_SIZE) {
685
+ q_shared[i] = CONVERT_Q_ACC4(q_ptr[i]);
267
686
  }
687
+ sub_group_barrier(CLK_LOCAL_MEM_FENCE);
268
688
 
269
689
  float slope = get_alibi_slope(max_bias, head_idx, n_head_log2, m0, m1);
270
690
 
@@ -273,14 +693,14 @@ __kernel void flash_attn_f32_f16_q1(
273
693
  sinks_ptr = (const global ACC_TYPE*)((const global char*)sinks_void + sinks_offset);
274
694
  }
275
695
 
276
- ACC_TYPE m_i = (sinks_ptr != NULL) ? sinks_ptr[head_idx] : -INFINITY;
696
+ ACC_TYPE m_i = (sinks_ptr != NULL) ? sinks_ptr[head_idx] : FA_M_INIT;
277
697
  for (int k_idx = tid; k_idx < n_kv; k_idx += Q1_WG_SIZE) {
278
698
  const ulong k_row_offset = batch_idx * k_nb3 + head_kv_idx * k_nb2 + k_idx * k_nb1;
279
699
  const global KV_DATA_TYPE4* k_ptr = (const global KV_DATA_TYPE4*)(k_base + k_row_offset);
280
700
  ACC_TYPE4 dot_acc = (ACC_TYPE4)(0.0f);
281
- #pragma unroll
701
+ FA_UNROLL
282
702
  for (int k = 0; k < DK_VEC; k++) {
283
- dot_acc = mad(q_priv[k], CONVERT_KV_ACC4(k_ptr[k]), dot_acc);
703
+ dot_acc = mad(q_shared[k], CONVERT_KV_ACC4(k_ptr[k]), dot_acc);
284
704
  }
285
705
  ACC_TYPE score = (dot_acc.s0 + dot_acc.s1 + dot_acc.s2 + dot_acc.s3) * scale;
286
706
  if (mask_base != NULL) {
@@ -293,18 +713,10 @@ __kernel void flash_attn_f32_f16_q1(
293
713
  m_i = max(m_i, score);
294
714
  }
295
715
 
296
- __local ACC_TYPE local_m[Q1_WG_SIZE];
297
- local_m[tid] = m_i;
298
- barrier(CLK_LOCAL_MEM_FENCE);
299
- #pragma unroll
300
- for (int s = Q1_WG_SIZE / 2; s > 0; s >>= 1) {
301
- if (tid < s) local_m[tid] = max(local_m[tid], local_m[tid + s]);
302
- barrier(CLK_LOCAL_MEM_FENCE);
303
- }
304
- const ACC_TYPE m_final = local_m[0];
716
+ const ACC_TYPE m_final = sub_group_reduce_max(m_i);
305
717
 
306
718
  ACC_TYPE4 o_acc[DV_VEC];
307
- #pragma unroll
719
+ FA_UNROLL
308
720
  for (int i = 0; i < DV_VEC; ++i) o_acc[i] = (ACC_TYPE4)(0.0f);
309
721
  ACC_TYPE l_i = 0.0f;
310
722
 
@@ -314,9 +726,9 @@ __kernel void flash_attn_f32_f16_q1(
314
726
  const global KV_DATA_TYPE4* k_ptr = (const global KV_DATA_TYPE4*)(k_base + k_row_offset);
315
727
  const global KV_DATA_TYPE4* v_ptr = (const global KV_DATA_TYPE4*)(v_base + v_row_offset);
316
728
  ACC_TYPE4 dot_acc = (ACC_TYPE4)(0.0f);
317
- #pragma unroll
729
+ FA_UNROLL
318
730
  for (int k = 0; k < DK_VEC; k++) {
319
- dot_acc = mad(q_priv[k], CONVERT_KV_ACC4(k_ptr[k]), dot_acc);
731
+ dot_acc = mad(q_shared[k], CONVERT_KV_ACC4(k_ptr[k]), dot_acc);
320
732
  }
321
733
  ACC_TYPE score = (dot_acc.s0 + dot_acc.s1 + dot_acc.s2 + dot_acc.s3) * scale;
322
734
  if (mask_base != NULL) {
@@ -328,25 +740,18 @@ __kernel void flash_attn_f32_f16_q1(
328
740
  }
329
741
  const ACC_TYPE p = exp(score - m_final);
330
742
  l_i += p;
331
- #pragma unroll
743
+ FA_UNROLL
332
744
  for (int i = 0; i < DV_VEC; i++) {
333
745
  o_acc[i] = mad(p, CONVERT_KV_ACC4(v_ptr[i]), o_acc[i]);
334
746
  }
335
747
  }
336
748
 
337
- __local ACC_TYPE local_l[Q1_WG_SIZE];
338
749
  __local ACC_TYPE4 local_o_comp[Q1_WG_SIZE];
339
- local_l[tid] = l_i;
340
- barrier(CLK_LOCAL_MEM_FENCE);
341
- #pragma unroll
342
- for (int s = Q1_WG_SIZE / 2; s > 0; s >>= 1) {
343
- if (tid < s) local_l[tid] += local_l[tid + s];
344
- barrier(CLK_LOCAL_MEM_FENCE);
345
- }
750
+ const ACC_TYPE l_red = sub_group_reduce_add(l_i);
346
751
 
347
752
  const ulong o_row_offset = batch_idx * o_nb3 + head_idx * o_nb1;
348
753
  global O_DATA_TYPE4 *o_row = (global O_DATA_TYPE4 *)(o_base + o_row_offset);
349
- ACC_TYPE l_final = local_l[0];
754
+ ACC_TYPE l_final = l_red;
350
755
 
351
756
  if (sinks_ptr != NULL) {
352
757
  l_final += exp(sinks_ptr[head_idx] - m_final);
@@ -356,18 +761,1512 @@ __kernel void flash_attn_f32_f16_q1(
356
761
  const ACC_TYPE l_inv = 1.0f / l_final;
357
762
  for (int i = 0; i < DV_VEC; i++) {
358
763
  local_o_comp[tid] = o_acc[i];
359
- barrier(CLK_LOCAL_MEM_FENCE);
360
- #pragma unroll
764
+ sub_group_barrier(CLK_LOCAL_MEM_FENCE);
765
+ FA_UNROLL
361
766
  for (int s = Q1_WG_SIZE / 2; s > 0; s >>= 1) {
362
767
  if (tid < s) local_o_comp[tid] += local_o_comp[tid + s];
363
- barrier(CLK_LOCAL_MEM_FENCE);
768
+ sub_group_barrier(CLK_LOCAL_MEM_FENCE);
364
769
  }
365
770
  if (tid == 0) {
366
771
  o_row[i] = CONVERT_O_DATA4(local_o_comp[0] * l_inv);
367
772
  }
368
773
  }
369
774
  } else if (tid == 0) {
370
- #pragma unroll
775
+ FA_UNROLL
371
776
  for (int i = 0; i < DV_VEC; ++i) o_row[i] = (O_DATA_TYPE4)(0.0f);
372
777
  }
373
778
  }
779
+
780
+ #endif // !FA_MQ_ONLY (q1)
781
+ // decode variant for large DV (e.g. Gemma-4 DK=DV=512 global layers).
782
+ #define VEC_NSG 4
783
+ #define VEC_WG_SIZE (Q1_WG_SIZE * VEC_NSG)
784
+ #define Q1V_DV_PER_THREAD ((DV_VEC + Q1_WG_SIZE - 1) / Q1_WG_SIZE)
785
+
786
+ // allow bypassing the kernel to avoid compiler crash for DK=512 on Adreno GPUs
787
+ #if !defined(FA_DECODE_MINIMAL) && !defined(FA_MQ_ONLY)
788
+ REQD_SUBGROUP_SIZE_64
789
+ __kernel void flash_attn_f32_f16_q1_vec(
790
+ const global void * q_void, ulong q_offset,
791
+ const global void * k_void, ulong k_offset,
792
+ const global void * v_void, ulong v_offset,
793
+ global void * o_void, ulong o_offset,
794
+ const float scale,
795
+ const int n_q,
796
+ const int n_kv,
797
+ const int is_causal,
798
+ const int n_head,
799
+ const ulong q_nb1, const ulong q_nb2, const ulong q_nb3,
800
+ const ulong k_nb1, const ulong k_nb2, const ulong k_nb3,
801
+ const ulong v_nb1, const ulong v_nb2, const ulong v_nb3,
802
+ const ulong o_nb1, const ulong o_nb2, const ulong o_nb3,
803
+ const float max_bias,
804
+ const float m0,
805
+ const float m1,
806
+ const int n_head_log2,
807
+ const float logit_softcap,
808
+ const int n_head_kv,
809
+ const global void* mask_void,
810
+ const ulong mask_offset,
811
+ const ulong mask_nb1,
812
+ const ulong mask_nb2,
813
+ const ulong mask_nb3,
814
+ const int mask_ne2,
815
+ const int mask_ne3,
816
+ const global void* sinks_void,
817
+ const ulong sinks_offset
818
+ ) {
819
+ const int tid = get_local_id(0);
820
+ const int sgid = tid / Q1_WG_SIZE; // subgroup index (0..VEC_NSG-1)
821
+ const int tid_sg = tid % Q1_WG_SIZE; // lane within subgroup
822
+ const int head_batch_idx = get_global_id(1);
823
+
824
+ const int batch_idx = head_batch_idx / n_head;
825
+ const int head_idx = head_batch_idx % n_head;
826
+
827
+ const int gqa_ratio = n_head / n_head_kv;
828
+ const int head_kv_idx = head_idx / gqa_ratio;
829
+
830
+ const global char * q_base = (const global char *) q_void + q_offset;
831
+ const global char * k_base = (const global char *) k_void + k_offset;
832
+ const global char * v_base = (const global char *) v_void + v_offset;
833
+ global char * o_base = (global char *) o_void + o_offset;
834
+
835
+ const global char * mask_base = NULL;
836
+ if (mask_void != NULL) {
837
+ const int mask_head_idx = head_idx % mask_ne2;
838
+ const int mask_batch_idx = batch_idx % mask_ne3;
839
+ mask_base = (const global char *) mask_void + mask_offset +
840
+ mask_batch_idx * mask_nb3 + mask_head_idx * mask_nb2;
841
+ }
842
+
843
+ // Q is uniform across the WG — stage in __local once. All WG threads load.
844
+ __local ACC_TYPE4 q_shared[DK_VEC];
845
+ {
846
+ const ulong q_row_offset = batch_idx * q_nb3 + head_idx * q_nb2;
847
+ const global Q_DATA_TYPE4 * q_ptr = (const global Q_DATA_TYPE4 *) (q_base + q_row_offset);
848
+ for (int i = tid; i < DK_VEC; i += VEC_WG_SIZE) {
849
+ q_shared[i] = CONVERT_Q_ACC4(q_ptr[i]);
850
+ }
851
+ }
852
+ barrier(CLK_LOCAL_MEM_FENCE);
853
+
854
+ const float slope = get_alibi_slope(max_bias, head_idx, n_head_log2, m0, m1);
855
+
856
+ const global ACC_TYPE * sinks_ptr = NULL;
857
+ if (sinks_void != NULL) {
858
+ sinks_ptr = (const global ACC_TYPE *) ((const global char *) sinks_void + sinks_offset);
859
+ }
860
+
861
+ // per-thread DV slice within its subgroup
862
+ // DV=512 -> 2x float4 = 32 bytes; DV=256 -> 1x float4 - no spill
863
+ ACC_TYPE4 o_acc[Q1V_DV_PER_THREAD];
864
+ #pragma unroll
865
+ for (int i = 0; i < Q1V_DV_PER_THREAD; ++i) o_acc[i] = (ACC_TYPE4)(0.0f);
866
+
867
+ // each subgroup independently runs the FA-2 online softmax over its slice of n_kv.
868
+ // sinks are not folded into per-subgroup m_i — they're added once in
869
+ // the cross-subgroup merge to avoid double-counting.
870
+ ACC_TYPE m_i = FA_M_INIT;
871
+ ACC_TYPE l_i = 0.0f;
872
+
873
+ const int kv_per_sg = (n_kv + VEC_NSG - 1) / VEC_NSG;
874
+ const int kv_start = sgid * kv_per_sg;
875
+ const int kv_end = min(n_kv, kv_start + kv_per_sg);
876
+
877
+ for (int k_idx = kv_start; k_idx < kv_end; ++k_idx) {
878
+ const ulong k_row_off = batch_idx * k_nb3 + head_kv_idx * k_nb2 + k_idx * k_nb1;
879
+ const ulong v_row_off = batch_idx * v_nb3 + head_kv_idx * v_nb2 + k_idx * v_nb1;
880
+ const global KV_DATA_TYPE4 * k_ptr = (const global KV_DATA_TYPE4 *) (k_base + k_row_off);
881
+ const global KV_DATA_TYPE4 * v_ptr = (const global KV_DATA_TYPE4 *) (v_base + v_row_off);
882
+
883
+ // Q*K^T: each thread accumulates its DK slice; subgroup-reduce the partial.
884
+ ACC_TYPE4 dot4 = (ACC_TYPE4)(0.0f);
885
+ for (int k = tid_sg; k < DK_VEC; k += Q1_WG_SIZE) {
886
+ dot4 = mad(q_shared[k], CONVERT_KV_ACC4(k_ptr[k]), dot4);
887
+ }
888
+ ACC_TYPE dot_partial = dot4.s0 + dot4.s1 + dot4.s2 + dot4.s3;
889
+ ACC_TYPE score = sub_group_reduce_add(dot_partial) * scale;
890
+
891
+ if (mask_base != NULL) {
892
+ const global MASK_DATA_TYPE * mask_ptr = (const global MASK_DATA_TYPE *) mask_base;
893
+ score += slope * (ACC_TYPE) mask_ptr[k_idx];
894
+ }
895
+ if (logit_softcap > 0.0f) {
896
+ score = logit_softcap * tanh(score / logit_softcap);
897
+ }
898
+
899
+ // FA-2 online update. All threads in the subgroup see the same score,
900
+ // so m_i and l_i evolve identically across lanes within the subgroup.
901
+ const ACC_TYPE m_new = max(m_i, score);
902
+ const ACC_TYPE scale_prev = native_exp(m_i - m_new);
903
+ const ACC_TYPE p = native_exp(score - m_new);
904
+
905
+ int idx = 0;
906
+ for (int dv_idx = tid_sg; dv_idx < DV_VEC; dv_idx += Q1_WG_SIZE, ++idx) {
907
+ o_acc[idx] = mad(p, CONVERT_KV_ACC4(v_ptr[dv_idx]), o_acc[idx] * scale_prev);
908
+ }
909
+ l_i = l_i * scale_prev + p;
910
+ m_i = m_new;
911
+ }
912
+
913
+ // Cross-subgroup merge via __local. Each subgroup publishes (m_i, l_i)
914
+ // and its o_acc slice; subgroup 0 then folds them into the final norm
915
+ // and writes the row.
916
+ __local ACC_TYPE sg_m[VEC_NSG];
917
+ __local ACC_TYPE sg_l[VEC_NSG];
918
+ __local ACC_TYPE4 sg_o[VEC_NSG][DV_VEC];
919
+
920
+ if (tid_sg == 0) {
921
+ sg_m[sgid] = m_i;
922
+ sg_l[sgid] = l_i;
923
+ }
924
+ {
925
+ int idx = 0;
926
+ for (int dv_idx = tid_sg; dv_idx < DV_VEC; dv_idx += Q1_WG_SIZE, ++idx) {
927
+ sg_o[sgid][dv_idx] = o_acc[idx];
928
+ }
929
+ }
930
+ barrier(CLK_LOCAL_MEM_FENCE);
931
+
932
+ if (sgid == 0) {
933
+ // m_final = max over all subgroups' m_i, plus the sink (if any).
934
+ ACC_TYPE m_final = sg_m[0];
935
+ #pragma unroll
936
+ for (int s = 1; s < VEC_NSG; ++s) {
937
+ m_final = max(m_final, sg_m[s]);
938
+ }
939
+ if (sinks_ptr != NULL) {
940
+ m_final = max(m_final, sinks_ptr[head_idx]);
941
+ }
942
+
943
+ ACC_TYPE l_final = 0.0f;
944
+ #pragma unroll
945
+ for (int s = 0; s < VEC_NSG; ++s) {
946
+ l_final += sg_l[s] * native_exp(sg_m[s] - m_final);
947
+ }
948
+ if (sinks_ptr != NULL) {
949
+ l_final += native_exp(sinks_ptr[head_idx] - m_final);
950
+ }
951
+ const ACC_TYPE l_inv = (l_final > 0.0f) ? (1.0f / l_final) : 0.0f;
952
+
953
+ const ulong o_row_offset = batch_idx * o_nb3 + head_idx * o_nb1;
954
+ global O_DATA_TYPE4 * o_row = (global O_DATA_TYPE4 *) (o_base + o_row_offset);
955
+
956
+ // Each thread in subgroup 0 writes its DV slice, folding all subgroups'
957
+ // contributions with the rescale factor.
958
+ int idx = 0;
959
+ for (int dv_idx = tid_sg; dv_idx < DV_VEC; dv_idx += Q1_WG_SIZE, ++idx) {
960
+ ACC_TYPE4 o_merged = (ACC_TYPE4)(0.0f);
961
+ #pragma unroll
962
+ for (int s = 0; s < VEC_NSG; ++s) {
963
+ const ACC_TYPE alpha = native_exp(sg_m[s] - m_final);
964
+ o_merged = mad((ACC_TYPE4)(alpha), sg_o[s][dv_idx], o_merged);
965
+ }
966
+ o_row[dv_idx] = CONVERT_O_DATA4(o_merged * l_inv);
967
+ }
968
+ }
969
+ }
970
+
971
+ #endif // !FA_DECODE_MINIMAL
972
+
973
+ #ifndef FA_DECODE_ONLY
974
+
975
+ // flash_attn_f32_f16_q1_local_tile
976
+ // one WG per (q_idx, q_head)
977
+
978
+ #define LT_KC 32
979
+ #define LT_WG 128
980
+
981
+ #ifndef FA_MQ_ONLY // q1_local_tile excluded from the MQ-only (g8) program
982
+ REQD_SUBGROUP_SIZE_64
983
+ __kernel void flash_attn_f32_f16_q1_local_tile(
984
+ const global void * q_void, ulong q_offset,
985
+ const global void * k_void, ulong k_offset,
986
+ const global void * v_void, ulong v_offset,
987
+ global void * o_void, ulong o_offset,
988
+ const float scale,
989
+ const int n_q,
990
+ const int n_kv,
991
+ const int is_causal,
992
+ const int n_head,
993
+ const ulong q_nb1, const ulong q_nb2, const ulong q_nb3,
994
+ const ulong k_nb1, const ulong k_nb2, const ulong k_nb3,
995
+ const ulong v_nb1, const ulong v_nb2, const ulong v_nb3,
996
+ const ulong o_nb1, const ulong o_nb2, const ulong o_nb3,
997
+ const float max_bias,
998
+ const float m0,
999
+ const float m1,
1000
+ const int n_head_log2,
1001
+ const float logit_softcap,
1002
+ const int n_head_kv,
1003
+ const global void * mask_void,
1004
+ const ulong mask_offset,
1005
+ const ulong mask_nb1,
1006
+ const ulong mask_nb2,
1007
+ const ulong mask_nb3,
1008
+ const int mask_ne2,
1009
+ const int mask_ne3,
1010
+ const global void * sinks_void,
1011
+ const ulong sinks_offset
1012
+ ) {
1013
+ const int q_idx = get_global_id(0) / LT_WG;
1014
+ const int head_idx = get_global_id(1);
1015
+ const int batch_idx = get_global_id(2);
1016
+ const int tid = get_local_id(0);
1017
+
1018
+ const int gqa_ratio = n_head_kv > 0 ? (n_head / n_head_kv) : 1;
1019
+ const int head_kv_idx = head_idx / gqa_ratio;
1020
+
1021
+ const float slope = get_alibi_slope(max_bias, head_idx, n_head_log2, m0, m1);
1022
+
1023
+ __local half k_tile[LT_KC * DK]; // 32*128*2 = 8 KB at DK=128
1024
+ __local half v_tile[LT_KC * DV]; // 8 KB
1025
+ __local float red[LT_WG]; // 512 B reduction scratch
1026
+ __local float score_shared; // broadcast score (each K-step)
1027
+
1028
+ // Each thread owns one float of Q at index `tid` (assumes LT_WG == DK).
1029
+ const global char * q_row_base = (const global char *) q_void + q_offset +
1030
+ batch_idx * q_nb3 + head_idx * q_nb2 + (ulong) q_idx * q_nb1;
1031
+ float q_val = ((const global float *) q_row_base)[tid];
1032
+
1033
+ const global char * mask_base = NULL;
1034
+ if (mask_void != NULL) {
1035
+ const int mask_head_idx = head_idx % mask_ne2;
1036
+ const int mask_batch_idx = batch_idx % mask_ne3;
1037
+ mask_base = (const global char *) mask_void + mask_offset +
1038
+ mask_batch_idx * mask_nb3 + mask_head_idx * mask_nb2 +
1039
+ (ulong) q_idx * mask_nb1;
1040
+ }
1041
+
1042
+ float o_val = 0.0f;
1043
+ float m_i = FA_M_INIT;
1044
+ float l_i = 0.0f;
1045
+
1046
+ for (int kb = 0; kb < n_kv; kb += LT_KC) {
1047
+ const int tile_len = min(LT_KC, n_kv - kb);
1048
+
1049
+ // Stage K and V tiles into __local.
1050
+ for (int i = tid; i < tile_len * DK; i += LT_WG) {
1051
+ const int j = i / DK;
1052
+ const int d = i % DK;
1053
+ const int kv_idx = kb + j;
1054
+ const global char * k_row = (const global char *) k_void + k_offset +
1055
+ batch_idx * k_nb3 + head_kv_idx * k_nb2 +
1056
+ (ulong) kv_idx * k_nb1;
1057
+ const global char * v_row = (const global char *) v_void + v_offset +
1058
+ batch_idx * v_nb3 + head_kv_idx * v_nb2 +
1059
+ (ulong) kv_idx * v_nb1;
1060
+ k_tile[j * DK + d] = ((const global half *) k_row)[d];
1061
+ v_tile[j * DV + d] = ((const global half *) v_row)[d];
1062
+ }
1063
+ barrier(CLK_LOCAL_MEM_FENCE);
1064
+
1065
+ for (int j = 0; j < tile_len; ++j) {
1066
+ const int kv_idx = kb + j;
1067
+
1068
+ // Q·K dot via __local tree-reduce.
1069
+ red[tid] = q_val * convert_float(k_tile[j * DK + tid]);
1070
+ barrier(CLK_LOCAL_MEM_FENCE);
1071
+ for (int stride = LT_WG >> 1; stride > 0; stride >>= 1) {
1072
+ if (tid < stride) {
1073
+ red[tid] += red[tid + stride];
1074
+ }
1075
+ barrier(CLK_LOCAL_MEM_FENCE);
1076
+ }
1077
+
1078
+ if (tid == 0) {
1079
+ float s = red[0] * scale;
1080
+ if (mask_base != NULL) {
1081
+ const global MASK_DATA_TYPE * mask_ptr = (const global MASK_DATA_TYPE *) mask_base;
1082
+ s += slope * (float) mask_ptr[kv_idx];
1083
+ }
1084
+ if (logit_softcap > 0.0f) {
1085
+ s = logit_softcap * tanh(s / logit_softcap);
1086
+ }
1087
+ score_shared = s;
1088
+ }
1089
+ barrier(CLK_LOCAL_MEM_FENCE);
1090
+
1091
+ const float s = score_shared;
1092
+ const float m_new = fmax(m_i, s);
1093
+ const float alpha = native_exp(m_i - m_new);
1094
+ const float beta = native_exp(s - m_new);
1095
+
1096
+ o_val = o_val * alpha + beta * convert_float(v_tile[j * DV + tid]);
1097
+ l_i = l_i * alpha + beta;
1098
+ m_i = m_new;
1099
+ }
1100
+ barrier(CLK_LOCAL_MEM_FENCE);
1101
+ }
1102
+
1103
+ // Fold attention sinks into the running (m, l, o), if present.
1104
+ if (sinks_void != NULL) {
1105
+ const global float * sinks_ptr =
1106
+ (const global float *) ((const global char *) sinks_void + sinks_offset);
1107
+ const float m_sink = sinks_ptr[head_idx];
1108
+ const float m_new = fmax(m_i, m_sink);
1109
+ const float alpha = native_exp(m_i - m_new);
1110
+ const float beta = native_exp(m_sink - m_new);
1111
+ o_val = o_val * alpha;
1112
+ l_i = l_i * alpha + beta;
1113
+ m_i = m_new;
1114
+ }
1115
+
1116
+ const float l_inv = (l_i > 0.0f) ? (1.0f / l_i) : 0.0f;
1117
+ global float * o_row = (global float *) ((global char *) o_void + o_offset +
1118
+ batch_idx * o_nb3 + head_idx * o_nb1 +
1119
+ (ulong) q_idx * o_nb2);
1120
+ o_row[tid] = o_val * l_inv;
1121
+ }
1122
+
1123
+ // flash_attn_f32_f16_q1_local_mq_split
1124
+
1125
+ #define LMQ_WG 64
1126
+ #define LMQ_KC 32
1127
+ #define LMQ_DPL 2 // DK / LMQ_WG at DK=128
1128
+
1129
+ #endif // !FA_MQ_ONLY (q1_local_tile)
1130
+ #ifndef MQ_GQA
1131
+ #define MQ_GQA 4
1132
+ #endif
1133
+
1134
+ #ifndef FA_PARTIAL_FLOATS
1135
+ #define FA_PARTIAL_FLOATS (2 + DV)
1136
+ #endif
1137
+
1138
+ #ifndef FA_MQ_ONLY // q1_local_mq_split excluded from the MQ-only (g8) program
1139
+ REQD_SUBGROUP_SIZE_64
1140
+ __kernel void flash_attn_f32_f16_q1_local_mq_split(
1141
+ const global void * q_void, ulong q_offset,
1142
+ const global void * k_void, ulong k_offset,
1143
+ const global void * v_void, ulong v_offset,
1144
+ const float scale,
1145
+ const int n_q,
1146
+ const int n_kv,
1147
+ const int n_head,
1148
+ const ulong q_nb1, const ulong q_nb2, const ulong q_nb3,
1149
+ const ulong k_nb1, const ulong k_nb2, const ulong k_nb3,
1150
+ const ulong v_nb1, const ulong v_nb2, const ulong v_nb3,
1151
+ const float max_bias,
1152
+ const float m0,
1153
+ const float m1,
1154
+ const int n_head_log2,
1155
+ const float logit_softcap,
1156
+ const int n_head_kv,
1157
+ const global void * mask_void,
1158
+ const ulong mask_offset,
1159
+ const ulong mask_nb1,
1160
+ const ulong mask_nb2,
1161
+ const ulong mask_nb3,
1162
+ const int mask_ne2,
1163
+ const int mask_ne3,
1164
+ global float * partial_void,
1165
+ const int n_splits,
1166
+ const int kv_per_split
1167
+ ) {
1168
+ const int tid = get_local_id(0); // 0..LMQ_WG-1
1169
+ const int kvhead_batch_idx = get_global_id(1);
1170
+ const int split_q_idx = get_global_id(2);
1171
+ const int split_idx = split_q_idx % n_splits;
1172
+ const int q_idx = split_q_idx / n_splits;
1173
+
1174
+ const int batch_idx = kvhead_batch_idx / n_head_kv;
1175
+ const int head_kv_idx = kvhead_batch_idx % n_head_kv;
1176
+
1177
+ const int kv_start = split_idx * kv_per_split;
1178
+ const int kv_end = min(kv_start + kv_per_split, n_kv);
1179
+
1180
+ const ulong record_stride = (ulong) FA_PARTIAL_FLOATS;
1181
+
1182
+ if (kv_start >= kv_end) {
1183
+ // Empty split — write sentinel for each Q-head so merge treats it as 0.
1184
+ if (tid == 0) {
1185
+ #pragma unroll
1186
+ for (int h = 0; h < MQ_GQA; ++h) {
1187
+ const int head_idx = head_kv_idx * MQ_GQA + h;
1188
+ const ulong rec_idx = ((((ulong) batch_idx * n_head + head_idx) * n_q + q_idx)
1189
+ * n_splits + split_idx);
1190
+ global float * rec = partial_void + rec_idx * record_stride;
1191
+ rec[0] = FA_M_INIT;
1192
+ rec[1] = 0.0f;
1193
+ }
1194
+ }
1195
+ return;
1196
+ }
1197
+
1198
+ const global char * q_base = (const global char *) q_void + q_offset;
1199
+ const global char * k_base = (const global char *) k_void + k_offset;
1200
+ const global char * v_base = (const global char *) v_void + v_offset;
1201
+
1202
+ // Stage MQ_GQA Q rows in __local (MQ_GQA × DK floats).
1203
+ __local float q_shared[MQ_GQA * DK];
1204
+ for (int i = tid; i < MQ_GQA * DK; i += LMQ_WG) {
1205
+ const int h = i / DK;
1206
+ const int d = i % DK;
1207
+ const int head_idx = head_kv_idx * MQ_GQA + h;
1208
+ const ulong q_row_off = batch_idx * q_nb3 + head_idx * q_nb2 + (ulong) q_idx * q_nb1;
1209
+ const global float * q_ptr = (const global float *) (q_base + q_row_off);
1210
+ q_shared[h * DK + d] = q_ptr[d];
1211
+ }
1212
+
1213
+ // K/V tile staging buffers (16 KB combined at DK=DV=128 KC=32).
1214
+ __local half k_tile[LMQ_KC * DK];
1215
+ __local half v_tile[LMQ_KC * DV];
1216
+
1217
+ // Per-h state held in private registers.
1218
+ float o_acc[MQ_GQA][LMQ_DPL];
1219
+ float m_i[MQ_GQA];
1220
+ float l_i[MQ_GQA];
1221
+ float slope[MQ_GQA];
1222
+ #pragma unroll
1223
+ for (int h = 0; h < MQ_GQA; ++h) {
1224
+ m_i[h] = FA_M_INIT;
1225
+ l_i[h] = 0.0f;
1226
+ slope[h] = get_alibi_slope(max_bias, head_kv_idx * MQ_GQA + h, n_head_log2, m0, m1);
1227
+ #pragma unroll
1228
+ for (int p = 0; p < LMQ_DPL; ++p) o_acc[h][p] = 0.0f;
1229
+ }
1230
+
1231
+ // Per-h mask pointers.
1232
+ const global char * mask_base[MQ_GQA];
1233
+ if (mask_void != NULL) {
1234
+ const int mask_batch_idx = batch_idx % mask_ne3;
1235
+ const global char * mask_base_b = (const global char *) mask_void + mask_offset +
1236
+ mask_batch_idx * mask_nb3 +
1237
+ (ulong) q_idx * mask_nb1;
1238
+ #pragma unroll
1239
+ for (int h = 0; h < MQ_GQA; ++h) {
1240
+ const int head_idx = head_kv_idx * MQ_GQA + h;
1241
+ const int mask_head_idx = head_idx % mask_ne2;
1242
+ mask_base[h] = mask_base_b + mask_head_idx * mask_nb2;
1243
+ }
1244
+ } else {
1245
+ #pragma unroll
1246
+ for (int h = 0; h < MQ_GQA; ++h) mask_base[h] = NULL;
1247
+ }
1248
+
1249
+ barrier(CLK_LOCAL_MEM_FENCE); // Ensure Q staged before first dot.
1250
+
1251
+ for (int kb = kv_start; kb < kv_end; kb += LMQ_KC) {
1252
+ const int tile_len = min((int) LMQ_KC, kv_end - kb);
1253
+
1254
+ // Cooperative load K + V tile.
1255
+ for (int i = tid; i < tile_len * DK; i += LMQ_WG) {
1256
+ const int j = i / DK;
1257
+ const int d = i % DK;
1258
+ const int kv_idx = kb + j;
1259
+ const global char * k_row = k_base + batch_idx * k_nb3 + head_kv_idx * k_nb2 + (ulong) kv_idx * k_nb1;
1260
+ const global char * v_row = v_base + batch_idx * v_nb3 + head_kv_idx * v_nb2 + (ulong) kv_idx * v_nb1;
1261
+ k_tile[j * DK + d] = ((const global half *) k_row)[d];
1262
+ v_tile[j * DV + d] = ((const global half *) v_row)[d];
1263
+ }
1264
+ barrier(CLK_LOCAL_MEM_FENCE);
1265
+
1266
+ // Process each cache row in the tile.
1267
+ for (int j = 0; j < tile_len; ++j) {
1268
+ const int kv_idx = kb + j;
1269
+
1270
+ // Dot product per h: lane owns LMQ_DPL D-elements at (tid*LMQ_DPL..).
1271
+ float score[MQ_GQA];
1272
+ #pragma unroll
1273
+ for (int h = 0; h < MQ_GQA; ++h) {
1274
+ float contrib = 0.0f;
1275
+ #pragma unroll
1276
+ for (int p = 0; p < LMQ_DPL; ++p) {
1277
+ const int d = tid * LMQ_DPL + p;
1278
+ contrib += q_shared[h * DK + d] * (float) k_tile[j * DK + d];
1279
+ }
1280
+ float s = sub_group_reduce_add(contrib) * scale;
1281
+ if (mask_base[h] != NULL) {
1282
+ const global MASK_DATA_TYPE * mask_ptr = (const global MASK_DATA_TYPE *) mask_base[h];
1283
+ s += slope[h] * (float) mask_ptr[kv_idx];
1284
+ }
1285
+ if (logit_softcap > 0.0f) {
1286
+ s = logit_softcap * tanh(s / logit_softcap);
1287
+ }
1288
+ score[h] = s;
1289
+ }
1290
+
1291
+ // Online softmax update + V accumulation per h.
1292
+ float p_h[MQ_GQA];
1293
+ float sp_h[MQ_GQA];
1294
+ #pragma unroll
1295
+ for (int h = 0; h < MQ_GQA; ++h) {
1296
+ const float m_new = fmax(m_i[h], score[h]);
1297
+ sp_h[h] = native_exp(m_i[h] - m_new);
1298
+ p_h[h] = native_exp(score[h] - m_new);
1299
+ l_i[h] = l_i[h] * sp_h[h] + p_h[h];
1300
+ m_i[h] = m_new;
1301
+ }
1302
+
1303
+ #pragma unroll
1304
+ for (int p = 0; p < LMQ_DPL; ++p) {
1305
+ const int d = tid * LMQ_DPL + p;
1306
+ const float v_val = (float) v_tile[j * DV + d];
1307
+ #pragma unroll
1308
+ for (int h = 0; h < MQ_GQA; ++h) {
1309
+ o_acc[h][p] = o_acc[h][p] * sp_h[h] + p_h[h] * v_val;
1310
+ }
1311
+ }
1312
+ }
1313
+ barrier(CLK_LOCAL_MEM_FENCE); // Before next tile load overwrites k/v_tile.
1314
+ }
1315
+
1316
+ // write partial records: one per (h, split)
1317
+ #pragma unroll
1318
+ for (int h = 0; h < MQ_GQA; ++h) {
1319
+ const int head_idx = head_kv_idx * MQ_GQA + h;
1320
+ const ulong rec_idx = ((((ulong) batch_idx * n_head + head_idx) * n_q + q_idx)
1321
+ * n_splits + split_idx);
1322
+ global float * rec = partial_void + rec_idx * record_stride;
1323
+ global float * rec_o = rec + 2;
1324
+
1325
+ if (tid == 0) {
1326
+ rec[0] = m_i[h];
1327
+ rec[1] = l_i[h];
1328
+ }
1329
+ #pragma unroll
1330
+ for (int p = 0; p < LMQ_DPL; ++p) {
1331
+ const int d = tid * LMQ_DPL + p;
1332
+ rec_o[d] = o_acc[h][p];
1333
+ }
1334
+ }
1335
+ }
1336
+
1337
+ #endif // !FA_MQ_ONLY (q1_local_mq_split)
1338
+ #ifndef MQ_NSG
1339
+ #define MQ_NSG 4
1340
+ #endif
1341
+ #define MQ_WG_SIZE (Q1_WG_SIZE * MQ_NSG)
1342
+
1343
+ REQD_SUBGROUP_SIZE_64
1344
+ __kernel void flash_attn_f32_f16_q1_vec_mq(
1345
+ const global void * q_void, ulong q_offset,
1346
+ const global void * k_void, ulong k_offset,
1347
+ const global void * v_void, ulong v_offset,
1348
+ global void * o_void, ulong o_offset,
1349
+ const float scale,
1350
+ const int n_q,
1351
+ const int n_kv,
1352
+ const int is_causal,
1353
+ const int n_head,
1354
+ const ulong q_nb1, const ulong q_nb2, const ulong q_nb3,
1355
+ const ulong k_nb1, const ulong k_nb2, const ulong k_nb3,
1356
+ const ulong v_nb1, const ulong v_nb2, const ulong v_nb3,
1357
+ const ulong o_nb1, const ulong o_nb2, const ulong o_nb3,
1358
+ const float max_bias,
1359
+ const float m0,
1360
+ const float m1,
1361
+ const int n_head_log2,
1362
+ const float logit_softcap,
1363
+ const int n_head_kv,
1364
+ const global void* mask_void,
1365
+ const ulong mask_offset,
1366
+ const ulong mask_nb1,
1367
+ const ulong mask_nb2,
1368
+ const ulong mask_nb3,
1369
+ const int mask_ne2,
1370
+ const int mask_ne3,
1371
+ const global void* sinks_void,
1372
+ const ulong sinks_offset
1373
+ ) {
1374
+ const int tid = get_local_id(0);
1375
+ const int sgid = tid / Q1_WG_SIZE; // subgroup 0..MQ_NSG-1
1376
+ const int tid_sg = tid % Q1_WG_SIZE; // lane 0..63
1377
+ const int kvhead_batch_idx = get_global_id(1);
1378
+
1379
+ const int batch_idx = kvhead_batch_idx / n_head_kv;
1380
+ const int head_kv_idx = kvhead_batch_idx % n_head_kv;
1381
+
1382
+ const global char * q_base = (const global char *) q_void + q_offset;
1383
+ const global char * k_base = (const global char *) k_void + k_offset;
1384
+ const global char * v_base = (const global char *) v_void + v_offset;
1385
+ global char * o_base = (global char *) o_void + o_offset;
1386
+
1387
+ __local ACC_TYPE4 q_shared[MQ_GQA * DK_VEC];
1388
+ for (int i = tid; i < MQ_GQA * DK_VEC; i += MQ_WG_SIZE) {
1389
+ const int h = i / DK_VEC;
1390
+ const int k = i % DK_VEC;
1391
+ const int head_idx = head_kv_idx * MQ_GQA + h;
1392
+ const ulong q_row_offset = batch_idx * q_nb3 + head_idx * q_nb2;
1393
+ const global Q_DATA_TYPE4 * q_ptr = (const global Q_DATA_TYPE4 *) (q_base + q_row_offset);
1394
+ q_shared[h * DK_VEC + k] = CONVERT_Q_ACC4(q_ptr[k]);
1395
+ }
1396
+ barrier(CLK_LOCAL_MEM_FENCE);
1397
+
1398
+ // per-h ALiBi slope
1399
+ float slope[MQ_GQA];
1400
+ #pragma unroll
1401
+ for (int h = 0; h < MQ_GQA; ++h) {
1402
+ slope[h] = get_alibi_slope(max_bias, head_kv_idx * MQ_GQA + h, n_head_log2, m0, m1);
1403
+ }
1404
+
1405
+ // per-h mask row pointer
1406
+ const global char * mask_base[MQ_GQA];
1407
+ if (mask_void != NULL) {
1408
+ const int mask_batch_idx = batch_idx % mask_ne3;
1409
+ const global char * mask_base_b = (const global char *) mask_void + mask_offset +
1410
+ mask_batch_idx * mask_nb3;
1411
+ #pragma unroll
1412
+ for (int h = 0; h < MQ_GQA; ++h) {
1413
+ const int head_idx = head_kv_idx * MQ_GQA + h;
1414
+ const int mask_head_idx = head_idx % mask_ne2;
1415
+ mask_base[h] = mask_base_b + mask_head_idx * mask_nb2;
1416
+ }
1417
+ } else {
1418
+ #pragma unroll
1419
+ for (int h = 0; h < MQ_GQA; ++h) mask_base[h] = NULL;
1420
+ }
1421
+
1422
+ const global ACC_TYPE * sinks_ptr = NULL;
1423
+ if (sinks_void != NULL) {
1424
+ sinks_ptr = (const global ACC_TYPE *) ((const global char *) sinks_void + sinks_offset);
1425
+ }
1426
+
1427
+ // per-thread per-h DV slice.
1428
+ ACC_TYPE4 o_acc[MQ_GQA][Q1V_DV_PER_THREAD];
1429
+ ACC_TYPE m_i[MQ_GQA];
1430
+ ACC_TYPE l_i[MQ_GQA];
1431
+ #pragma unroll
1432
+ for (int h = 0; h < MQ_GQA; ++h) {
1433
+ m_i[h] = FA_M_INIT;
1434
+ l_i[h] = 0.0f;
1435
+ #pragma unroll
1436
+ for (int i = 0; i < Q1V_DV_PER_THREAD; ++i) o_acc[h][i] = (ACC_TYPE4)(0.0f);
1437
+ }
1438
+
1439
+ // each subgroup independently sweeps its slice of n_kv.
1440
+ const int kv_per_sg = (n_kv + MQ_NSG - 1) / MQ_NSG;
1441
+ const int kv_start = sgid * kv_per_sg;
1442
+ const int kv_end = min(n_kv, kv_start + kv_per_sg);
1443
+
1444
+ for (int k_idx = kv_start; k_idx < kv_end; ++k_idx) {
1445
+ const ulong k_row_off = batch_idx * k_nb3 + head_kv_idx * k_nb2 + k_idx * k_nb1;
1446
+ const ulong v_row_off = batch_idx * v_nb3 + head_kv_idx * v_nb2 + k_idx * v_nb1;
1447
+ const global KV_DATA_TYPE4 * k_ptr = (const global KV_DATA_TYPE4 *) (k_base + k_row_off);
1448
+ const global KV_DATA_TYPE4 * v_ptr = (const global KV_DATA_TYPE4 *) (v_base + v_row_off);
1449
+
1450
+ // Q*K^T: load each K stride once, dot against all MQ_GQA Q rows.
1451
+ ACC_TYPE4 dot4[MQ_GQA];
1452
+ #pragma unroll
1453
+ for (int h = 0; h < MQ_GQA; ++h) dot4[h] = (ACC_TYPE4)(0.0f);
1454
+ for (int k = tid_sg; k < DK_VEC; k += Q1_WG_SIZE) {
1455
+ const ACC_TYPE4 k_vec = CONVERT_KV_ACC4(k_ptr[k]);
1456
+ #pragma unroll
1457
+ for (int h = 0; h < MQ_GQA; ++h) {
1458
+ dot4[h] = mad(q_shared[h * DK_VEC + k], k_vec, dot4[h]);
1459
+ }
1460
+ }
1461
+
1462
+ ACC_TYPE score[MQ_GQA];
1463
+ #pragma unroll
1464
+ for (int h = 0; h < MQ_GQA; ++h) {
1465
+ const ACC_TYPE dot_partial = dot4[h].s0 + dot4[h].s1 + dot4[h].s2 + dot4[h].s3;
1466
+ ACC_TYPE s = sub_group_reduce_add(dot_partial) * scale;
1467
+ if (mask_base[h] != NULL) {
1468
+ const global MASK_DATA_TYPE * mask_ptr = (const global MASK_DATA_TYPE *) mask_base[h];
1469
+ s += slope[h] * (ACC_TYPE) mask_ptr[k_idx];
1470
+ }
1471
+ if (logit_softcap > 0.0f) {
1472
+ s = logit_softcap * tanh(s / logit_softcap);
1473
+ }
1474
+ score[h] = s;
1475
+ }
1476
+
1477
+ // FA-2 online softmax update — V load amortized across MQ_GQA heads.
1478
+ // p, scale_prev are computed per h; the V vector is loaded once
1479
+ // per dv stride and reused MQ_GQA times.
1480
+ ACC_TYPE p_h[MQ_GQA];
1481
+ ACC_TYPE sp_h[MQ_GQA];
1482
+ #pragma unroll
1483
+ for (int h = 0; h < MQ_GQA; ++h) {
1484
+ const ACC_TYPE m_new = max(m_i[h], score[h]);
1485
+ sp_h[h] = native_exp(m_i[h] - m_new);
1486
+ p_h[h] = native_exp(score[h] - m_new);
1487
+ l_i[h] = l_i[h] * sp_h[h] + p_h[h];
1488
+ m_i[h] = m_new;
1489
+ }
1490
+
1491
+ int idx = 0;
1492
+ for (int dv_idx = tid_sg; dv_idx < DV_VEC; dv_idx += Q1_WG_SIZE, ++idx) {
1493
+ const ACC_TYPE4 v_vec = CONVERT_KV_ACC4(v_ptr[dv_idx]);
1494
+ #pragma unroll
1495
+ for (int h = 0; h < MQ_GQA; ++h) {
1496
+ o_acc[h][idx] = mad(p_h[h], v_vec, o_acc[h][idx] * sp_h[h]);
1497
+ }
1498
+ }
1499
+ }
1500
+
1501
+ // cross subgroup merge
1502
+ __local ACC_TYPE sg_m[MQ_GQA][MQ_NSG];
1503
+ __local ACC_TYPE sg_l[MQ_GQA][MQ_NSG];
1504
+ __local ACC_TYPE4 sg_o[MQ_NSG][DV_VEC];
1505
+
1506
+ if (tid_sg == 0) {
1507
+ #pragma unroll
1508
+ for (int h = 0; h < MQ_GQA; ++h) {
1509
+ sg_m[h][sgid] = m_i[h];
1510
+ sg_l[h][sgid] = l_i[h];
1511
+ }
1512
+ }
1513
+
1514
+ #pragma unroll
1515
+ for (int h = 0; h < MQ_GQA; ++h) {
1516
+ // each subgroup publishes its o_acc slice for head h.
1517
+ {
1518
+ int idx = 0;
1519
+ for (int dv_idx = tid_sg; dv_idx < DV_VEC; dv_idx += Q1_WG_SIZE, ++idx) {
1520
+ sg_o[sgid][dv_idx] = o_acc[h][idx];
1521
+ }
1522
+ }
1523
+ barrier(CLK_LOCAL_MEM_FENCE);
1524
+
1525
+ if (sgid == 0) {
1526
+ const int head_idx = head_kv_idx * MQ_GQA + h;
1527
+
1528
+ ACC_TYPE m_final = sg_m[h][0];
1529
+ #pragma unroll
1530
+ for (int s = 1; s < MQ_NSG; ++s) {
1531
+ m_final = max(m_final, sg_m[h][s]);
1532
+ }
1533
+ if (sinks_ptr != NULL) {
1534
+ m_final = max(m_final, sinks_ptr[head_idx]);
1535
+ }
1536
+
1537
+ ACC_TYPE l_final = 0.0f;
1538
+ #pragma unroll
1539
+ for (int s = 0; s < MQ_NSG; ++s) {
1540
+ l_final += sg_l[h][s] * native_exp(sg_m[h][s] - m_final);
1541
+ }
1542
+ if (sinks_ptr != NULL) {
1543
+ l_final += native_exp(sinks_ptr[head_idx] - m_final);
1544
+ }
1545
+ const ACC_TYPE l_inv = (l_final > 0.0f) ? (1.0f / l_final) : 0.0f;
1546
+
1547
+ const ulong o_row_offset = batch_idx * o_nb3 + head_idx * o_nb1;
1548
+ global O_DATA_TYPE4 * o_row = (global O_DATA_TYPE4 *) (o_base + o_row_offset);
1549
+ for (int dv_idx = tid_sg; dv_idx < DV_VEC; dv_idx += Q1_WG_SIZE) {
1550
+ ACC_TYPE4 o_merged = (ACC_TYPE4)(0.0f);
1551
+ #pragma unroll
1552
+ for (int s = 0; s < MQ_NSG; ++s) {
1553
+ const ACC_TYPE alpha = native_exp(sg_m[h][s] - m_final);
1554
+ o_merged = mad((ACC_TYPE4)(alpha), sg_o[s][dv_idx], o_merged);
1555
+ }
1556
+ o_row[dv_idx] = CONVERT_O_DATA4(o_merged * l_inv);
1557
+ }
1558
+ }
1559
+ // Barrier guards next h's overwrite of sg_o.
1560
+ barrier(CLK_LOCAL_MEM_FENCE);
1561
+ }
1562
+ }
1563
+
1564
+ #ifndef MQ_NSG_SPLIT
1565
+ #define MQ_NSG_SPLIT 4
1566
+ #endif
1567
+ #define MQ_SPLIT_WG_SIZE (Q1_WG_SIZE * MQ_NSG_SPLIT)
1568
+
1569
+ #ifndef FA_PARTIAL_FLOATS
1570
+ #define FA_PARTIAL_FLOATS (2 + DV)
1571
+ #endif
1572
+
1573
+ REQD_SUBGROUP_SIZE_64
1574
+ __kernel void flash_attn_f32_f16_q1_vec_mq_split(
1575
+ const global void * q_void, ulong q_offset,
1576
+ const global void * k_void, ulong k_offset,
1577
+ const global void * v_void, ulong v_offset,
1578
+ const float scale,
1579
+ const int n_q,
1580
+ const int n_kv,
1581
+ const int n_head,
1582
+ const ulong q_nb1, const ulong q_nb2, const ulong q_nb3,
1583
+ const ulong k_nb1, const ulong k_nb2, const ulong k_nb3,
1584
+ const ulong v_nb1, const ulong v_nb2, const ulong v_nb3,
1585
+ const float max_bias,
1586
+ const float m0,
1587
+ const float m1,
1588
+ const int n_head_log2,
1589
+ const float logit_softcap,
1590
+ const int n_head_kv,
1591
+ const global void * mask_void,
1592
+ const ulong mask_offset,
1593
+ const ulong mask_nb1,
1594
+ const ulong mask_nb2,
1595
+ const ulong mask_nb3,
1596
+ const int mask_ne2,
1597
+ const int mask_ne3,
1598
+ global float * partial_void,
1599
+ const int n_splits,
1600
+ const int kv_per_split
1601
+ ) {
1602
+ const int tid = get_local_id(0);
1603
+ const int sgid = tid / Q1_WG_SIZE;
1604
+ const int tid_sg = tid % Q1_WG_SIZE;
1605
+ const int kvhead_batch_idx = get_global_id(1);
1606
+ const int split_q_idx = get_global_id(2);
1607
+ const int split_idx = split_q_idx % n_splits;
1608
+ const int q_idx = split_q_idx / n_splits;
1609
+
1610
+ const int batch_idx = kvhead_batch_idx / n_head_kv;
1611
+ const int head_kv_idx = kvhead_batch_idx % n_head_kv;
1612
+
1613
+ const int kv_start = split_idx * kv_per_split;
1614
+ const int kv_end = min(kv_start + kv_per_split, n_kv);
1615
+
1616
+ const ulong record_stride = (ulong) FA_PARTIAL_FLOATS;
1617
+
1618
+ if (kv_start >= kv_end) {
1619
+ // write sentinel for each of the MQ_GQA Q-heads so the
1620
+ // merge pass treats this slot as dropped
1621
+ if (tid == 0) {
1622
+ #pragma unroll
1623
+ for (int h = 0; h < MQ_GQA; ++h) {
1624
+ const int head_idx = head_kv_idx * MQ_GQA + h;
1625
+ const ulong rec_idx = ((((ulong) batch_idx * n_head + head_idx) * n_q + q_idx)
1626
+ * n_splits + split_idx);
1627
+ global float * rec = partial_void + rec_idx * record_stride;
1628
+ rec[0] = FA_M_INIT;
1629
+ rec[1] = 0.0f;
1630
+ }
1631
+ }
1632
+ return;
1633
+ }
1634
+
1635
+ const global char * q_base = (const global char *) q_void + q_offset;
1636
+ const global char * k_base = (const global char *) k_void + k_offset;
1637
+ const global char * v_base = (const global char *) v_void + v_offset;
1638
+
1639
+ // stage MQ_GQA Q rows in __local once (uniform across WG)
1640
+ __local ACC_TYPE4 q_shared[MQ_GQA * DK_VEC];
1641
+ for (int i = tid; i < MQ_GQA * DK_VEC; i += MQ_SPLIT_WG_SIZE) {
1642
+ const int h = i / DK_VEC;
1643
+ const int k = i % DK_VEC;
1644
+ const int head_idx = head_kv_idx * MQ_GQA + h;
1645
+ const ulong q_row_offset = batch_idx * q_nb3 + head_idx * q_nb2 + (ulong) q_idx * q_nb1;
1646
+ const global Q_DATA_TYPE4 * q_ptr = (const global Q_DATA_TYPE4 *) (q_base + q_row_offset);
1647
+ q_shared[h * DK_VEC + k] = CONVERT_Q_ACC4(q_ptr[k]);
1648
+ }
1649
+ barrier(CLK_LOCAL_MEM_FENCE);
1650
+
1651
+ float slope[MQ_GQA];
1652
+ #pragma unroll
1653
+ for (int h = 0; h < MQ_GQA; ++h) {
1654
+ slope[h] = get_alibi_slope(max_bias, head_kv_idx * MQ_GQA + h, n_head_log2, m0, m1);
1655
+ }
1656
+
1657
+ const global char * mask_base[MQ_GQA];
1658
+ if (mask_void != NULL) {
1659
+ const int mask_batch_idx = batch_idx % mask_ne3;
1660
+ const global char * mask_base_b = (const global char *) mask_void + mask_offset +
1661
+ mask_batch_idx * mask_nb3 +
1662
+ (ulong) q_idx * mask_nb1;
1663
+ #pragma unroll
1664
+ for (int h = 0; h < MQ_GQA; ++h) {
1665
+ const int head_idx = head_kv_idx * MQ_GQA + h;
1666
+ const int mask_head_idx = head_idx % mask_ne2;
1667
+ mask_base[h] = mask_base_b + mask_head_idx * mask_nb2;
1668
+ }
1669
+ } else {
1670
+ #pragma unroll
1671
+ for (int h = 0; h < MQ_GQA; ++h) mask_base[h] = NULL;
1672
+ }
1673
+
1674
+ ACC_TYPE4 o_acc[MQ_GQA][Q1V_DV_PER_THREAD];
1675
+ ACC_TYPE m_i[MQ_GQA];
1676
+ ACC_TYPE l_i[MQ_GQA];
1677
+ #pragma unroll
1678
+ for (int h = 0; h < MQ_GQA; ++h) {
1679
+ m_i[h] = FA_M_INIT;
1680
+ l_i[h] = 0.0f;
1681
+ #pragma unroll
1682
+ for (int i = 0; i < Q1V_DV_PER_THREAD; ++i) o_acc[h][i] = (ACC_TYPE4)(0.0f);
1683
+ }
1684
+
1685
+ // each subgroup independently sweeps its slice of the split's kv range.
1686
+ const int kv_len = kv_end - kv_start;
1687
+ const int kv_per_sg = (kv_len + MQ_NSG_SPLIT - 1) / MQ_NSG_SPLIT;
1688
+ const int kv_lo = kv_start + sgid * kv_per_sg;
1689
+ const int kv_hi = min(kv_end, kv_lo + kv_per_sg);
1690
+
1691
+ for (int k_idx = kv_lo; k_idx < kv_hi; ++k_idx) {
1692
+ const ulong k_row_off = batch_idx * k_nb3 + head_kv_idx * k_nb2 + k_idx * k_nb1;
1693
+ const ulong v_row_off = batch_idx * v_nb3 + head_kv_idx * v_nb2 + k_idx * v_nb1;
1694
+ const global KV_DATA_TYPE4 * k_ptr = (const global KV_DATA_TYPE4 *) (k_base + k_row_off);
1695
+ const global KV_DATA_TYPE4 * v_ptr = (const global KV_DATA_TYPE4 *) (v_base + v_row_off);
1696
+
1697
+ ACC_TYPE4 dot4[MQ_GQA];
1698
+ #pragma unroll
1699
+ for (int h = 0; h < MQ_GQA; ++h) dot4[h] = (ACC_TYPE4)(0.0f);
1700
+ for (int k = tid_sg; k < DK_VEC; k += Q1_WG_SIZE) {
1701
+ const ACC_TYPE4 k_vec = CONVERT_KV_ACC4(k_ptr[k]);
1702
+ #pragma unroll
1703
+ for (int h = 0; h < MQ_GQA; ++h) {
1704
+ dot4[h] = mad(q_shared[h * DK_VEC + k], k_vec, dot4[h]);
1705
+ }
1706
+ }
1707
+
1708
+ ACC_TYPE score[MQ_GQA];
1709
+ #pragma unroll
1710
+ for (int h = 0; h < MQ_GQA; ++h) {
1711
+ const ACC_TYPE dot_partial = dot4[h].s0 + dot4[h].s1 + dot4[h].s2 + dot4[h].s3;
1712
+ ACC_TYPE s = sub_group_reduce_add(dot_partial) * scale;
1713
+ if (mask_base[h] != NULL) {
1714
+ const global MASK_DATA_TYPE * mask_ptr = (const global MASK_DATA_TYPE *) mask_base[h];
1715
+ s += slope[h] * (ACC_TYPE) mask_ptr[k_idx];
1716
+ }
1717
+ if (logit_softcap > 0.0f) {
1718
+ s = logit_softcap * tanh(s / logit_softcap);
1719
+ }
1720
+ score[h] = s;
1721
+ }
1722
+
1723
+ ACC_TYPE p_h[MQ_GQA];
1724
+ ACC_TYPE sp_h[MQ_GQA];
1725
+ #pragma unroll
1726
+ for (int h = 0; h < MQ_GQA; ++h) {
1727
+ const ACC_TYPE m_new = max(m_i[h], score[h]);
1728
+ sp_h[h] = native_exp(m_i[h] - m_new);
1729
+ p_h[h] = native_exp(score[h] - m_new);
1730
+ l_i[h] = l_i[h] * sp_h[h] + p_h[h];
1731
+ m_i[h] = m_new;
1732
+ }
1733
+
1734
+ int idx = 0;
1735
+ for (int dv_idx = tid_sg; dv_idx < DV_VEC; dv_idx += Q1_WG_SIZE, ++idx) {
1736
+ const ACC_TYPE4 v_vec = CONVERT_KV_ACC4(v_ptr[dv_idx]);
1737
+ #pragma unroll
1738
+ for (int h = 0; h < MQ_GQA; ++h) {
1739
+ o_acc[h][idx] = mad(p_h[h], v_vec, o_acc[h][idx] * sp_h[h]);
1740
+ }
1741
+ }
1742
+ }
1743
+
1744
+ // per-h cross-subgroup merge
1745
+ __local ACC_TYPE sg_m[MQ_GQA][MQ_NSG_SPLIT];
1746
+ __local ACC_TYPE sg_l[MQ_GQA][MQ_NSG_SPLIT];
1747
+ __local ACC_TYPE4 sg_o[MQ_NSG_SPLIT][DV_VEC];
1748
+
1749
+ if (tid_sg == 0) {
1750
+ #pragma unroll
1751
+ for (int h = 0; h < MQ_GQA; ++h) {
1752
+ sg_m[h][sgid] = m_i[h];
1753
+ sg_l[h][sgid] = l_i[h];
1754
+ }
1755
+ }
1756
+
1757
+ #pragma unroll
1758
+ for (int h = 0; h < MQ_GQA; ++h) {
1759
+ {
1760
+ int idx = 0;
1761
+ for (int dv_idx = tid_sg; dv_idx < DV_VEC; dv_idx += Q1_WG_SIZE, ++idx) {
1762
+ sg_o[sgid][dv_idx] = o_acc[h][idx];
1763
+ }
1764
+ }
1765
+ barrier(CLK_LOCAL_MEM_FENCE);
1766
+
1767
+ if (sgid == 0) {
1768
+ const int head_idx = head_kv_idx * MQ_GQA + h;
1769
+
1770
+ // fold per-subgroup (m, l) into split-level (m_c, l_c)
1771
+ ACC_TYPE m_c = sg_m[h][0];
1772
+ #pragma unroll
1773
+ for (int s = 1; s < MQ_NSG_SPLIT; ++s) {
1774
+ m_c = max(m_c, sg_m[h][s]);
1775
+ }
1776
+ ACC_TYPE l_c = 0.0f;
1777
+ #pragma unroll
1778
+ for (int s = 0; s < MQ_NSG_SPLIT; ++s) {
1779
+ l_c += sg_l[h][s] * native_exp(sg_m[h][s] - m_c);
1780
+ }
1781
+
1782
+ const ulong rec_idx = ((((ulong) batch_idx * n_head + head_idx) * n_q + q_idx)
1783
+ * n_splits + split_idx);
1784
+ global float * rec = partial_void + rec_idx * record_stride;
1785
+ global float4 * rec_o = (global float4 *) (rec + 2);
1786
+
1787
+ if (tid_sg == 0) {
1788
+ rec[0] = (float) m_c;
1789
+ rec[1] = (float) l_c;
1790
+ }
1791
+ // each thread writes its DV slice of the merged O.
1792
+ for (int dv_idx = tid_sg; dv_idx < DV_VEC; dv_idx += Q1_WG_SIZE) {
1793
+ ACC_TYPE4 o_merged = (ACC_TYPE4)(0.0f);
1794
+ #pragma unroll
1795
+ for (int s = 0; s < MQ_NSG_SPLIT; ++s) {
1796
+ const ACC_TYPE alpha = native_exp(sg_m[h][s] - m_c);
1797
+ o_merged = mad((ACC_TYPE4)(alpha), sg_o[s][dv_idx], o_merged);
1798
+ }
1799
+ rec_o[dv_idx] = o_merged;
1800
+ }
1801
+ }
1802
+ barrier(CLK_LOCAL_MEM_FENCE);
1803
+ }
1804
+ }
1805
+
1806
+ REQD_SUBGROUP_SIZE_64
1807
+ __kernel void flash_attn_f32_f16_q1_vec_mq_split_k_img(
1808
+ const global void * q_void, ulong q_offset,
1809
+ __read_only image1d_buffer_t k_img,
1810
+ const global void * v_void, ulong v_offset,
1811
+ const float scale,
1812
+ const int n_q,
1813
+ const int n_kv,
1814
+ const int n_head,
1815
+ const ulong q_nb1, const ulong q_nb2, const ulong q_nb3,
1816
+ const ulong k_nb1, const ulong k_nb2, const ulong k_nb3,
1817
+ const ulong v_nb1, const ulong v_nb2, const ulong v_nb3,
1818
+ const float max_bias,
1819
+ const float m0,
1820
+ const float m1,
1821
+ const int n_head_log2,
1822
+ const float logit_softcap,
1823
+ const int n_head_kv,
1824
+ const global void * mask_void,
1825
+ const ulong mask_offset,
1826
+ const ulong mask_nb1,
1827
+ const ulong mask_nb2,
1828
+ const ulong mask_nb3,
1829
+ const int mask_ne2,
1830
+ const int mask_ne3,
1831
+ global float * partial_void,
1832
+ const int n_splits,
1833
+ const int kv_per_split
1834
+ ) {
1835
+ const int tid = get_local_id(0);
1836
+ const int sgid = tid / Q1_WG_SIZE;
1837
+ const int tid_sg = tid % Q1_WG_SIZE;
1838
+ const int kvhead_batch_idx = get_global_id(1);
1839
+ const int split_q_idx = get_global_id(2);
1840
+ const int split_idx = split_q_idx % n_splits;
1841
+ const int q_idx = split_q_idx / n_splits;
1842
+
1843
+ const int batch_idx = kvhead_batch_idx / n_head_kv;
1844
+ const int head_kv_idx = kvhead_batch_idx % n_head_kv;
1845
+
1846
+ const int kv_start = split_idx * kv_per_split;
1847
+ const int kv_end = min(kv_start + kv_per_split, n_kv);
1848
+
1849
+ const ulong record_stride = (ulong) FA_PARTIAL_FLOATS;
1850
+
1851
+ if (kv_start >= kv_end) {
1852
+ if (tid == 0) {
1853
+ #pragma unroll
1854
+ for (int h = 0; h < MQ_GQA; ++h) {
1855
+ const int head_idx = head_kv_idx * MQ_GQA + h;
1856
+ const ulong rec_idx = ((((ulong) batch_idx * n_head + head_idx) * n_q + q_idx)
1857
+ * n_splits + split_idx);
1858
+ global float * rec = partial_void + rec_idx * record_stride;
1859
+ rec[0] = FA_M_INIT;
1860
+ rec[1] = 0.0f;
1861
+ }
1862
+ }
1863
+ return;
1864
+ }
1865
+
1866
+ const global char * q_base = (const global char *) q_void + q_offset;
1867
+ const global char * v_base = (const global char *) v_void + v_offset;
1868
+
1869
+ __local ACC_TYPE4 q_shared[MQ_GQA * DK_VEC];
1870
+ for (int i = tid; i < MQ_GQA * DK_VEC; i += MQ_SPLIT_WG_SIZE) {
1871
+ const int h = i / DK_VEC;
1872
+ const int k = i % DK_VEC;
1873
+ const int head_idx = head_kv_idx * MQ_GQA + h;
1874
+ const ulong q_row_offset = batch_idx * q_nb3 + head_idx * q_nb2 + (ulong) q_idx * q_nb1;
1875
+ const global Q_DATA_TYPE4 * q_ptr = (const global Q_DATA_TYPE4 *) (q_base + q_row_offset);
1876
+ q_shared[h * DK_VEC + k] = CONVERT_Q_ACC4(q_ptr[k]);
1877
+ }
1878
+ barrier(CLK_LOCAL_MEM_FENCE);
1879
+
1880
+ float slope[MQ_GQA];
1881
+ #pragma unroll
1882
+ for (int h = 0; h < MQ_GQA; ++h) {
1883
+ slope[h] = get_alibi_slope(max_bias, head_kv_idx * MQ_GQA + h, n_head_log2, m0, m1);
1884
+ }
1885
+
1886
+ const global char * mask_base[MQ_GQA];
1887
+ if (mask_void != NULL) {
1888
+ const int mask_batch_idx = batch_idx % mask_ne3;
1889
+ const global char * mask_base_b = (const global char *) mask_void + mask_offset +
1890
+ mask_batch_idx * mask_nb3 +
1891
+ (ulong) q_idx * mask_nb1;
1892
+ #pragma unroll
1893
+ for (int h = 0; h < MQ_GQA; ++h) {
1894
+ const int head_idx = head_kv_idx * MQ_GQA + h;
1895
+ const int mask_head_idx = head_idx % mask_ne2;
1896
+ mask_base[h] = mask_base_b + mask_head_idx * mask_nb2;
1897
+ }
1898
+ } else {
1899
+ #pragma unroll
1900
+ for (int h = 0; h < MQ_GQA; ++h) mask_base[h] = NULL;
1901
+ }
1902
+
1903
+ ACC_TYPE4 o_acc[MQ_GQA][Q1V_DV_PER_THREAD];
1904
+ ACC_TYPE m_i[MQ_GQA];
1905
+ ACC_TYPE l_i[MQ_GQA];
1906
+ #pragma unroll
1907
+ for (int h = 0; h < MQ_GQA; ++h) {
1908
+ m_i[h] = FA_M_INIT;
1909
+ l_i[h] = 0.0f;
1910
+ #pragma unroll
1911
+ for (int i = 0; i < Q1V_DV_PER_THREAD; ++i) o_acc[h][i] = (ACC_TYPE4)(0.0f);
1912
+ }
1913
+
1914
+ // K pitches in pixel units, pixel = 1 half4 = 8 B -> byte_stride >> 3.
1915
+ const int pitch_px_row = (int)(k_nb1 >> 3);
1916
+ const int pitch_px_head = (int)(k_nb2 >> 3);
1917
+ const int pitch_px_batch = (int)(k_nb3 >> 3);
1918
+
1919
+ const int kv_len = kv_end - kv_start;
1920
+ const int kv_per_sg = (kv_len + MQ_NSG_SPLIT - 1) / MQ_NSG_SPLIT;
1921
+ const int kv_lo = kv_start + sgid * kv_per_sg;
1922
+ const int kv_hi = min(kv_end, kv_lo + kv_per_sg);
1923
+
1924
+ for (int k_idx = kv_lo; k_idx < kv_hi; ++k_idx) {
1925
+ const int k_row_px = batch_idx * pitch_px_batch +
1926
+ head_kv_idx * pitch_px_head +
1927
+ k_idx * pitch_px_row;
1928
+
1929
+ const ulong v_row_off = batch_idx * v_nb3 + head_kv_idx * v_nb2 + k_idx * v_nb1;
1930
+ const global KV_DATA_TYPE4 * v_ptr = (const global KV_DATA_TYPE4 *) (v_base + v_row_off);
1931
+
1932
+ ACC_TYPE4 dot4[MQ_GQA];
1933
+ #pragma unroll
1934
+ for (int h = 0; h < MQ_GQA; ++h) dot4[h] = (ACC_TYPE4)(0.0f);
1935
+ for (int k = tid_sg; k < DK_VEC; k += Q1_WG_SIZE) {
1936
+ const half4 k_h4 = read_imageh(k_img, k_row_px + k);
1937
+ const ACC_TYPE4 k_vec = CONVERT_KV_ACC4(k_h4);
1938
+ #pragma unroll
1939
+ for (int h = 0; h < MQ_GQA; ++h) {
1940
+ dot4[h] = mad(q_shared[h * DK_VEC + k], k_vec, dot4[h]);
1941
+ }
1942
+ }
1943
+
1944
+ ACC_TYPE score[MQ_GQA];
1945
+ #pragma unroll
1946
+ for (int h = 0; h < MQ_GQA; ++h) {
1947
+ const ACC_TYPE dot_partial = dot4[h].s0 + dot4[h].s1 + dot4[h].s2 + dot4[h].s3;
1948
+ ACC_TYPE s = sub_group_reduce_add(dot_partial) * scale;
1949
+ if (mask_base[h] != NULL) {
1950
+ const global MASK_DATA_TYPE * mask_ptr = (const global MASK_DATA_TYPE *) mask_base[h];
1951
+ s += slope[h] * (ACC_TYPE) mask_ptr[k_idx];
1952
+ }
1953
+ if (logit_softcap > 0.0f) {
1954
+ s = logit_softcap * tanh(s / logit_softcap);
1955
+ }
1956
+ score[h] = s;
1957
+ }
1958
+
1959
+ ACC_TYPE p_h[MQ_GQA];
1960
+ ACC_TYPE sp_h[MQ_GQA];
1961
+ #pragma unroll
1962
+ for (int h = 0; h < MQ_GQA; ++h) {
1963
+ const ACC_TYPE m_new = max(m_i[h], score[h]);
1964
+ sp_h[h] = native_exp(m_i[h] - m_new);
1965
+ p_h[h] = native_exp(score[h] - m_new);
1966
+ l_i[h] = l_i[h] * sp_h[h] + p_h[h];
1967
+ m_i[h] = m_new;
1968
+ }
1969
+
1970
+ int idx = 0;
1971
+ for (int dv_idx = tid_sg; dv_idx < DV_VEC; dv_idx += Q1_WG_SIZE, ++idx) {
1972
+ const ACC_TYPE4 v_vec = CONVERT_KV_ACC4(v_ptr[dv_idx]);
1973
+ #pragma unroll
1974
+ for (int h = 0; h < MQ_GQA; ++h) {
1975
+ o_acc[h][idx] = mad(p_h[h], v_vec, o_acc[h][idx] * sp_h[h]);
1976
+ }
1977
+ }
1978
+ }
1979
+
1980
+ __local ACC_TYPE sg_m[MQ_GQA][MQ_NSG_SPLIT];
1981
+ __local ACC_TYPE sg_l[MQ_GQA][MQ_NSG_SPLIT];
1982
+ __local ACC_TYPE4 sg_o[MQ_NSG_SPLIT][DV_VEC];
1983
+
1984
+ if (tid_sg == 0) {
1985
+ #pragma unroll
1986
+ for (int h = 0; h < MQ_GQA; ++h) {
1987
+ sg_m[h][sgid] = m_i[h];
1988
+ sg_l[h][sgid] = l_i[h];
1989
+ }
1990
+ }
1991
+
1992
+ #pragma unroll
1993
+ for (int h = 0; h < MQ_GQA; ++h) {
1994
+ {
1995
+ int idx = 0;
1996
+ for (int dv_idx = tid_sg; dv_idx < DV_VEC; dv_idx += Q1_WG_SIZE, ++idx) {
1997
+ sg_o[sgid][dv_idx] = o_acc[h][idx];
1998
+ }
1999
+ }
2000
+ barrier(CLK_LOCAL_MEM_FENCE);
2001
+
2002
+ if (sgid == 0) {
2003
+ const int head_idx = head_kv_idx * MQ_GQA + h;
2004
+
2005
+ ACC_TYPE m_c = sg_m[h][0];
2006
+ #pragma unroll
2007
+ for (int s = 1; s < MQ_NSG_SPLIT; ++s) {
2008
+ m_c = max(m_c, sg_m[h][s]);
2009
+ }
2010
+ ACC_TYPE l_c = 0.0f;
2011
+ #pragma unroll
2012
+ for (int s = 0; s < MQ_NSG_SPLIT; ++s) {
2013
+ l_c += sg_l[h][s] * native_exp(sg_m[h][s] - m_c);
2014
+ }
2015
+
2016
+ const ulong rec_idx = ((((ulong) batch_idx * n_head + head_idx) * n_q + q_idx)
2017
+ * n_splits + split_idx);
2018
+ global float * rec = partial_void + rec_idx * record_stride;
2019
+ global float4 * rec_o = (global float4 *) (rec + 2);
2020
+
2021
+ if (tid_sg == 0) {
2022
+ rec[0] = (float) m_c;
2023
+ rec[1] = (float) l_c;
2024
+ }
2025
+ for (int dv_idx = tid_sg; dv_idx < DV_VEC; dv_idx += Q1_WG_SIZE) {
2026
+ ACC_TYPE4 o_merged = (ACC_TYPE4)(0.0f);
2027
+ #pragma unroll
2028
+ for (int s = 0; s < MQ_NSG_SPLIT; ++s) {
2029
+ const ACC_TYPE alpha = native_exp(sg_m[h][s] - m_c);
2030
+ o_merged = mad((ACC_TYPE4)(alpha), sg_o[s][dv_idx], o_merged);
2031
+ }
2032
+ rec_o[dv_idx] = o_merged;
2033
+ }
2034
+ }
2035
+ barrier(CLK_LOCAL_MEM_FENCE);
2036
+ }
2037
+ }
2038
+ #endif // !FA_DECODE_ONLY
2039
+
2040
+ #ifndef FA_MQ_ONLY // q1_split + merge excluded from the MQ-only (g8) program
2041
+ __kernel void flash_attn_f32_f16_q1_split(
2042
+ const global void * q_void, ulong q_offset,
2043
+ const global void * k_void, ulong k_offset,
2044
+ const global void * v_void, ulong v_offset,
2045
+ const float scale,
2046
+ const int n_q,
2047
+ const int n_kv,
2048
+ const int n_head,
2049
+ const ulong q_nb1, const ulong q_nb2, const ulong q_nb3,
2050
+ const ulong k_nb1, const ulong k_nb2, const ulong k_nb3,
2051
+ const ulong v_nb1, const ulong v_nb2, const ulong v_nb3,
2052
+ const float max_bias,
2053
+ const float m0,
2054
+ const float m1,
2055
+ const int n_head_log2,
2056
+ const float logit_softcap,
2057
+ const int n_head_kv,
2058
+ const global void * mask_void,
2059
+ const ulong mask_offset,
2060
+ const ulong mask_nb1,
2061
+ const ulong mask_nb2,
2062
+ const ulong mask_nb3,
2063
+ const int mask_ne2,
2064
+ const int mask_ne3,
2065
+ global float * partial_void,
2066
+ const int n_splits,
2067
+ const int kv_per_split
2068
+ ) {
2069
+ const int tid = get_local_id(0);
2070
+ const int head_batch_idx = get_global_id(1);
2071
+ const int split_q_idx = get_global_id(2);
2072
+ const int split_idx = split_q_idx % n_splits;
2073
+ const int q_idx = split_q_idx / n_splits;
2074
+ const int batch_idx = head_batch_idx / n_head;
2075
+ const int head_idx = head_batch_idx % n_head;
2076
+ const int gqa_ratio = n_head / n_head_kv;
2077
+ const int head_kv_idx = head_idx / gqa_ratio;
2078
+
2079
+ const int kv_start = split_idx * kv_per_split;
2080
+ const int kv_end = min(kv_start + kv_per_split, n_kv);
2081
+
2082
+ const ulong record_stride = (ulong) FA_PARTIAL_FLOATS;
2083
+ const ulong record_idx = ((((ulong) batch_idx * n_head + head_idx) * n_q + q_idx)
2084
+ * n_splits + split_idx);
2085
+ global float * rec = partial_void + record_idx * record_stride;
2086
+ global float4 * rec_o = (global float4 *) (rec + 2);
2087
+
2088
+ if (kv_start >= kv_end) {
2089
+ // Empty split: leave sentinel partial for merge.
2090
+ if (tid == 0) {
2091
+ rec[0] = FA_M_INIT;
2092
+ rec[1] = 0.0f;
2093
+ }
2094
+ return;
2095
+ }
2096
+
2097
+ const global char * q_base = (const global char *) q_void + q_offset;
2098
+ const global char * k_base = (const global char *) k_void + k_offset;
2099
+ const global char * v_base = (const global char *) v_void + v_offset;
2100
+
2101
+ const global char * mask_base = NULL;
2102
+ if (mask_void != NULL) {
2103
+ const int mask_head_idx = head_idx % mask_ne2;
2104
+ const int mask_batch_idx = batch_idx % mask_ne3;
2105
+ mask_base = (const global char *) mask_void + mask_offset +
2106
+ mask_batch_idx * mask_nb3 + mask_head_idx * mask_nb2 +
2107
+ (ulong) q_idx * mask_nb1;
2108
+ }
2109
+
2110
+ // share Q via local memory (n_q=1 per split -> uniform across WG).
2111
+ __local ACC_TYPE4 q_shared[DK_VEC];
2112
+ const ulong q_row_offset = batch_idx * q_nb3 + head_idx * q_nb2 + (ulong) q_idx * q_nb1;
2113
+ const global Q_DATA_TYPE4 * q_ptr = (const global Q_DATA_TYPE4 *) (q_base + q_row_offset);
2114
+ for (int i = tid; i < DK_VEC; i += Q1_WG_SIZE) {
2115
+ q_shared[i] = CONVERT_Q_ACC4(q_ptr[i]);
2116
+ }
2117
+ sub_group_barrier(CLK_LOCAL_MEM_FENCE);
2118
+
2119
+ const float slope = get_alibi_slope(max_bias, head_idx, n_head_log2, m0, m1);
2120
+
2121
+ // pass 1a — split-local max.
2122
+ ACC_TYPE m_i = FA_M_INIT;
2123
+ for (int k_idx = kv_start + tid; k_idx < kv_end; k_idx += Q1_WG_SIZE) {
2124
+ const ulong k_row_offset = batch_idx * k_nb3 + head_kv_idx * k_nb2 + k_idx * k_nb1;
2125
+ const global KV_DATA_TYPE4 * k_ptr = (const global KV_DATA_TYPE4 *) (k_base + k_row_offset);
2126
+ ACC_TYPE4 dot_acc = (ACC_TYPE4)(0.0f);
2127
+ #pragma unroll
2128
+ for (int k = 0; k < DK_VEC; ++k) {
2129
+ dot_acc = mad(q_shared[k], CONVERT_KV_ACC4(k_ptr[k]), dot_acc);
2130
+ }
2131
+ ACC_TYPE score = (dot_acc.s0 + dot_acc.s1 + dot_acc.s2 + dot_acc.s3) * scale;
2132
+ if (mask_base != NULL) {
2133
+ const global MASK_DATA_TYPE * mask_ptr = (const global MASK_DATA_TYPE *) (mask_base);
2134
+ score += slope * (ACC_TYPE) mask_ptr[k_idx];
2135
+ }
2136
+ if (logit_softcap > 0.0f) {
2137
+ score = logit_softcap * tanh(score / logit_softcap);
2138
+ }
2139
+ m_i = max(m_i, score);
2140
+ }
2141
+
2142
+ const ACC_TYPE m_c = sub_group_reduce_max(m_i);
2143
+
2144
+ // pass 1b — softmax-weighted V accumulate.
2145
+ ACC_TYPE4 o_acc[DV_VEC];
2146
+ #pragma unroll
2147
+ for (int i = 0; i < DV_VEC; ++i) o_acc[i] = (ACC_TYPE4)(0.0f);
2148
+ ACC_TYPE l_i = 0.0f;
2149
+
2150
+ for (int k_idx = kv_start + tid; k_idx < kv_end; k_idx += Q1_WG_SIZE) {
2151
+ const ulong k_row_offset = batch_idx * k_nb3 + head_kv_idx * k_nb2 + k_idx * k_nb1;
2152
+ const ulong v_row_offset = batch_idx * v_nb3 + head_kv_idx * v_nb2 + k_idx * v_nb1;
2153
+ const global KV_DATA_TYPE4 * k_ptr = (const global KV_DATA_TYPE4 *) (k_base + k_row_offset);
2154
+ const global KV_DATA_TYPE4 * v_ptr = (const global KV_DATA_TYPE4 *) (v_base + v_row_offset);
2155
+ ACC_TYPE4 dot_acc = (ACC_TYPE4)(0.0f);
2156
+ #pragma unroll
2157
+ for (int k = 0; k < DK_VEC; ++k) {
2158
+ dot_acc = mad(q_shared[k], CONVERT_KV_ACC4(k_ptr[k]), dot_acc);
2159
+ }
2160
+ ACC_TYPE score = (dot_acc.s0 + dot_acc.s1 + dot_acc.s2 + dot_acc.s3) * scale;
2161
+ if (mask_base != NULL) {
2162
+ const global MASK_DATA_TYPE * mask_ptr = (const global MASK_DATA_TYPE *) (mask_base);
2163
+ score += slope * (ACC_TYPE) mask_ptr[k_idx];
2164
+ }
2165
+ if (logit_softcap > 0.0f) {
2166
+ score = logit_softcap * tanh(score / logit_softcap);
2167
+ }
2168
+ const ACC_TYPE p = exp(score - m_c);
2169
+ l_i += p;
2170
+ #pragma unroll
2171
+ for (int i = 0; i < DV_VEC; ++i) {
2172
+ o_acc[i] = mad(p, CONVERT_KV_ACC4(v_ptr[i]), o_acc[i]);
2173
+ }
2174
+ }
2175
+
2176
+ __local ACC_TYPE4 local_o[Q1_WG_SIZE];
2177
+ const ACC_TYPE l_c = sub_group_reduce_add(l_i);
2178
+
2179
+ if (tid == 0) {
2180
+ rec[0] = (float) m_c;
2181
+ rec[1] = (float) l_c;
2182
+ }
2183
+ for (int i = 0; i < DV_VEC; ++i) {
2184
+ local_o[tid] = o_acc[i];
2185
+ sub_group_barrier(CLK_LOCAL_MEM_FENCE);
2186
+ #pragma unroll
2187
+ for (int s = Q1_WG_SIZE / 2; s > 0; s >>= 1) {
2188
+ if (tid < s) local_o[tid] += local_o[tid + s];
2189
+ sub_group_barrier(CLK_LOCAL_MEM_FENCE);
2190
+ }
2191
+ if (tid == 0) {
2192
+ rec_o[i] = local_o[0];
2193
+ }
2194
+ }
2195
+ }
2196
+
2197
+ // FD Pass 2: merge per-split partials into final O
2198
+ // empty splits drop via exp(-INF)=0.
2199
+ __kernel void flash_attn_f32_merge(
2200
+ const global float * partial_void,
2201
+ global void * o_void,
2202
+ const ulong o_offset,
2203
+ const int n_head,
2204
+ const int n_splits,
2205
+ const ulong o_nb1, const ulong o_nb2, const ulong o_nb3,
2206
+ const global void * sinks_void,
2207
+ const ulong sinks_offset,
2208
+ const int n_q
2209
+ ) {
2210
+ const int lane = get_local_id(0); // 0..DV_VEC-1
2211
+ const int head_batch_idx = get_global_id(1);
2212
+ const int q_idx = get_global_id(2);
2213
+ const int batch_idx = head_batch_idx / n_head;
2214
+ const int head_idx = head_batch_idx % n_head;
2215
+
2216
+ const ulong record_stride = (ulong) FA_PARTIAL_FLOATS;
2217
+ const ulong record_idx_0 = (((ulong) batch_idx * n_head + head_idx) * n_q + q_idx) * n_splits;
2218
+ const global float * rec0 = partial_void + record_idx_0 * record_stride;
2219
+
2220
+ __local ACC_TYPE m_final_shared;
2221
+ __local ACC_TYPE l_final_shared;
2222
+ if (lane == 0) {
2223
+ ACC_TYPE m = FA_M_INIT;
2224
+ for (int c = 0; c < n_splits; ++c) {
2225
+ const ACC_TYPE m_c = rec0[c * record_stride + 0];
2226
+ m = max(m, m_c);
2227
+ }
2228
+ ACC_TYPE m_sink = 0.0f;
2229
+ bool has_sink = false;
2230
+ if (sinks_void != NULL) {
2231
+ const global ACC_TYPE * sinks_ptr =
2232
+ (const global ACC_TYPE *) ((const global char *) sinks_void + sinks_offset);
2233
+ m_sink = sinks_ptr[head_idx];
2234
+ has_sink = true;
2235
+ m = max(m, m_sink);
2236
+ }
2237
+ ACC_TYPE l = 0.0f;
2238
+ for (int c = 0; c < n_splits; ++c) {
2239
+ const ACC_TYPE m_c = rec0[c * record_stride + 0];
2240
+ const ACC_TYPE l_c = rec0[c * record_stride + 1];
2241
+ if (m_c > FA_M_INIT) {
2242
+ l += l_c * exp(m_c - m);
2243
+ }
2244
+ }
2245
+ if (has_sink) {
2246
+ l += exp(m_sink - m);
2247
+ }
2248
+ m_final_shared = m;
2249
+ l_final_shared = l;
2250
+ }
2251
+ barrier(CLK_LOCAL_MEM_FENCE);
2252
+ const ACC_TYPE m_final = m_final_shared;
2253
+ const ACC_TYPE l_final = l_final_shared;
2254
+ const ACC_TYPE l_inv = (l_final > 0.0f) ? (1.0f / l_final) : 0.0f;
2255
+
2256
+ ACC_TYPE4 o = (ACC_TYPE4)(0.0f);
2257
+ for (int c = 0; c < n_splits; ++c) {
2258
+ const global float * rec_c = rec0 + c * record_stride;
2259
+ const ACC_TYPE m_c = rec_c[0];
2260
+ if (m_c <= FA_M_INIT) continue;
2261
+ const global float4 * rec_oc = (const global float4 *) (rec_c + 2);
2262
+ const ACC_TYPE scale_c = exp(m_c - m_final);
2263
+ o = mad((ACC_TYPE4)(scale_c), rec_oc[lane], o);
2264
+ }
2265
+ o = o * l_inv;
2266
+
2267
+ const ulong o_row_offset = (ulong) batch_idx * o_nb3 + (ulong) q_idx * o_nb2 + (ulong) head_idx * o_nb1;
2268
+ global O_DATA_TYPE4 * o_row = (global O_DATA_TYPE4 *) ((global char *) o_void + o_offset + o_row_offset);
2269
+ o_row[lane] = CONVERT_O_DATA4(o);
2270
+ }
2271
+ #endif // !FA_MQ_ONLY (q1_split + merge)
2272
+ #endif // !FA_PREFILL_ONLY (decode kernels)