whispercpp 1.3.7 → 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 (308) hide show
  1. checksums.yaml +4 -4
  2. data/README.md +5 -4
  3. data/ext/options.rb +1 -1
  4. data/ext/ruby_whisper.c +0 -1
  5. data/ext/ruby_whisper.h +7 -1
  6. data/ext/ruby_whisper_context.c +50 -1
  7. data/ext/ruby_whisper_log_settable.h +1 -2
  8. data/ext/ruby_whisper_params.c +9 -8
  9. data/ext/ruby_whisper_transcribe.cpp +0 -19
  10. data/ext/ruby_whisper_vad_context.c +30 -10
  11. data/ext/ruby_whisper_vad_context_detect.cpp +8 -9
  12. data/ext/ruby_whisper_vad_params.c +4 -4
  13. data/ext/ruby_whisper_vad_segment.c +2 -2
  14. data/ext/sources/CMakeLists.txt +2 -1
  15. data/ext/sources/cmake/parakeet.pc.in +2 -2
  16. data/ext/sources/cmake/whisper.pc.in +2 -2
  17. data/ext/sources/examples/cli/cli.cpp +9 -1
  18. data/ext/sources/examples/common-ggml.cpp +2 -0
  19. data/ext/sources/examples/vad-speech-segments/speech.cpp +3 -2
  20. data/ext/sources/ggml/CMakeLists.txt +3 -4
  21. data/ext/sources/ggml/include/ggml-cuda.h +0 -3
  22. data/ext/sources/ggml/include/ggml-sycl.h +8 -0
  23. data/ext/sources/ggml/include/ggml.h +3 -1
  24. data/ext/sources/ggml/src/CMakeLists.txt +8 -1
  25. data/ext/sources/ggml/src/ggml-backend-meta.cpp +7 -4
  26. data/ext/sources/ggml/src/ggml-common.h +13 -2
  27. data/ext/sources/ggml/src/ggml-cpu/CMakeLists.txt +1 -1
  28. data/ext/sources/ggml/src/ggml-cpu/amx/mmq.cpp +5 -6
  29. data/ext/sources/ggml/src/ggml-cpu/arch/arm/quants.c +78 -4
  30. data/ext/sources/ggml/src/ggml-cpu/arch/x86/quants.c +142 -4
  31. data/ext/sources/ggml/src/ggml-cpu/arch-fallback.h +7 -2
  32. data/ext/sources/ggml/src/ggml-cpu/ggml-cpu.c +14 -0
  33. data/ext/sources/ggml/src/ggml-cpu/llamafile/sgemm.cpp +26 -19
  34. data/ext/sources/ggml/src/ggml-cpu/ops.cpp +129 -46
  35. data/ext/sources/ggml/src/ggml-cpu/quants.c +51 -0
  36. data/ext/sources/ggml/src/ggml-cpu/quants.h +3 -0
  37. data/ext/sources/ggml/src/ggml-cpu/simd-gemm.h +1 -1
  38. data/ext/sources/ggml/src/ggml-cpu/simd-mappings.h +11 -0
  39. data/ext/sources/ggml/src/ggml-cpu/vec.cpp +2 -2
  40. data/ext/sources/ggml/src/ggml-cuda/binbcast.cu +90 -46
  41. data/ext/sources/ggml/src/ggml-cuda/col2im-1d.cu +81 -0
  42. data/ext/sources/ggml/src/ggml-cuda/col2im-1d.cuh +3 -0
  43. data/ext/sources/ggml/src/ggml-cuda/common.cuh +4 -0
  44. data/ext/sources/ggml/src/ggml-cuda/concat.cu +33 -21
  45. data/ext/sources/ggml/src/ggml-cuda/conv-transpose-1d.cu +14 -12
  46. data/ext/sources/ggml/src/ggml-cuda/convert.cu +86 -34
  47. data/ext/sources/ggml/src/ggml-cuda/cpy.cu +80 -29
  48. data/ext/sources/ggml/src/ggml-cuda/fattn-common.cuh +9 -5
  49. data/ext/sources/ggml/src/ggml-cuda/fattn-mma-f16.cuh +4 -0
  50. data/ext/sources/ggml/src/ggml-cuda/fattn-tile.cuh +9 -5
  51. data/ext/sources/ggml/src/ggml-cuda/fattn.cu +27 -21
  52. data/ext/sources/ggml/src/ggml-cuda/gated_delta_net.cu +40 -25
  53. data/ext/sources/ggml/src/ggml-cuda/gated_delta_net.cuh +10 -0
  54. data/ext/sources/ggml/src/ggml-cuda/getrows.cu +15 -12
  55. data/ext/sources/ggml/src/ggml-cuda/ggml-cuda.cu +718 -1248
  56. data/ext/sources/ggml/src/ggml-cuda/mmq.cu +7 -0
  57. data/ext/sources/ggml/src/ggml-cuda/mmvq.cu +77 -40
  58. data/ext/sources/ggml/src/ggml-cuda/out-prod.cu +55 -12
  59. data/ext/sources/ggml/src/ggml-cuda/set-rows.cu +64 -4
  60. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_16-ncols2_2.cu +1 -0
  61. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_32-ncols2_2.cu +1 -0
  62. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_2.cu +1 -0
  63. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_8-ncols2_2.cu +1 -0
  64. data/ext/sources/ggml/src/ggml-cuda/topk-moe.cu +7 -1
  65. data/ext/sources/ggml/src/ggml-cuda/vendors/hip.h +1 -0
  66. data/ext/sources/ggml/src/ggml-cuda/vendors/musa.h +1 -0
  67. data/ext/sources/ggml/src/ggml-hexagon/CMakeLists.txt +0 -5
  68. data/ext/sources/ggml/src/ggml-hexagon/ggml-hexagon.cpp +1634 -1293
  69. data/ext/sources/ggml/src/ggml-hexagon/htp/CMakeLists.txt +11 -40
  70. data/ext/sources/ggml/src/ggml-hexagon/htp/cmake-toolchain.cmake +13 -15
  71. data/ext/sources/ggml/src/ggml-hexagon/htp/concat-ops.c +1 -1
  72. data/ext/sources/ggml/src/ggml-hexagon/htp/flash-attn-ops.c +1749 -399
  73. data/ext/sources/ggml/src/ggml-hexagon/htp/flash-attn-ops.h +303 -0
  74. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-common.h +80 -0
  75. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-dma.h +26 -23
  76. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-profile.h +64 -0
  77. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-utils.h +1 -83
  78. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-fa-kernels.h +555 -0
  79. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-mm-kernels-tiled.h +1303 -0
  80. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-queue.c +9 -0
  81. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-queue.h +27 -4
  82. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-utils.h +59 -37
  83. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-ctx.h +11 -3
  84. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-ops.h +52 -12
  85. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-vtcm.h +19 -0
  86. data/ext/sources/ggml/src/ggml-hexagon/htp/htp_iface.idl +2 -1
  87. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-base.h +14 -30
  88. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-exp.h +39 -0
  89. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-fa-kernels.h +232 -0
  90. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-flat.h +1511 -0
  91. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-tiled.h +1200 -0
  92. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-sigmoid.h +39 -0
  93. data/ext/sources/ggml/src/ggml-hexagon/htp/main.c +127 -32
  94. data/ext/sources/ggml/src/ggml-hexagon/htp/matmul-ops.c +3023 -4425
  95. data/ext/sources/ggml/src/ggml-hexagon/htp/matmul-ops.h +650 -0
  96. data/ext/sources/ggml/src/ggml-hexagon/htp/rope-ops.c +48 -13
  97. data/ext/sources/ggml/src/ggml-hexagon/htp/ssm-conv.c +10 -9
  98. data/ext/sources/ggml/src/ggml-hexagon/htp/worker-pool.c +15 -3
  99. data/ext/sources/ggml/src/ggml-hexagon/htp/worker-pool.h +8 -0
  100. data/ext/sources/ggml/src/ggml-hexagon/htp-opnode.h +168 -50
  101. data/ext/sources/ggml/src/ggml-hexagon/libggml-htp.inf +0 -4
  102. data/ext/sources/ggml/src/ggml-hip/CMakeLists.txt +5 -0
  103. data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.cpp +69 -5
  104. data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.h +4 -1
  105. data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.m +27 -6
  106. data/ext/sources/ggml/src/ggml-metal/ggml-metal-impl.h +38 -0
  107. data/ext/sources/ggml/src/ggml-metal/ggml-metal-ops.cpp +132 -2
  108. data/ext/sources/ggml/src/ggml-metal/ggml-metal-ops.h +2 -0
  109. data/ext/sources/ggml/src/ggml-metal/ggml-metal.metal +345 -87
  110. data/ext/sources/ggml/src/ggml-opencl/CMakeLists.txt +13 -0
  111. data/ext/sources/ggml/src/ggml-opencl/fa_tune.h +92 -0
  112. data/ext/sources/ggml/src/ggml-opencl/ggml-opencl.cpp +4060 -357
  113. data/ext/sources/ggml/src/ggml-opencl/kernels/cvt.cl +198 -0
  114. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f16.cl +81 -41
  115. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32.cl +88 -39
  116. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_f16.cl +1995 -96
  117. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_q4_0.cl +1615 -0
  118. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_q8_0.cl +1486 -0
  119. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_pre_f16.cl +156 -0
  120. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_mxfp4_f32_ns.cl +74 -6
  121. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_0_f32_ns.cl +74 -6
  122. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_1_f32_ns.cl +74 -6
  123. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_k_f32_ns.cl +71 -6
  124. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_0_f32_ns.cl +74 -6
  125. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_1_f32_ns.cl +74 -6
  126. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_k_f32_ns.cl +74 -6
  127. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q6_k_f32_ns.cl +74 -6
  128. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q1_0_f32.cl +94 -0
  129. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q1_0_f32.cl +121 -0
  130. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q8_0_f32.cl +1 -1
  131. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q1_0_f32_l4_lm.cl +156 -0
  132. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_f16_f32_l4.cl +1149 -0
  133. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q1_0_f32.cl +141 -0
  134. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q1_0_f32_flat.cl +190 -0
  135. data/ext/sources/ggml/src/ggml-opencl/kernels/norm.cl +5 -2
  136. data/ext/sources/ggml/src/ggml-opencl/kernels/set_rows.cl +500 -0
  137. data/ext/sources/ggml/src/ggml-opencl/libdl.h +79 -0
  138. data/ext/sources/ggml/src/ggml-openvino/.clang-format +0 -5
  139. data/ext/sources/ggml/src/ggml-openvino/CMakeLists.txt +2 -4
  140. data/ext/sources/ggml/src/ggml-openvino/ggml-decoder.cpp +733 -130
  141. data/ext/sources/ggml/src/ggml-openvino/ggml-decoder.h +76 -23
  142. data/ext/sources/ggml/src/ggml-openvino/ggml-openvino-extra.cpp +57 -3
  143. data/ext/sources/ggml/src/ggml-openvino/ggml-openvino-extra.h +29 -8
  144. data/ext/sources/ggml/src/ggml-openvino/ggml-openvino.cpp +307 -59
  145. data/ext/sources/ggml/src/ggml-openvino/ggml-quants.cpp +66 -0
  146. data/ext/sources/ggml/src/ggml-openvino/ggml-quants.h +10 -4
  147. data/ext/sources/ggml/src/ggml-openvino/openvino/decoder.h +56 -16
  148. data/ext/sources/ggml/src/ggml-openvino/openvino/frontend.h +1 -1
  149. data/ext/sources/ggml/src/ggml-openvino/openvino/input_model.h +4 -4
  150. data/ext/sources/ggml/src/ggml-openvino/openvino/node_context.h +94 -37
  151. data/ext/sources/ggml/src/ggml-openvino/openvino/op/add_id.cpp +76 -0
  152. data/ext/sources/ggml/src/ggml-openvino/openvino/op/argsort.cpp +47 -0
  153. data/ext/sources/ggml/src/ggml-openvino/openvino/op/clamp.cpp +33 -0
  154. data/ext/sources/ggml/src/ggml-openvino/openvino/op/concat.cpp +48 -0
  155. data/ext/sources/ggml/src/ggml-openvino/openvino/op/cont.cpp +8 -16
  156. data/ext/sources/ggml/src/ggml-openvino/openvino/op/cpy.cpp +14 -1
  157. data/ext/sources/ggml/src/ggml-openvino/openvino/op/div.cpp +146 -0
  158. data/ext/sources/ggml/src/ggml-openvino/openvino/op/flash_attn_ext.cpp +108 -21
  159. data/ext/sources/ggml/src/ggml-openvino/openvino/op/gated_delta_net.cpp +282 -0
  160. data/ext/sources/ggml/src/ggml-openvino/openvino/op/gated_delta_net.hpp +65 -0
  161. data/ext/sources/ggml/src/ggml-openvino/openvino/op/get_rows.cpp +2 -9
  162. data/ext/sources/ggml/src/ggml-openvino/openvino/op/glu_geglu.cpp +21 -7
  163. data/ext/sources/ggml/src/ggml-openvino/openvino/op/glu_swiglu.cpp +41 -8
  164. data/ext/sources/ggml/src/ggml-openvino/openvino/op/im2col.cpp +120 -0
  165. data/ext/sources/ggml/src/ggml-openvino/openvino/op/l2_norm.cpp +44 -0
  166. data/ext/sources/ggml/src/ggml-openvino/openvino/op/mul_mat_id.cpp +226 -0
  167. data/ext/sources/ggml/src/ggml-openvino/openvino/op/mulmat.cpp +19 -9
  168. data/ext/sources/ggml/src/ggml-openvino/openvino/op/norm.cpp +58 -0
  169. data/ext/sources/ggml/src/ggml-openvino/openvino/op/pad.cpp +95 -0
  170. data/ext/sources/ggml/src/ggml-openvino/openvino/op/permute.cpp +58 -13
  171. data/ext/sources/ggml/src/ggml-openvino/openvino/op/repeat.cpp +74 -0
  172. data/ext/sources/ggml/src/ggml-openvino/openvino/op/reshape.cpp +13 -6
  173. data/ext/sources/ggml/src/ggml-openvino/openvino/op/rms_norm.cpp +1 -1
  174. data/ext/sources/ggml/src/ggml-openvino/openvino/op/rope.cpp +134 -38
  175. data/ext/sources/ggml/src/ggml-openvino/openvino/op/set_rows.cpp +3 -3
  176. data/ext/sources/ggml/src/ggml-openvino/openvino/op/softmax.cpp +126 -49
  177. data/ext/sources/ggml/src/ggml-openvino/openvino/op/ssm_conv.cpp +59 -0
  178. data/ext/sources/ggml/src/ggml-openvino/openvino/op/sum_rows.cpp +27 -0
  179. data/ext/sources/ggml/src/ggml-openvino/openvino/op/transpose.cpp +32 -1
  180. data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_silu.cpp +1 -1
  181. data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_softplus.cpp +38 -0
  182. data/ext/sources/ggml/src/ggml-openvino/openvino/op/view.cpp +90 -25
  183. data/ext/sources/ggml/src/ggml-openvino/openvino/op_table.cpp +41 -23
  184. data/ext/sources/ggml/src/ggml-openvino/openvino/op_table.h +18 -5
  185. data/ext/sources/ggml/src/ggml-openvino/openvino/pass/mark_decompression_convert_constant_folding.h +1 -1
  186. data/ext/sources/ggml/src/ggml-openvino/openvino/translate_session.cpp +43 -40
  187. data/ext/sources/ggml/src/ggml-openvino/openvino/translate_session.h +5 -4
  188. data/ext/sources/ggml/src/ggml-openvino/openvino/utils.cpp +548 -3
  189. data/ext/sources/ggml/src/ggml-openvino/openvino/utils.h +28 -26
  190. data/ext/sources/ggml/src/ggml-openvino/utils.cpp +383 -94
  191. data/ext/sources/ggml/src/ggml-openvino/utils.h +11 -8
  192. data/ext/sources/ggml/src/ggml-quants.c +76 -0
  193. data/ext/sources/ggml/src/ggml-quants.h +3 -0
  194. data/ext/sources/ggml/src/ggml-sycl/CMakeLists.txt +5 -5
  195. data/ext/sources/ggml/src/ggml-sycl/backend.hpp +2 -0
  196. data/ext/sources/ggml/src/ggml-sycl/binbcast.cpp +12 -0
  197. data/ext/sources/ggml/src/ggml-sycl/col2im-1d.cpp +102 -0
  198. data/ext/sources/ggml/src/ggml-sycl/col2im-1d.hpp +8 -0
  199. data/ext/sources/ggml/src/ggml-sycl/common.cpp +6 -8
  200. data/ext/sources/ggml/src/ggml-sycl/common.hpp +19 -2
  201. data/ext/sources/ggml/src/ggml-sycl/concat.cpp +21 -1
  202. data/ext/sources/ggml/src/ggml-sycl/conv2d-dw.cpp +158 -0
  203. data/ext/sources/ggml/src/ggml-sycl/conv2d-dw.hpp +10 -0
  204. data/ext/sources/ggml/src/ggml-sycl/conv2d-transpose.cpp +125 -0
  205. data/ext/sources/ggml/src/ggml-sycl/conv2d-transpose.hpp +10 -0
  206. data/ext/sources/ggml/src/ggml-sycl/conv2d.cpp +150 -0
  207. data/ext/sources/ggml/src/ggml-sycl/conv2d.hpp +10 -0
  208. data/ext/sources/ggml/src/ggml-sycl/conv3d.cpp +224 -0
  209. data/ext/sources/ggml/src/ggml-sycl/conv3d.hpp +8 -0
  210. data/ext/sources/ggml/src/ggml-sycl/convert.cpp +6 -0
  211. data/ext/sources/ggml/src/ggml-sycl/cpy.cpp +706 -0
  212. data/ext/sources/ggml/src/ggml-sycl/cpy.hpp +281 -0
  213. data/ext/sources/ggml/src/ggml-sycl/cross_entropy_loss.cpp +255 -0
  214. data/ext/sources/ggml/src/ggml-sycl/cross_entropy_loss.hpp +7 -0
  215. data/ext/sources/ggml/src/ggml-sycl/dequantize.hpp +15 -0
  216. data/ext/sources/ggml/src/ggml-sycl/dmmv.cpp +492 -319
  217. data/ext/sources/ggml/src/ggml-sycl/dpct/helper.hpp +15 -7
  218. data/ext/sources/ggml/src/ggml-sycl/element_wise.cpp +215 -115
  219. data/ext/sources/ggml/src/ggml-sycl/element_wise.hpp +2 -0
  220. data/ext/sources/ggml/src/ggml-sycl/ggml-sycl.cpp +1006 -336
  221. data/ext/sources/ggml/src/ggml-sycl/mmvq.cpp +252 -67
  222. data/ext/sources/ggml/src/ggml-sycl/mmvq.hpp +17 -0
  223. data/ext/sources/ggml/src/ggml-sycl/norm.cpp +103 -49
  224. data/ext/sources/ggml/src/ggml-sycl/outprod.cpp +45 -9
  225. data/ext/sources/ggml/src/ggml-sycl/pool.cpp +185 -0
  226. data/ext/sources/ggml/src/ggml-sycl/pool.hpp +22 -0
  227. data/ext/sources/ggml/src/ggml-sycl/presets.hpp +3 -1
  228. data/ext/sources/ggml/src/ggml-sycl/set_rows.cpp +10 -2
  229. data/ext/sources/ggml/src/ggml-sycl/softmax.cpp +9 -10
  230. data/ext/sources/ggml/src/ggml-sycl/vecdotq.hpp +35 -0
  231. data/ext/sources/ggml/src/ggml-vulkan/CMakeLists.txt +5 -0
  232. data/ext/sources/ggml/src/ggml-vulkan/ggml-vulkan.cpp +833 -215
  233. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/col2im_1d.comp +61 -0
  234. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/conv2d_mm.comp +1 -1
  235. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/conv3d_mm.comp +431 -0
  236. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/diag.comp +3 -3
  237. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp +1 -0
  238. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp +1 -0
  239. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/generic_unary_head.glsl +21 -19
  240. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/get_rows_back.comp +25 -0
  241. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/glu_head.glsl +23 -4
  242. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/glu_main.glsl +14 -18
  243. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/l2_norm.comp +4 -7
  244. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq.comp +21 -24
  245. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp +31 -23
  246. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp +6 -5
  247. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl +84 -67
  248. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/norm.comp +10 -10
  249. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/repeat_back.comp +3 -3
  250. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/roll.comp +3 -3
  251. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/tri.comp +3 -3
  252. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/unary.comp +168 -0
  253. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +121 -74
  254. data/ext/sources/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp +26 -19
  255. data/ext/sources/ggml/src/ggml-webgpu/ggml-webgpu.cpp +31 -36
  256. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/common_decls.tmpl +16 -2
  257. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_split.wgsl +7 -7
  258. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/get_rows.wgsl +21 -0
  259. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_decls.tmpl +439 -320
  260. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_id_vec.wgsl +2 -2
  261. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec.wgsl +45 -39
  262. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_acc.tmpl +586 -465
  263. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_q_acc.tmpl +63 -69
  264. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/quantize_q8.wgsl +14 -9
  265. data/ext/sources/ggml/src/ggml.c +36 -14
  266. data/ext/sources/include/whisper.h +21 -0
  267. data/ext/sources/src/whisper.cpp +164 -14
  268. data/lib/whisper/log_settable.rb +5 -8
  269. data/lib/whisper/model/uri.rb +0 -7
  270. data/sig/whisper.rbs +6 -0
  271. data/test/test_vad.rb +9 -0
  272. data/test/test_vad_context.rb +2 -2
  273. data/whispercpp.gemspec +1 -1
  274. metadata +62 -37
  275. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-flash-attn-ops.c +0 -1878
  276. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-matmul-ops.c +0 -2066
  277. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-ops.c +0 -6
  278. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-ops.h +0 -88
  279. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-profile.h +0 -34
  280. data/ext/sources/ggml/src/ggml-hexagon/htp/vtcm-utils.h +0 -16
  281. data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_gelu.cpp +0 -25
  282. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/abs.comp +0 -21
  283. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/ceil.comp +0 -22
  284. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/clamp.comp +0 -17
  285. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/cos.comp +0 -17
  286. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/elu.comp +0 -27
  287. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/exp.comp +0 -20
  288. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/floor.comp +0 -22
  289. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu.comp +0 -25
  290. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu_erf.comp +0 -39
  291. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu_quick.comp +0 -23
  292. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/hardsigmoid.comp +0 -22
  293. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/hardswish.comp +0 -22
  294. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/leaky_relu.comp +0 -22
  295. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/neg.comp +0 -20
  296. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/relu.comp +0 -21
  297. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/round.comp +0 -29
  298. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sgn.comp +0 -21
  299. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sigmoid.comp +0 -20
  300. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/silu.comp +0 -22
  301. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sin.comp +0 -17
  302. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/softplus.comp +0 -23
  303. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sqrt.comp +0 -17
  304. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/square.comp +0 -17
  305. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/step.comp +0 -22
  306. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/tanh.comp +0 -20
  307. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/trunc.comp +0 -22
  308. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/xielu.comp +0 -35
@@ -18,6 +18,14 @@
18
18
  #define REQD_SUBGROUP_SIZE_128 __attribute__((qcom_reqd_sub_group_size("full")))
19
19
  #endif
20
20
 
21
+ #ifdef cl_khr_subgroup_shuffle
22
+ #pragma OPENCL EXTENSION cl_khr_subgroup_shuffle : enable
23
+ #define HAS_SUBGROUP_SHUFFLE 1
24
+ #elif defined(cl_qcom_subgroup_shuffle)
25
+ #pragma OPENCL EXTENSION cl_qcom_subgroup_shuffle : enable
26
+ #define HAS_SUBGROUP_SHUFFLE 1
27
+ #endif
28
+
21
29
  // Assumes row size (ne00) is a multiple of 4
22
30
  #ifdef ADRENO_GPU
23
31
  REQD_SUBGROUP_SIZE_64
@@ -82,3 +90,1144 @@ kernel void kernel_mul_mat_f16_f32_l4(
82
90
  }
83
91
  }
84
92
  }
93
+
94
+ // Each subgroup produces DR_NDST outputs, assumes ne11 == 1
95
+ #define MUL_MAT_F16_F32_L4_DR_NDST 4
96
+
97
+ #ifdef ADRENO_GPU
98
+ REQD_SUBGROUP_SIZE_64
99
+ #endif
100
+ kernel void kernel_mul_mat_f16_f32_l4_dr(
101
+ global char * src0,
102
+ ulong offset0,
103
+ global char * src1,
104
+ ulong offset1,
105
+ global float * dst,
106
+ ulong offsetd,
107
+ int ne00,
108
+ int ne01,
109
+ int ne02,
110
+ ulong nb00,
111
+ ulong nb01,
112
+ ulong nb02,
113
+ ulong nb03,
114
+ int ne10,
115
+ int ne11,
116
+ int ne12,
117
+ ulong nb10,
118
+ ulong nb11,
119
+ ulong nb12,
120
+ ulong nb13,
121
+ int ne0,
122
+ int ne1,
123
+ int r2,
124
+ int r3
125
+ ) {
126
+ src0 = (global char*)((global char*)src0 + offset0);
127
+ src1 = (global char*)((global char*)src1 + offset1);
128
+ dst = (global float*)((global char*)dst + offsetd);
129
+
130
+ const int r0_base = get_group_id(0) * MUL_MAT_F16_F32_L4_DR_NDST;
131
+ const int im = get_group_id(2);
132
+
133
+ const int i12 = im % ne12;
134
+ const int i13 = im / ne12;
135
+
136
+ // assume ne11 == 1
137
+ const ulong offset_src1 = i12*nb12 + i13*nb13;
138
+ global float4 * y4 = (global float4 *)(src1 + offset_src1);
139
+
140
+ global half4 * x4[MUL_MAT_F16_F32_L4_DR_NDST];
141
+ float sumf[MUL_MAT_F16_F32_L4_DR_NDST];
142
+
143
+ const ulong k_head_off = (i12/r2)*nb02 + (i13/r3)*nb03;
144
+
145
+ #pragma unroll
146
+ for (int n = 0; n < MUL_MAT_F16_F32_L4_DR_NDST; ++n) {
147
+ int r0 = r0_base + n;
148
+ int r0c = r0 < ne01 ? r0 : 0;
149
+ ulong off = (ulong)r0c*nb01 + k_head_off;
150
+ x4[n] = (global half4 *)(src0 + off);
151
+ sumf[n] = 0.0f;
152
+ }
153
+
154
+ const int n_chunks = ne00 / 4;
155
+ const int sg_size = get_max_sub_group_size();
156
+ const int lid = get_sub_group_local_id();
157
+
158
+ for (int i = lid; i < n_chunks; i += sg_size) {
159
+ float4 q = y4[i];
160
+ #pragma unroll
161
+ for (int n = 0; n < MUL_MAT_F16_F32_L4_DR_NDST; ++n) {
162
+ float4 k = convert_float4(x4[n][i]);
163
+ sumf[n] = mad(k.s0, q.s0, sumf[n]);
164
+ sumf[n] = mad(k.s1, q.s1, sumf[n]);
165
+ sumf[n] = mad(k.s2, q.s2, sumf[n]);
166
+ sumf[n] = mad(k.s3, q.s3, sumf[n]);
167
+ }
168
+ }
169
+
170
+ #pragma unroll
171
+ for (int n = 0; n < MUL_MAT_F16_F32_L4_DR_NDST; ++n) {
172
+ float reduced = sub_group_reduce_add(sumf[n]);
173
+ int r0 = r0_base + n;
174
+ if (lid == 0 && r0 < ne01) {
175
+ dst[im*ne1*ne0 + r0] = reduced;
176
+ }
177
+ }
178
+ }
179
+
180
+ // Kernels for decoding, Adreno only for now
181
+ #define MUL_MAT_F16_F32_L4_DR_LS_R2_MAX 8
182
+
183
+ #ifdef ADRENO_GPU
184
+ #pragma OPENCL EXTENSION cl_qcom_subgroup_shuffle : enable
185
+ #define sub_group_shuffle_xor(val, mask) qcom_sub_group_shuffle_xor((val), (mask), CLK_SUB_GROUP_SHUFFLE_WIDTH_WAVE_SIZE_QCOM, 0.0f)
186
+
187
+ REQD_SUBGROUP_SIZE_64
188
+ kernel void kernel_mul_mat_f16_f32_l4_dr_ls(
189
+ global char * src0,
190
+ ulong offset0,
191
+ global char * src1,
192
+ ulong offset1,
193
+ global float * dst,
194
+ ulong offsetd,
195
+ int ne00,
196
+ int ne01,
197
+ int ne02,
198
+ ulong nb00,
199
+ ulong nb01,
200
+ ulong nb02,
201
+ ulong nb03,
202
+ int ne10,
203
+ int ne11,
204
+ int ne12,
205
+ ulong nb10,
206
+ ulong nb11,
207
+ ulong nb12,
208
+ ulong nb13,
209
+ int ne0,
210
+ int ne1,
211
+ int r2,
212
+ int r3
213
+ ) {
214
+ src0 = (global char*)((global char*)src0 + offset0);
215
+ src1 = (global char*)((global char*)src1 + offset1);
216
+ dst = (global float*)((global char*)dst + offsetd);
217
+
218
+ const int r0_base = get_group_id(0) * 2;
219
+ const int kv_grp = get_group_id(2); // KV head group; im = kv_grp*r2 + q
220
+
221
+ const int i12_kv = kv_grp % ne02;
222
+ const int i13_kv = kv_grp / ne02;
223
+
224
+ const int lid = get_sub_group_local_id();
225
+ const int subhalf = lid >> 5; // 0 or 1 (which K row in the WG)
226
+ const int intra = lid & 31; // 0..31 (lane within the half)
227
+
228
+ const int r0 = r0_base + subhalf;
229
+ const int r0c = r0 < ne01 ? r0 : 0; // clamp OOB to row 0; skip write below
230
+
231
+ // K row pointer for this lane (one K row per half-wave).
232
+ const ulong k_off = (ulong)r0c*nb01 + (ulong)i12_kv*nb02 + (ulong)i13_kv*nb03;
233
+ global half4 * x4 = (global half4 *)(src0 + k_off);
234
+
235
+ global float4 * y4[MUL_MAT_F16_F32_L4_DR_LS_R2_MAX];
236
+ #pragma unroll
237
+ for (int q = 0; q < MUL_MAT_F16_F32_L4_DR_LS_R2_MAX; ++q) {
238
+ const int i12_q = i12_kv*r2 + q;
239
+ const ulong q_off = (ulong)i12_q*nb12 + (ulong)i13_kv*nb13;
240
+ y4[q] = (global float4 *)(src1 + q_off);
241
+ }
242
+
243
+ float partial[MUL_MAT_F16_F32_L4_DR_LS_R2_MAX];
244
+ #pragma unroll
245
+ for (int q = 0; q < MUL_MAT_F16_F32_L4_DR_LS_R2_MAX; ++q) {
246
+ partial[q] = 0.0f;
247
+ }
248
+
249
+ const int n_chunks = ne00 / 4;
250
+
251
+ for (int i = intra; i < n_chunks; i += 32) {
252
+ float4 k = convert_float4(x4[i]);
253
+
254
+ #pragma unroll
255
+ for (int q = 0; q < MUL_MAT_F16_F32_L4_DR_LS_R2_MAX; ++q) {
256
+ if (q < r2) {
257
+ float4 v = y4[q][i];
258
+ partial[q] = mad(k.s0, v.s0, partial[q]);
259
+ partial[q] = mad(k.s1, v.s1, partial[q]);
260
+ partial[q] = mad(k.s2, v.s2, partial[q]);
261
+ partial[q] = mad(k.s3, v.s3, partial[q]);
262
+ }
263
+ }
264
+ }
265
+
266
+ // half-wave reduction
267
+ #pragma unroll
268
+ for (int q = 0; q < MUL_MAT_F16_F32_L4_DR_LS_R2_MAX; ++q) {
269
+ if (q < r2) {
270
+ partial[q] += sub_group_shuffle_xor(partial[q], 1u);
271
+ partial[q] += sub_group_shuffle_xor(partial[q], 2u);
272
+ partial[q] += sub_group_shuffle_xor(partial[q], 4u);
273
+ partial[q] += sub_group_shuffle_xor(partial[q], 8u);
274
+ partial[q] += sub_group_shuffle_xor(partial[q], 16u);
275
+ }
276
+ }
277
+
278
+ if (intra == 0 && r0 < ne01) {
279
+ #pragma unroll
280
+ for (int q = 0; q < MUL_MAT_F16_F32_L4_DR_LS_R2_MAX; ++q) {
281
+ if (q < r2) {
282
+ const int im = i12_kv*r2 + q + i13_kv*ne12;
283
+ dst[im*ne1*ne0 + r0] = partial[q];
284
+ }
285
+ }
286
+ }
287
+ }
288
+
289
+ REQD_SUBGROUP_SIZE_64
290
+ kernel void kernel_mul_mat_f16_f32_l4_dr_lq(
291
+ global char * src0,
292
+ ulong offset0,
293
+ global char * src1,
294
+ ulong offset1,
295
+ global float * dst,
296
+ ulong offsetd,
297
+ int ne00,
298
+ int ne01,
299
+ int ne02,
300
+ ulong nb00,
301
+ ulong nb01,
302
+ ulong nb02,
303
+ ulong nb03,
304
+ int ne10,
305
+ int ne11,
306
+ int ne12,
307
+ ulong nb10,
308
+ ulong nb11,
309
+ ulong nb12,
310
+ ulong nb13,
311
+ int ne0,
312
+ int ne1,
313
+ int r2,
314
+ int r3
315
+ ) {
316
+ src0 = (global char*)((global char*)src0 + offset0);
317
+ src1 = (global char*)((global char*)src1 + offset1);
318
+ dst = (global float*)((global char*)dst + offsetd);
319
+
320
+ const int r0_base = get_group_id(0) * 4;
321
+ const int kv_grp = get_group_id(2);
322
+
323
+ const int i12_kv = kv_grp % ne02;
324
+ const int i13_kv = kv_grp / ne02;
325
+
326
+ const int lid = get_sub_group_local_id();
327
+ const int subq = lid >> 4; // 0..3 (which K row)
328
+ const int intra = lid & 15; // 0..15 (lane within quarter)
329
+
330
+ const int r0 = r0_base + subq;
331
+ const int r0c = r0 < ne01 ? r0 : 0;
332
+
333
+ const ulong k_off = (ulong)r0c*nb01 + (ulong)i12_kv*nb02 + (ulong)i13_kv*nb03;
334
+ global half4 * x4 = (global half4 *)(src0 + k_off);
335
+
336
+ global float4 * y4[MUL_MAT_F16_F32_L4_DR_LS_R2_MAX];
337
+ #pragma unroll
338
+ for (int q = 0; q < MUL_MAT_F16_F32_L4_DR_LS_R2_MAX; ++q) {
339
+ const int i12_q = i12_kv*r2 + q;
340
+ const ulong q_off = (ulong)i12_q*nb12 + (ulong)i13_kv*nb13;
341
+ y4[q] = (global float4 *)(src1 + q_off);
342
+ }
343
+
344
+ float partial[MUL_MAT_F16_F32_L4_DR_LS_R2_MAX];
345
+ #pragma unroll
346
+ for (int q = 0; q < MUL_MAT_F16_F32_L4_DR_LS_R2_MAX; ++q) {
347
+ partial[q] = 0.0f;
348
+ }
349
+
350
+ const int n_chunks = ne00 / 4;
351
+
352
+ for (int i = intra; i < n_chunks; i += 16) {
353
+ float4 k = convert_float4(x4[i]);
354
+
355
+ #pragma unroll
356
+ for (int q = 0; q < MUL_MAT_F16_F32_L4_DR_LS_R2_MAX; ++q) {
357
+ if (q < r2) {
358
+ float4 v = y4[q][i];
359
+ partial[q] = mad(k.s0, v.s0, partial[q]);
360
+ partial[q] = mad(k.s1, v.s1, partial[q]);
361
+ partial[q] = mad(k.s2, v.s2, partial[q]);
362
+ partial[q] = mad(k.s3, v.s3, partial[q]);
363
+ }
364
+ }
365
+ }
366
+
367
+ // quarter-wave reduction
368
+ #pragma unroll
369
+ for (int q = 0; q < MUL_MAT_F16_F32_L4_DR_LS_R2_MAX; ++q) {
370
+ if (q < r2) {
371
+ partial[q] += sub_group_shuffle_xor(partial[q], 1u);
372
+ partial[q] += sub_group_shuffle_xor(partial[q], 2u);
373
+ partial[q] += sub_group_shuffle_xor(partial[q], 4u);
374
+ partial[q] += sub_group_shuffle_xor(partial[q], 8u);
375
+ }
376
+ }
377
+
378
+ if (intra == 0 && r0 < ne01) {
379
+ #pragma unroll
380
+ for (int q = 0; q < MUL_MAT_F16_F32_L4_DR_LS_R2_MAX; ++q) {
381
+ if (q < r2) {
382
+ const int im = i12_kv*r2 + q + i13_kv*ne12;
383
+ dst[im*ne1*ne0 + r0] = partial[q];
384
+ }
385
+ }
386
+ }
387
+ }
388
+ #endif // ADRENO_GPU
389
+
390
+ #define N_ROWS_PER_WG 8
391
+ #define N_OUTS_PER_WG 8
392
+
393
+ #ifdef ADRENO_GPU
394
+ REQD_SUBGROUP_SIZE_64
395
+ #endif
396
+ kernel void kernel_mul_mat_f16_f32_l4_x8(
397
+ global char * src0,
398
+ ulong offset0,
399
+ global char * src1,
400
+ ulong offset1,
401
+ global float * dst,
402
+ ulong offsetd,
403
+ int ne00,
404
+ int ne01,
405
+ int ne02,
406
+ ulong nb00,
407
+ ulong nb01,
408
+ ulong nb02,
409
+ ulong nb03,
410
+ int ne10,
411
+ int ne11,
412
+ int ne12,
413
+ ulong nb10,
414
+ ulong nb11,
415
+ ulong nb12,
416
+ ulong nb13,
417
+ int ne0,
418
+ int ne1,
419
+ int r2,
420
+ int r3
421
+ ) {
422
+ src0 = (global char *)((global char *)src0 + offset0);
423
+ src1 = (global char *)((global char *)src1 + offset1);
424
+ dst = (global float*)((global char *)dst + offsetd);
425
+
426
+ const int sgs_lid = get_sub_group_local_id();
427
+ const int sgs_sz = get_max_sub_group_size();
428
+
429
+ const int r0_base = get_group_id(0) * N_ROWS_PER_WG;
430
+ const int im = get_group_id(2);
431
+
432
+ const int i12 = im % ne12;
433
+ const int i13 = im / ne12;
434
+
435
+ const ulong offset_src1 = (i12) * nb12 + (i13) * nb13;
436
+ global float4 * y4 = (global float4 *)(src1 + offset_src1);
437
+
438
+ __local float4 q_loc[64]; // ne00/4 max for sub_group_size 64
439
+ if (sgs_lid < ne00 / 4) {
440
+ q_loc[sgs_lid] = y4[sgs_lid];
441
+ }
442
+ barrier(CLK_LOCAL_MEM_FENCE);
443
+
444
+ #pragma unroll
445
+ for (int dr = 0; dr < N_ROWS_PER_WG; ++dr) {
446
+ const int r0 = r0_base + dr;
447
+ if (r0 >= ne01) return;
448
+
449
+ const ulong offset_src0 = r0 * nb01 + (i12 / r2) * nb02 + (i13 / r3) * nb03;
450
+ global half4 * x4 = (global half4 *)(src0 + offset_src0);
451
+
452
+ float sumf = 0.0f;
453
+ for (int i = sgs_lid; i < ne00 / 4; i += sgs_sz) {
454
+ const half4 k4 = x4[i];
455
+ const float4 q = q_loc[i];
456
+ sumf += convert_float(k4.s0) * q.s0
457
+ + convert_float(k4.s1) * q.s1
458
+ + convert_float(k4.s2) * q.s2
459
+ + convert_float(k4.s3) * q.s3;
460
+ }
461
+
462
+ const float all_sum = sub_group_reduce_add(sumf);
463
+ if (sgs_lid == 0) {
464
+ dst[im * ne1 * ne0 + r0] = all_sum; // ne11 == 1, so r1==0
465
+ }
466
+ }
467
+ }
468
+
469
+ #ifdef ADRENO_GPU
470
+ REQD_SUBGROUP_SIZE_64
471
+ #endif
472
+ kernel void kernel_mul_mat_f16_f32_l4_y8(
473
+ global char * src0,
474
+ ulong offset0,
475
+ global char * src1,
476
+ ulong offset1,
477
+ global float * dst,
478
+ ulong offsetd,
479
+ int ne00,
480
+ int ne01,
481
+ int ne02,
482
+ ulong nb00,
483
+ ulong nb01,
484
+ ulong nb02,
485
+ ulong nb03,
486
+ int ne10,
487
+ int ne11,
488
+ int ne12,
489
+ ulong nb10,
490
+ ulong nb11,
491
+ ulong nb12,
492
+ ulong nb13,
493
+ int ne0,
494
+ int ne1,
495
+ int r2,
496
+ int r3
497
+ ) {
498
+ src0 = (global char *)((global char *)src0 + offset0);
499
+ src1 = (global char *)((global char *)src1 + offset1);
500
+ dst = (global float*)((global char *)dst + offsetd);
501
+
502
+ const int sgs_lid = get_sub_group_local_id();
503
+ const int sgs_sz = get_max_sub_group_size();
504
+
505
+ const int r0_base = get_group_id(0) * N_OUTS_PER_WG;
506
+ const int im = get_group_id(2);
507
+
508
+ const int i12 = im % ne12;
509
+ const int i13 = im / ne12;
510
+
511
+ const ulong offset_src1 = (i12) * nb12 + (i13) * nb13;
512
+ global float4 * y4 = (global float4 *)(src1 + offset_src1);
513
+
514
+ global half4 * x4_o[N_OUTS_PER_WG];
515
+ #pragma unroll
516
+ for (int o = 0; o < N_OUTS_PER_WG; ++o) {
517
+ const int r0 = r0_base + o;
518
+ const int r0c = (r0 < ne01) ? r0 : 0;
519
+ const ulong off = r0c * nb01 + (i12 / r2) * nb02 + (i13 / r3) * nb03;
520
+ x4_o[o] = (global half4 *)(src0 + off);
521
+ }
522
+
523
+ float sum[N_OUTS_PER_WG] = { 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f };
524
+
525
+ for (int i = sgs_lid; i < ne00 / 4; i += sgs_sz) {
526
+ const float4 q4 = y4[i];
527
+ #pragma unroll
528
+ for (int o = 0; o < N_OUTS_PER_WG; ++o) {
529
+ const half4 v4 = x4_o[o][i];
530
+ sum[o] += convert_float(v4.s0) * q4.s0
531
+ + convert_float(v4.s1) * q4.s1
532
+ + convert_float(v4.s2) * q4.s2
533
+ + convert_float(v4.s3) * q4.s3;
534
+ }
535
+ }
536
+
537
+ #pragma unroll
538
+ for (int o = 0; o < N_OUTS_PER_WG; ++o) {
539
+ const int r0 = r0_base + o;
540
+ const float s = sub_group_reduce_add(sum[o]);
541
+ if (sgs_lid == 0 && r0 < ne01) {
542
+ dst[im * ne1 * ne0 + r0] = s;
543
+ }
544
+ }
545
+ }
546
+
547
+ #define N_OUTS_PAIR 8
548
+ #define N_PAIRS_PAIR (N_OUTS_PAIR / 2)
549
+
550
+ #ifdef ADRENO_GPU
551
+ REQD_SUBGROUP_SIZE_64
552
+ #endif
553
+ kernel void kernel_mul_mat_f16_f32_l4_x8_pair(
554
+ global char * src0,
555
+ ulong offset0,
556
+ global char * src1,
557
+ ulong offset1,
558
+ global float * dst,
559
+ ulong offsetd,
560
+ int ne00,
561
+ int ne01,
562
+ int ne02,
563
+ ulong nb00,
564
+ ulong nb01,
565
+ ulong nb02,
566
+ ulong nb03,
567
+ int ne10,
568
+ int ne11,
569
+ int ne12,
570
+ ulong nb10,
571
+ ulong nb11,
572
+ ulong nb12,
573
+ ulong nb13,
574
+ int ne0,
575
+ int ne1,
576
+ int r2,
577
+ int r3
578
+ ) {
579
+ src0 = (global char *)((global char *)src0 + offset0);
580
+ src1 = (global char *)((global char *)src1 + offset1);
581
+ dst = (global float*)((global char *)dst + offsetd);
582
+
583
+ const int sgs_lid = get_sub_group_local_id();
584
+ const int half_id = sgs_lid >> 5; // 0 = lower half, 1 = upper half
585
+ const int lane_h = sgs_lid & 31; // lane 0..31 within half
586
+
587
+ const int r0_base = get_group_id(0) * N_OUTS_PAIR;
588
+ const int im = get_group_id(2);
589
+
590
+ const int i12 = im % ne12;
591
+ const int i13 = im / ne12;
592
+
593
+ const ulong offset_src1 = (i12) * nb12 + (i13) * nb13;
594
+ global float4 * y4 = (global float4 *)(src1 + offset_src1);
595
+
596
+ __local float4 q_loc[64]; // ne00/4 max for sub_group_size 64
597
+ if (sgs_lid < ne00 / 4) {
598
+ q_loc[sgs_lid] = y4[sgs_lid];
599
+ }
600
+ barrier(CLK_LOCAL_MEM_FENCE);
601
+
602
+ const int dk_vec = ne00 / 4;
603
+
604
+ #pragma unroll
605
+ for (int p = 0; p < N_PAIRS_PAIR; ++p) {
606
+ const int r0 = r0_base + 2 * p + half_id;
607
+
608
+ const ulong offset_src0 = r0 * nb01 + (i12 / r2) * nb02 + (i13 / r3) * nb03;
609
+ global half4 * x4 = (global half4 *)(src0 + offset_src0);
610
+
611
+ float sumf = 0.0f;
612
+ for (int i = lane_h; i < dk_vec; i += 32) {
613
+ const half4 k4 = x4[i];
614
+ const float4 q = q_loc[i];
615
+ sumf += convert_float(k4.s0) * q.s0
616
+ + convert_float(k4.s1) * q.s1
617
+ + convert_float(k4.s2) * q.s2
618
+ + convert_float(k4.s3) * q.s3;
619
+ }
620
+
621
+ sumf += sub_group_shuffle_xor(sumf, 16);
622
+ sumf += sub_group_shuffle_xor(sumf, 8);
623
+ sumf += sub_group_shuffle_xor(sumf, 4);
624
+ sumf += sub_group_shuffle_xor(sumf, 2);
625
+ sumf += sub_group_shuffle_xor(sumf, 1);
626
+
627
+ if (lane_h == 0) {
628
+ dst[im * ne1 * ne0 + r0] = sumf;
629
+ }
630
+ }
631
+ }
632
+
633
+ #define N_K_ROWS_GQA 16
634
+ #define GQA_RATIO_GQA 8
635
+ #define LANES_PER_QH 8 // 64 / GQA_RATIO_GQA
636
+ #define DK_VEC_GQA 32 // DK / 4 for DK=128
637
+
638
+ #ifdef ADRENO_GPU
639
+ REQD_SUBGROUP_SIZE_64
640
+ #endif
641
+ kernel void kernel_mul_mat_f16_f32_l4_x8_gqa4(
642
+ global char * src0,
643
+ ulong offset0,
644
+ global char * src1,
645
+ ulong offset1,
646
+ global float * dst,
647
+ ulong offsetd,
648
+ int ne00,
649
+ int ne01,
650
+ int ne02,
651
+ ulong nb00,
652
+ ulong nb01,
653
+ ulong nb02,
654
+ ulong nb03,
655
+ int ne10,
656
+ int ne11,
657
+ int ne12,
658
+ ulong nb10,
659
+ ulong nb11,
660
+ ulong nb12,
661
+ ulong nb13,
662
+ int ne0,
663
+ int ne1,
664
+ int r2,
665
+ int r3
666
+ ) {
667
+ src0 = (global char *)((global char *)src0 + offset0);
668
+ src1 = (global char *)((global char *)src1 + offset1);
669
+ dst = (global float*)((global char *)dst + offsetd);
670
+
671
+ const int sgs_lid = get_sub_group_local_id();
672
+ const int q_id = sgs_lid >> 3; // 0..7: which Q-head (8 per WG)
673
+ const int lane_q = sgs_lid & 7; // 0..7: lane within Q-head partition
674
+
675
+ const int r0_base = get_group_id(0) * N_K_ROWS_GQA;
676
+ const int im_kv = get_group_id(2);
677
+
678
+ const int i02 = im_kv % ne02; // K-head index (also K2 batch)
679
+ const int i03 = im_kv / ne02; // n13 batch index
680
+
681
+ const int q_head_lo = i02 * GQA_RATIO_GQA;
682
+
683
+ __local float4 q_loc[GQA_RATIO_GQA * DK_VEC_GQA]; // 4 × 32 = 128 float4
684
+ #pragma unroll
685
+ for (int qh = 0; qh < GQA_RATIO_GQA; ++qh) {
686
+ const int qh_idx = q_head_lo + qh;
687
+ global float4 * y4 = (global float4 *)(src1 + qh_idx * nb12 + i03 * nb13);
688
+
689
+ if (sgs_lid < DK_VEC_GQA) {
690
+ q_loc[qh * DK_VEC_GQA + sgs_lid] = y4[sgs_lid];
691
+ }
692
+ }
693
+ barrier(CLK_LOCAL_MEM_FENCE);
694
+
695
+ // K base offset for this WG. All 8 K-rows × 4 Q-heads share this K-head.
696
+ const ulong offset_src0_base = (i02) * nb02 + (i03 / r3) * nb03;
697
+
698
+ #pragma unroll
699
+ for (int dr = 0; dr < N_K_ROWS_GQA; ++dr) {
700
+ const int r0 = r0_base + dr;
701
+
702
+ const ulong offset_src0 = r0 * nb01 + offset_src0_base;
703
+ global half4 * x4 = (global half4 *)(src0 + offset_src0);
704
+
705
+ float sumf = 0.0f;
706
+ #pragma unroll
707
+ for (int t = 0; t < 4; ++t) {
708
+ const int i = lane_q + t * LANES_PER_QH; // 8, 16, 24-step
709
+ const half4 k4 = x4[i];
710
+ const float4 q = q_loc[q_id * DK_VEC_GQA + i];
711
+ sumf += convert_float(k4.s0) * q.s0
712
+ + convert_float(k4.s1) * q.s1
713
+ + convert_float(k4.s2) * q.s2
714
+ + convert_float(k4.s3) * q.s3;
715
+ }
716
+
717
+ sumf += sub_group_shuffle_xor(sumf, 4);
718
+ sumf += sub_group_shuffle_xor(sumf, 2);
719
+ sumf += sub_group_shuffle_xor(sumf, 1);
720
+
721
+ if (lane_q == 0) {
722
+ const int im_out = i03 * ne12 + (q_head_lo + q_id);
723
+ dst[im_out * ne1 * ne0 + r0] = sumf;
724
+ }
725
+ }
726
+ }
727
+
728
+ #define N_DV_ROWS_Y8GQA 8
729
+ #define GQA_RATIO_Y8GQA 8
730
+
731
+ #ifdef ADRENO_GPU
732
+ REQD_SUBGROUP_SIZE_64
733
+ #endif
734
+ kernel void kernel_mul_mat_f16_f32_l4_y8_gqa(
735
+ global char * src0,
736
+ ulong offset0,
737
+ global char * src1,
738
+ ulong offset1,
739
+ global float * dst,
740
+ ulong offsetd,
741
+ int ne00,
742
+ int ne01,
743
+ int ne02,
744
+ ulong nb00,
745
+ ulong nb01,
746
+ ulong nb02,
747
+ ulong nb03,
748
+ int ne10,
749
+ int ne11,
750
+ int ne12,
751
+ ulong nb10,
752
+ ulong nb11,
753
+ ulong nb12,
754
+ ulong nb13,
755
+ int ne0,
756
+ int ne1,
757
+ int r2,
758
+ int r3
759
+ ) {
760
+ src0 = (global char *)((global char *)src0 + offset0);
761
+ src1 = (global char *)((global char *)src1 + offset1);
762
+ dst = (global float*)((global char *)dst + offsetd);
763
+
764
+ const int sgs_lid = get_sub_group_local_id();
765
+ const int sgs_sz = get_max_sub_group_size();
766
+
767
+ const int r0_base = get_group_id(0) * N_DV_ROWS_Y8GQA;
768
+ const int im_kv = get_group_id(2);
769
+
770
+ const int i02 = im_kv % ne02; // K-head index
771
+ const int i03 = im_kv / ne02; // n13 batch index
772
+
773
+ // GQA Q-heads sharing this K-head.
774
+ const int q_head_lo = i02 * GQA_RATIO_Y8GQA;
775
+
776
+ global float4 * y4_q[GQA_RATIO_Y8GQA];
777
+ #pragma unroll
778
+ for (int qh = 0; qh < GQA_RATIO_Y8GQA; ++qh) {
779
+ const int qh_idx = q_head_lo + qh;
780
+ y4_q[qh] = (global float4 *)(src1 + qh_idx * nb12 + i03 * nb13);
781
+ }
782
+
783
+ global half4 * x4_o[N_DV_ROWS_Y8GQA];
784
+ #pragma unroll
785
+ for (int o = 0; o < N_DV_ROWS_Y8GQA; ++o) {
786
+ const int r0 = r0_base + o;
787
+ const int r0c = (r0 < ne01) ? r0 : 0;
788
+ const ulong off = r0c * nb01 + (i02) * nb02 + (i03 / r3) * nb03;
789
+ x4_o[o] = (global half4 *)(src0 + off);
790
+ }
791
+
792
+ float sum[N_DV_ROWS_Y8GQA][GQA_RATIO_Y8GQA] = { {0.0f} };
793
+
794
+ for (int i = sgs_lid; i < ne00 / 4; i += sgs_sz) {
795
+ // load 8 V values (one per DV row), same K-head, K-pos = i.
796
+ half4 v[N_DV_ROWS_Y8GQA];
797
+ #pragma unroll
798
+ for (int o = 0; o < N_DV_ROWS_Y8GQA; ++o) {
799
+ v[o] = x4_o[o][i];
800
+ }
801
+
802
+ // load 8 softmax values (one per Q-head).
803
+ float4 q[GQA_RATIO_Y8GQA];
804
+ #pragma unroll
805
+ for (int qh = 0; qh < GQA_RATIO_Y8GQA; ++qh) {
806
+ q[qh] = y4_q[qh][i];
807
+ }
808
+
809
+ #pragma unroll
810
+ for (int o = 0; o < N_DV_ROWS_Y8GQA; ++o) {
811
+ const float4 vf = (float4)(convert_float(v[o].s0),
812
+ convert_float(v[o].s1),
813
+ convert_float(v[o].s2),
814
+ convert_float(v[o].s3));
815
+ #pragma unroll
816
+ for (int qh = 0; qh < GQA_RATIO_Y8GQA; ++qh) {
817
+ sum[o][qh] += vf.s0 * q[qh].s0
818
+ + vf.s1 * q[qh].s1
819
+ + vf.s2 * q[qh].s2
820
+ + vf.s3 * q[qh].s3;
821
+ }
822
+ }
823
+ }
824
+
825
+ #pragma unroll
826
+ for (int o = 0; o < N_DV_ROWS_Y8GQA; ++o) {
827
+ const int r0 = r0_base + o;
828
+ #pragma unroll
829
+ for (int qh = 0; qh < GQA_RATIO_Y8GQA; ++qh) {
830
+ const float s = sub_group_reduce_add(sum[o][qh]);
831
+ if (sgs_lid == 0 && r0 < ne01) {
832
+ const int im_out = i03 * ne12 + (q_head_lo + qh);
833
+ dst[im_out * ne1 * ne0 + r0] = s;
834
+ }
835
+ }
836
+ }
837
+ }
838
+
839
+ #ifdef ADRENO_GPU
840
+ REQD_SUBGROUP_SIZE_64
841
+ #endif
842
+ kernel void kernel_mul_mat_f16_f32_l4_x8_gqa4_img(
843
+ __read_only image1d_buffer_t src0_img,
844
+ global char * src1,
845
+ ulong offset1,
846
+ global float * dst,
847
+ ulong offsetd,
848
+ int ne00,
849
+ int ne01,
850
+ int ne02,
851
+ ulong nb01,
852
+ ulong nb02,
853
+ ulong nb03,
854
+ int ne10,
855
+ int ne11,
856
+ int ne12,
857
+ ulong nb10,
858
+ ulong nb11,
859
+ ulong nb12,
860
+ ulong nb13,
861
+ int ne0,
862
+ int ne1,
863
+ int r2,
864
+ int r3
865
+ ) {
866
+ src1 = (global char *)((global char *)src1 + offset1);
867
+ dst = (global float*)((global char *)dst + offsetd);
868
+
869
+ const int sgs_lid = get_sub_group_local_id();
870
+ const int q_id = sgs_lid >> 3; // 0..7: which Q-head (8 per WG)
871
+ const int lane_q = sgs_lid & 7; // 0..7: lane within Q-head partition
872
+
873
+ const int r0_base = get_group_id(0) * N_K_ROWS_GQA;
874
+ const int im_kv = get_group_id(2);
875
+
876
+ const int i02 = im_kv % ne02;
877
+ const int i03 = im_kv / ne02;
878
+
879
+ const int q_head_lo = i02 * GQA_RATIO_GQA;
880
+
881
+ __local float4 q_loc[GQA_RATIO_GQA * DK_VEC_GQA];
882
+ #pragma unroll
883
+ for (int qh = 0; qh < GQA_RATIO_GQA; ++qh) {
884
+ const int qh_idx = q_head_lo + qh;
885
+ global float4 * y4 = (global float4 *)(src1 + qh_idx * nb12 + i03 * nb13);
886
+ if (sgs_lid < DK_VEC_GQA) {
887
+ q_loc[qh * DK_VEC_GQA + sgs_lid] = y4[sgs_lid];
888
+ }
889
+ }
890
+ barrier(CLK_LOCAL_MEM_FENCE);
891
+
892
+ const int pitch_px_row = (int)(nb01 >> 4);
893
+ const int pitch_px_head = (int)(nb02 >> 4);
894
+ const int pitch_px_n13 = (int)(nb03 >> 4);
895
+
896
+ const int head_px_base = i02 * pitch_px_head + (i03 / r3) * pitch_px_n13;
897
+
898
+ #pragma unroll
899
+ for (int dr = 0; dr < N_K_ROWS_GQA; ++dr) {
900
+ const int r0 = r0_base + dr;
901
+ const int row_px_base = r0 * pitch_px_row + head_px_base;
902
+
903
+ float sumf = 0.0f;
904
+ #pragma unroll
905
+ for (int t = 0; t < 2; ++t) {
906
+ const int p = lane_q + t * LANES_PER_QH; // pixel idx in row, 0..15
907
+ const half8 k8 = as_half8(read_imagef(src0_img, row_px_base + p));
908
+ const int i0 = 2 * p; // first half4 idx
909
+ const float4 qa = q_loc[q_id * DK_VEC_GQA + i0 ];
910
+ const float4 qb = q_loc[q_id * DK_VEC_GQA + i0 + 1];
911
+ sumf += convert_float(k8.s0) * qa.s0
912
+ + convert_float(k8.s1) * qa.s1
913
+ + convert_float(k8.s2) * qa.s2
914
+ + convert_float(k8.s3) * qa.s3
915
+ + convert_float(k8.s4) * qb.s0
916
+ + convert_float(k8.s5) * qb.s1
917
+ + convert_float(k8.s6) * qb.s2
918
+ + convert_float(k8.s7) * qb.s3;
919
+ }
920
+
921
+ sumf += sub_group_shuffle_xor(sumf, 4);
922
+ sumf += sub_group_shuffle_xor(sumf, 2);
923
+ sumf += sub_group_shuffle_xor(sumf, 1);
924
+
925
+ if (lane_q == 0) {
926
+ const int im_out = i03 * ne12 + (q_head_lo + q_id);
927
+ dst[im_out * ne1 * ne0 + r0] = sumf;
928
+ }
929
+ }
930
+ }
931
+
932
+ #ifdef ADRENO_GPU
933
+ REQD_SUBGROUP_SIZE_64
934
+ #endif
935
+ kernel void kernel_mul_mat_f16_f32_l4_y8_gqa_img(
936
+ __read_only image1d_buffer_t src0_img,
937
+ global char * src1,
938
+ ulong offset1,
939
+ global float * dst,
940
+ ulong offsetd,
941
+ int ne00,
942
+ int ne01,
943
+ int ne02,
944
+ ulong nb01,
945
+ ulong nb02,
946
+ ulong nb03,
947
+ int ne10,
948
+ int ne11,
949
+ int ne12,
950
+ ulong nb10,
951
+ ulong nb11,
952
+ ulong nb12,
953
+ ulong nb13,
954
+ int ne0,
955
+ int ne1,
956
+ int r2,
957
+ int r3
958
+ ) {
959
+ src1 = (global char *)((global char *)src1 + offset1);
960
+ dst = (global float*)((global char *)dst + offsetd);
961
+
962
+ const int sgs_lid = get_sub_group_local_id();
963
+ const int sgs_sz = get_max_sub_group_size();
964
+
965
+ const int r0_base = get_group_id(0) * N_DV_ROWS_Y8GQA;
966
+ const int im_kv = get_group_id(2);
967
+
968
+ const int i02 = im_kv % ne02;
969
+ const int i03 = im_kv / ne02;
970
+
971
+ const int q_head_lo = i02 * GQA_RATIO_Y8GQA;
972
+
973
+ // Q (= softmax(KQ)) base pointers per Q-head
974
+ global float4 * y4_q[GQA_RATIO_Y8GQA];
975
+ #pragma unroll
976
+ for (int qh = 0; qh < GQA_RATIO_Y8GQA; ++qh) {
977
+ const int qh_idx = q_head_lo + qh;
978
+ y4_q[qh] = (global float4 *)(src1 + qh_idx * nb12 + i03 * nb13);
979
+ }
980
+
981
+ const int pitch_px_row = (int)(nb01 >> 3);
982
+ const int pitch_px_head = (int)(nb02 >> 3);
983
+ const int pitch_px_n13 = (int)(nb03 >> 3);
984
+
985
+ const int head_px_base = i02 * pitch_px_head + (i03 / r3) * pitch_px_n13;
986
+
987
+ // per-DV-row pixel base
988
+ int row_px_base[N_DV_ROWS_Y8GQA];
989
+ #pragma unroll
990
+ for (int o = 0; o < N_DV_ROWS_Y8GQA; ++o) {
991
+ const int r0 = r0_base + o;
992
+ const int r0c = (r0 < ne01) ? r0 : 0;
993
+ row_px_base[o] = r0c * pitch_px_row + head_px_base;
994
+ }
995
+
996
+ float sum[N_DV_ROWS_Y8GQA][GQA_RATIO_Y8GQA] = { {0.0f} };
997
+
998
+ for (int i = sgs_lid; i < ne00 / 4; i += sgs_sz) {
999
+ half4 v[N_DV_ROWS_Y8GQA];
1000
+
1001
+ #pragma unroll
1002
+ for (int o = 0; o < N_DV_ROWS_Y8GQA; ++o) {
1003
+ v[o] = read_imageh(src0_img, row_px_base[o] + i);
1004
+ }
1005
+
1006
+ float4 q[GQA_RATIO_Y8GQA];
1007
+ #pragma unroll
1008
+ for (int qh = 0; qh < GQA_RATIO_Y8GQA; ++qh) {
1009
+ q[qh] = y4_q[qh][i];
1010
+ }
1011
+ // 64 mads.
1012
+ #pragma unroll
1013
+ for (int o = 0; o < N_DV_ROWS_Y8GQA; ++o) {
1014
+ const float4 vf = (float4)(convert_float(v[o].s0),
1015
+ convert_float(v[o].s1),
1016
+ convert_float(v[o].s2),
1017
+ convert_float(v[o].s3));
1018
+ #pragma unroll
1019
+ for (int qh = 0; qh < GQA_RATIO_Y8GQA; ++qh) {
1020
+ sum[o][qh] += vf.s0 * q[qh].s0
1021
+ + vf.s1 * q[qh].s1
1022
+ + vf.s2 * q[qh].s2
1023
+ + vf.s3 * q[qh].s3;
1024
+ }
1025
+ }
1026
+ }
1027
+
1028
+ #pragma unroll
1029
+ for (int o = 0; o < N_DV_ROWS_Y8GQA; ++o) {
1030
+ const int r0 = r0_base + o;
1031
+ #pragma unroll
1032
+ for (int qh = 0; qh < GQA_RATIO_Y8GQA; ++qh) {
1033
+ const float s = sub_group_reduce_add(sum[o][qh]);
1034
+ if (sgs_lid == 0 && r0 < ne01) {
1035
+ const int im_out = i03 * ne12 + (q_head_lo + qh);
1036
+ dst[im_out * ne1 * ne0 + r0] = s;
1037
+ }
1038
+ }
1039
+ }
1040
+ }
1041
+
1042
+ #define N_K_ROWS_GQA_R4 16
1043
+ #define GQA_RATIO_R4 4
1044
+ #define LANES_PER_QH_R4 16 // = 64 / GQA_RATIO_R4
1045
+ #define DK_VEC_R4 32 // DK / 4 for DK=128
1046
+
1047
+ #ifdef ADRENO_GPU
1048
+ REQD_SUBGROUP_SIZE_64
1049
+ #endif
1050
+ kernel void kernel_mul_mat_f16_f32_l4_x8_gqa_r4_img(
1051
+ __read_only image1d_buffer_t src0_img,
1052
+ global char * src1,
1053
+ ulong offset1,
1054
+ global float * dst,
1055
+ ulong offsetd,
1056
+ int ne00,
1057
+ int ne01,
1058
+ int ne02,
1059
+ ulong nb01,
1060
+ ulong nb02,
1061
+ ulong nb03,
1062
+ int ne10,
1063
+ int ne11,
1064
+ int ne12,
1065
+ ulong nb10,
1066
+ ulong nb11,
1067
+ ulong nb12,
1068
+ ulong nb13,
1069
+ int ne0,
1070
+ int ne1,
1071
+ int r2,
1072
+ int r3
1073
+ ) {
1074
+ src1 = (global char *)((global char *)src1 + offset1);
1075
+ dst = (global float*)((global char *)dst + offsetd);
1076
+
1077
+ const int sgs_lid = get_sub_group_local_id();
1078
+ const int q_id = sgs_lid >> 4; // 0..3
1079
+ const int lane_q = sgs_lid & 15; // 0..15
1080
+
1081
+ const int r0_base = get_group_id(0) * N_K_ROWS_GQA_R4;
1082
+ const int im_kv = get_group_id(2);
1083
+
1084
+ const int i02 = im_kv % ne02;
1085
+ const int i03 = im_kv / ne02;
1086
+
1087
+ const int q_head_lo = i02 * GQA_RATIO_R4;
1088
+
1089
+ __local float4 q_loc[GQA_RATIO_R4 * DK_VEC_R4];
1090
+ #pragma unroll
1091
+ for (int qh = 0; qh < GQA_RATIO_R4; ++qh) {
1092
+ const int qh_idx = q_head_lo + qh;
1093
+ global float4 * y4 = (global float4 *)(src1 + qh_idx * nb12 + i03 * nb13);
1094
+ if (sgs_lid < DK_VEC_R4) {
1095
+ q_loc[qh * DK_VEC_R4 + sgs_lid] = y4[sgs_lid];
1096
+ }
1097
+ }
1098
+ barrier(CLK_LOCAL_MEM_FENCE);
1099
+
1100
+ const int pitch_px_row = (int)(nb01 >> 4);
1101
+ const int pitch_px_head = (int)(nb02 >> 4);
1102
+ const int pitch_px_n13 = (int)(nb03 >> 4);
1103
+
1104
+ const int head_px_base = i02 * pitch_px_head + (i03 / r3) * pitch_px_n13;
1105
+
1106
+ #pragma unroll
1107
+ for (int dr = 0; dr < N_K_ROWS_GQA_R4; ++dr) {
1108
+ const int r0 = r0_base + dr;
1109
+ const int row_px_base = r0 * pitch_px_row + head_px_base;
1110
+
1111
+ const int p = lane_q;
1112
+ const half8 k8 = as_half8(read_imagef(src0_img, row_px_base + p));
1113
+ const int i0 = 2 * p;
1114
+ const float4 qa = q_loc[q_id * DK_VEC_R4 + i0 ];
1115
+ const float4 qb = q_loc[q_id * DK_VEC_R4 + i0 + 1];
1116
+
1117
+ float sumf =
1118
+ convert_float(k8.s0) * qa.s0
1119
+ + convert_float(k8.s1) * qa.s1
1120
+ + convert_float(k8.s2) * qa.s2
1121
+ + convert_float(k8.s3) * qa.s3
1122
+ + convert_float(k8.s4) * qb.s0
1123
+ + convert_float(k8.s5) * qb.s1
1124
+ + convert_float(k8.s6) * qb.s2
1125
+ + convert_float(k8.s7) * qb.s3;
1126
+
1127
+ sumf += sub_group_shuffle_xor(sumf, 8);
1128
+ sumf += sub_group_shuffle_xor(sumf, 4);
1129
+ sumf += sub_group_shuffle_xor(sumf, 2);
1130
+ sumf += sub_group_shuffle_xor(sumf, 1);
1131
+
1132
+ if (lane_q == 0) {
1133
+ const int im_out = i03 * ne12 + (q_head_lo + q_id);
1134
+ dst[im_out * ne1 * ne0 + r0] = sumf;
1135
+ }
1136
+ }
1137
+ }
1138
+
1139
+ #define N_K_ROWS_GQA_R2_DK256 16
1140
+ #define GQA_RATIO_R2 2
1141
+ #define LANES_PER_QH_R2 32 // = 64 / GQA_RATIO_R2
1142
+ #define DK_VEC_DK256 64 // DK / 4 for DK=256
1143
+
1144
+ #ifdef ADRENO_GPU
1145
+ REQD_SUBGROUP_SIZE_64
1146
+ #endif
1147
+ kernel void kernel_mul_mat_f16_f32_l4_x8_gqa_r2_dk256_img(
1148
+ __read_only image1d_buffer_t src0_img,
1149
+ global char * src1,
1150
+ ulong offset1,
1151
+ global float * dst,
1152
+ ulong offsetd,
1153
+ int ne00,
1154
+ int ne01,
1155
+ int ne02,
1156
+ ulong nb01,
1157
+ ulong nb02,
1158
+ ulong nb03,
1159
+ int ne10,
1160
+ int ne11,
1161
+ int ne12,
1162
+ ulong nb10,
1163
+ ulong nb11,
1164
+ ulong nb12,
1165
+ ulong nb13,
1166
+ int ne0,
1167
+ int ne1,
1168
+ int r2,
1169
+ int r3
1170
+ ) {
1171
+ src1 = (global char *)((global char *)src1 + offset1);
1172
+ dst = (global float*)((global char *)dst + offsetd);
1173
+
1174
+ const int sgs_lid = get_sub_group_local_id();
1175
+ const int q_id = sgs_lid >> 5; // 0..1
1176
+ const int lane_q = sgs_lid & 31; // 0..31
1177
+
1178
+ const int r0_base = get_group_id(0) * N_K_ROWS_GQA_R2_DK256;
1179
+ const int im_kv = get_group_id(2);
1180
+
1181
+ const int i02 = im_kv % ne02;
1182
+ const int i03 = im_kv / ne02;
1183
+
1184
+ const int q_head_lo = i02 * GQA_RATIO_R2;
1185
+
1186
+ __local float4 q_loc[GQA_RATIO_R2 * DK_VEC_DK256];
1187
+ #pragma unroll
1188
+ for (int qh = 0; qh < GQA_RATIO_R2; ++qh) {
1189
+ const int qh_idx = q_head_lo + qh;
1190
+ global float4 * y4 = (global float4 *)(src1 + qh_idx * nb12 + i03 * nb13);
1191
+ q_loc[qh * DK_VEC_DK256 + sgs_lid] = y4[sgs_lid];
1192
+ }
1193
+ barrier(CLK_LOCAL_MEM_FENCE);
1194
+
1195
+ const int pitch_px_row = (int)(nb01 >> 4);
1196
+ const int pitch_px_head = (int)(nb02 >> 4);
1197
+ const int pitch_px_n13 = (int)(nb03 >> 4);
1198
+
1199
+ const int head_px_base = i02 * pitch_px_head + (i03 / r3) * pitch_px_n13;
1200
+
1201
+ #pragma unroll
1202
+ for (int dr = 0; dr < N_K_ROWS_GQA_R2_DK256; ++dr) {
1203
+ const int r0 = r0_base + dr;
1204
+ const int row_px_base = r0 * pitch_px_row + head_px_base;
1205
+
1206
+ const int p = lane_q;
1207
+ const half8 k8 = as_half8(read_imagef(src0_img, row_px_base + p));
1208
+ const int i0 = 2 * p;
1209
+ const float4 qa = q_loc[q_id * DK_VEC_DK256 + i0 ];
1210
+ const float4 qb = q_loc[q_id * DK_VEC_DK256 + i0 + 1];
1211
+
1212
+ float sumf =
1213
+ convert_float(k8.s0) * qa.s0
1214
+ + convert_float(k8.s1) * qa.s1
1215
+ + convert_float(k8.s2) * qa.s2
1216
+ + convert_float(k8.s3) * qa.s3
1217
+ + convert_float(k8.s4) * qb.s0
1218
+ + convert_float(k8.s5) * qb.s1
1219
+ + convert_float(k8.s6) * qb.s2
1220
+ + convert_float(k8.s7) * qb.s3;
1221
+
1222
+ sumf += sub_group_shuffle_xor(sumf, 16);
1223
+ sumf += sub_group_shuffle_xor(sumf, 8);
1224
+ sumf += sub_group_shuffle_xor(sumf, 4);
1225
+ sumf += sub_group_shuffle_xor(sumf, 2);
1226
+ sumf += sub_group_shuffle_xor(sumf, 1);
1227
+
1228
+ if (lane_q == 0) {
1229
+ const int im_out = i03 * ne12 + (q_head_lo + q_id);
1230
+ dst[im_out * ne1 * ne0 + r0] = sumf;
1231
+ }
1232
+ }
1233
+ }