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
@@ -0,0 +1,555 @@
1
+ #ifndef HMX_FA_KERNELS_H
2
+ #define HMX_FA_KERNELS_H
3
+
4
+ #include <stdint.h>
5
+ #include <stddef.h>
6
+ #include <stdbool.h>
7
+ #include "hvx-utils.h"
8
+ #include "hmx-utils.h"
9
+ #include "hex-fastdiv.h"
10
+
11
+ // HMX-specific parameters, offsets and inner kernels for Flash Attention
12
+
13
+ // Scatter offsets for diagonal tile: entry[2i] = i*136, entry[2i+1] = i*136+6
14
+ // 136 = 4 * 32 + 8 = byte offset to diagonal in a 32x32 fp16 interleaved tile
15
+ static const int16_t d_tile_scatter_offsets[64] __attribute__((aligned(128))) = {
16
+ 0 * 136, 0 * 136 + 6,
17
+ 1 * 136, 1 * 136 + 6,
18
+ 2 * 136, 2 * 136 + 6,
19
+ 3 * 136, 3 * 136 + 6,
20
+ 4 * 136, 4 * 136 + 6,
21
+ 5 * 136, 5 * 136 + 6,
22
+ 6 * 136, 6 * 136 + 6,
23
+ 7 * 136, 7 * 136 + 6,
24
+ 8 * 136, 8 * 136 + 6,
25
+ 9 * 136, 9 * 136 + 6,
26
+ 10 * 136, 10 * 136 + 6,
27
+ 11 * 136, 11 * 136 + 6,
28
+ 12 * 136, 12 * 136 + 6,
29
+ 13 * 136, 13 * 136 + 6,
30
+ 14 * 136, 14 * 136 + 6,
31
+ 15 * 136, 15 * 136 + 6,
32
+ 0, 0,
33
+ 0, 0,
34
+ 0, 0,
35
+ 0, 0,
36
+ 0, 0,
37
+ 0, 0,
38
+ 0, 0,
39
+ 0, 0,
40
+ 0, 0,
41
+ 0, 0,
42
+ 0, 0,
43
+ 0, 0,
44
+ 0, 0,
45
+ 0, 0,
46
+ 0, 0,
47
+ 0, 0,
48
+ };
49
+ // Inner HMX tile computation kernels
50
+
51
+ static void hmx_fa_qk_dot_tile(
52
+ const __fp16 * row_tiles,
53
+ const __fp16 * col_tiles,
54
+ __fp16 * out_tile,
55
+ size_t n_dot_tiles
56
+ ) {
57
+ if (n_dot_tiles == 2) {
58
+ asm volatile(
59
+ HMX_LOAD_MPY_F16("%1", "%2", "%0")
60
+ HMX_LOAD_MPY_F16("%3", "%4", "%0")
61
+ :
62
+ : "r"(2047),
63
+ "r"(row_tiles + 0 * HMX_FP16_TILE_N_ELMS), "r"(col_tiles + 0 * HMX_FP16_TILE_N_ELMS),
64
+ "r"(row_tiles + 1 * HMX_FP16_TILE_N_ELMS), "r"(col_tiles + 1 * HMX_FP16_TILE_N_ELMS)
65
+ );
66
+ } else if (n_dot_tiles == 4) {
67
+ asm volatile(
68
+ HMX_LOAD_MPY_F16("%1", "%2", "%0")
69
+ HMX_LOAD_MPY_F16("%3", "%4", "%0")
70
+ HMX_LOAD_MPY_F16("%5", "%6", "%0")
71
+ HMX_LOAD_MPY_F16("%7", "%8", "%0")
72
+ :
73
+ : "r"(2047),
74
+ "r"(row_tiles + 0 * HMX_FP16_TILE_N_ELMS), "r"(col_tiles + 0 * HMX_FP16_TILE_N_ELMS),
75
+ "r"(row_tiles + 1 * HMX_FP16_TILE_N_ELMS), "r"(col_tiles + 1 * HMX_FP16_TILE_N_ELMS),
76
+ "r"(row_tiles + 2 * HMX_FP16_TILE_N_ELMS), "r"(col_tiles + 2 * HMX_FP16_TILE_N_ELMS),
77
+ "r"(row_tiles + 3 * HMX_FP16_TILE_N_ELMS), "r"(col_tiles + 3 * HMX_FP16_TILE_N_ELMS)
78
+ );
79
+ } else if (n_dot_tiles == 8) {
80
+ asm volatile(
81
+ HMX_LOAD_MPY_F16("%1", "%2", "%0")
82
+ HMX_LOAD_MPY_F16("%3", "%4", "%0")
83
+ HMX_LOAD_MPY_F16("%5", "%6", "%0")
84
+ HMX_LOAD_MPY_F16("%7", "%8", "%0")
85
+ HMX_LOAD_MPY_F16("%9", "%10", "%0")
86
+ HMX_LOAD_MPY_F16("%11", "%12", "%0")
87
+ HMX_LOAD_MPY_F16("%13", "%14", "%0")
88
+ HMX_LOAD_MPY_F16("%15", "%16", "%0")
89
+ :
90
+ : "r"(2047),
91
+ "r"(row_tiles + 0 * HMX_FP16_TILE_N_ELMS), "r"(col_tiles + 0 * HMX_FP16_TILE_N_ELMS),
92
+ "r"(row_tiles + 1 * HMX_FP16_TILE_N_ELMS), "r"(col_tiles + 1 * HMX_FP16_TILE_N_ELMS),
93
+ "r"(row_tiles + 2 * HMX_FP16_TILE_N_ELMS), "r"(col_tiles + 2 * HMX_FP16_TILE_N_ELMS),
94
+ "r"(row_tiles + 3 * HMX_FP16_TILE_N_ELMS), "r"(col_tiles + 3 * HMX_FP16_TILE_N_ELMS),
95
+ "r"(row_tiles + 4 * HMX_FP16_TILE_N_ELMS), "r"(col_tiles + 4 * HMX_FP16_TILE_N_ELMS),
96
+ "r"(row_tiles + 5 * HMX_FP16_TILE_N_ELMS), "r"(col_tiles + 5 * HMX_FP16_TILE_N_ELMS),
97
+ "r"(row_tiles + 6 * HMX_FP16_TILE_N_ELMS), "r"(col_tiles + 6 * HMX_FP16_TILE_N_ELMS),
98
+ "r"(row_tiles + 7 * HMX_FP16_TILE_N_ELMS), "r"(col_tiles + 7 * HMX_FP16_TILE_N_ELMS)
99
+ );
100
+ } else {
101
+ for (size_t k = 0; k < n_dot_tiles; ++k) {
102
+ asm volatile(
103
+ HMX_LOAD_MPY_F16("%1", "%2", "%0")
104
+ :
105
+ : "r"(2047), "r"(row_tiles), "r"(col_tiles)
106
+ );
107
+ row_tiles += HMX_FP16_TILE_N_ELMS;
108
+ col_tiles += HMX_FP16_TILE_N_ELMS;
109
+ }
110
+ }
111
+ asm volatile(
112
+ HMX_STORE_AFTER_F16("%0", "%1")
113
+ :
114
+ : "r"(out_tile), "r"(0)
115
+ : "memory"
116
+ );
117
+ }
118
+
119
+ static void hmx_fa_o_update_tile(
120
+ const __fp16 * d_diag,
121
+ const __fp16 * o_rc,
122
+ const __fp16 * p_tile_in,
123
+ const __fp16 * v_tile_in,
124
+ __fp16 * o_tile_out,
125
+ size_t n_col_tiles
126
+ ) {
127
+ asm volatile(
128
+ HMX_LOAD_MPY_F16("%1", "%2", "%0")
129
+ :
130
+ : "r"(2047), "r"(d_diag), "r"(o_rc)
131
+ );
132
+ if (n_col_tiles == 2) {
133
+ asm volatile(
134
+ HMX_LOAD_MPY_F16("%1", "%2", "%0")
135
+ HMX_LOAD_MPY_F16("%3", "%4", "%0")
136
+ :
137
+ : "r"(2047),
138
+ "r"(p_tile_in + 0 * HMX_FP16_TILE_N_ELMS), "r"(v_tile_in + 0 * HMX_FP16_TILE_N_ELMS),
139
+ "r"(p_tile_in + 1 * HMX_FP16_TILE_N_ELMS), "r"(v_tile_in + 1 * HMX_FP16_TILE_N_ELMS)
140
+ );
141
+ } else if (n_col_tiles == 4) {
142
+ asm volatile(
143
+ HMX_LOAD_MPY_F16("%1", "%2", "%0")
144
+ HMX_LOAD_MPY_F16("%3", "%4", "%0")
145
+ HMX_LOAD_MPY_F16("%5", "%6", "%0")
146
+ HMX_LOAD_MPY_F16("%7", "%8", "%0")
147
+ :
148
+ : "r"(2047),
149
+ "r"(p_tile_in + 0 * HMX_FP16_TILE_N_ELMS), "r"(v_tile_in + 0 * HMX_FP16_TILE_N_ELMS),
150
+ "r"(p_tile_in + 1 * HMX_FP16_TILE_N_ELMS), "r"(v_tile_in + 1 * HMX_FP16_TILE_N_ELMS),
151
+ "r"(p_tile_in + 2 * HMX_FP16_TILE_N_ELMS), "r"(v_tile_in + 2 * HMX_FP16_TILE_N_ELMS),
152
+ "r"(p_tile_in + 3 * HMX_FP16_TILE_N_ELMS), "r"(v_tile_in + 3 * HMX_FP16_TILE_N_ELMS)
153
+ );
154
+ } else if (n_col_tiles == 8) {
155
+ asm volatile(
156
+ HMX_LOAD_MPY_F16("%1", "%2", "%0")
157
+ HMX_LOAD_MPY_F16("%3", "%4", "%0")
158
+ HMX_LOAD_MPY_F16("%5", "%6", "%0")
159
+ HMX_LOAD_MPY_F16("%7", "%8", "%0")
160
+ HMX_LOAD_MPY_F16("%9", "%10", "%0")
161
+ HMX_LOAD_MPY_F16("%11", "%12", "%0")
162
+ HMX_LOAD_MPY_F16("%13", "%14", "%0")
163
+ HMX_LOAD_MPY_F16("%15", "%16", "%0")
164
+ :
165
+ : "r"(2047),
166
+ "r"(p_tile_in + 0 * HMX_FP16_TILE_N_ELMS), "r"(v_tile_in + 0 * HMX_FP16_TILE_N_ELMS),
167
+ "r"(p_tile_in + 1 * HMX_FP16_TILE_N_ELMS), "r"(v_tile_in + 1 * HMX_FP16_TILE_N_ELMS),
168
+ "r"(p_tile_in + 2 * HMX_FP16_TILE_N_ELMS), "r"(v_tile_in + 2 * HMX_FP16_TILE_N_ELMS),
169
+ "r"(p_tile_in + 3 * HMX_FP16_TILE_N_ELMS), "r"(v_tile_in + 3 * HMX_FP16_TILE_N_ELMS),
170
+ "r"(p_tile_in + 4 * HMX_FP16_TILE_N_ELMS), "r"(v_tile_in + 4 * HMX_FP16_TILE_N_ELMS),
171
+ "r"(p_tile_in + 5 * HMX_FP16_TILE_N_ELMS), "r"(v_tile_in + 5 * HMX_FP16_TILE_N_ELMS),
172
+ "r"(p_tile_in + 6 * HMX_FP16_TILE_N_ELMS), "r"(v_tile_in + 6 * HMX_FP16_TILE_N_ELMS),
173
+ "r"(p_tile_in + 7 * HMX_FP16_TILE_N_ELMS), "r"(v_tile_in + 7 * HMX_FP16_TILE_N_ELMS)
174
+ );
175
+ } else {
176
+ for (size_t k = 0; k < n_col_tiles; ++k) {
177
+ asm volatile(
178
+ HMX_LOAD_MPY_F16("%1", "%2", "%0")
179
+ :
180
+ : "r"(2047), "r"(p_tile_in), "r"(v_tile_in)
181
+ );
182
+ p_tile_in += HMX_FP16_TILE_N_ELMS;
183
+ v_tile_in += HMX_FP16_TILE_N_ELMS;
184
+ }
185
+ }
186
+ asm volatile(
187
+ HMX_STORE_AFTER_F16("%0", "%1")
188
+ :
189
+ : "r"(o_tile_out), "r"(0)
190
+ : "memory"
191
+ );
192
+ }
193
+
194
+ static inline void hmx_fa_o_norm_tile(
195
+ const __fp16 * d_diag,
196
+ const __fp16 * o_rc,
197
+ __fp16 * o_out
198
+ ) {
199
+ asm volatile(
200
+ HMX_LOAD_MPY_F16("%1", "%2", "%0")
201
+ :
202
+ : "r"(2047), "r"(d_diag), "r"(o_rc)
203
+ );
204
+ asm volatile(
205
+ HMX_STORE_AFTER_F16("%0", "%1")
206
+ :
207
+ : "r"(o_out), "r"(0)
208
+ : "memory"
209
+ );
210
+ }
211
+
212
+ static inline void hmx_fa_q_prep_fp32_d2(
213
+ __fp16 * vtcm_q_tiles, const uint8_t * temp_q_vtcm,
214
+ size_t start, size_t end, size_t g_rows_end,
215
+ size_t DK, size_t G, size_t n_rows_q,
216
+ const struct fastdiv_values * div_G, bool q_transposed
217
+ ) {
218
+ for (size_t r = start; r < end; r += 2) {
219
+ size_t r0 = r / HMX_FP16_TILE_N_ROWS;
220
+ size_t r1 = r % HMX_FP16_TILE_N_ROWS;
221
+ __fp16 * out_base = vtcm_q_tiles + r0 * HMX_FP16_TILE_N_ROWS * DK;
222
+
223
+ if (r >= g_rows_end) {
224
+ ((HVX_Vector *) (out_base + 0 * HMX_FP16_TILE_N_ELMS))[r1 / 2] = Q6_V_vzero();
225
+ ((HVX_Vector *) (out_base + 1 * HMX_FP16_TILE_N_ELMS))[r1 / 2] = Q6_V_vzero();
226
+ continue;
227
+ }
228
+
229
+ const size_t q_idx0 = fastdiv(r + 0, div_G);
230
+ const size_t h_idx0 = fastmodulo(r + 0, G, div_G);
231
+ const size_t q_idx1 = fastdiv(r + 1, div_G);
232
+ const size_t h_idx1 = fastmodulo(r + 1, G, div_G);
233
+
234
+ const size_t offset0 = q_transposed ? (h_idx0 * n_rows_q + q_idx0) : (q_idx0 * G + h_idx0);
235
+ const size_t offset1 = q_transposed ? (h_idx1 * n_rows_q + q_idx1) : (q_idx1 * G + h_idx1);
236
+
237
+ const HVX_Vector * pv_in0 = (const HVX_Vector *) (temp_q_vtcm + offset0 * DK * sizeof(float));
238
+ const HVX_Vector * pv_in1 = (r + 1 < g_rows_end)
239
+ ? (const HVX_Vector *) (temp_q_vtcm + offset1 * DK * sizeof(float))
240
+ : NULL;
241
+
242
+ {
243
+ HVX_Vector v0 = pv_in0[0];
244
+ HVX_Vector v1 = pv_in1 ? pv_in1[0] : Q6_V_vzero();
245
+ HVX_Vector v_hf = hvx_vec_f32_to_f16_shuff(v0, v1);
246
+ ((HVX_Vector *) (out_base + 0 * HMX_FP16_TILE_N_ELMS))[r1 / 2] = v_hf;
247
+ }
248
+ {
249
+ HVX_Vector v0 = pv_in0[1];
250
+ HVX_Vector v1 = pv_in1 ? pv_in1[1] : Q6_V_vzero();
251
+ HVX_Vector v_hf = hvx_vec_f32_to_f16_shuff(v0, v1);
252
+ ((HVX_Vector *) (out_base + 1 * HMX_FP16_TILE_N_ELMS))[r1 / 2] = v_hf;
253
+ }
254
+ }
255
+ }
256
+
257
+ static inline void hmx_fa_q_prep_fp32_d4(
258
+ __fp16 * vtcm_q_tiles, const uint8_t * temp_q_vtcm,
259
+ size_t start, size_t end, size_t g_rows_end,
260
+ size_t DK, size_t G, size_t n_rows_q,
261
+ const struct fastdiv_values * div_G, bool q_transposed
262
+ ) {
263
+ for (size_t r = start; r < end; r += 2) {
264
+ size_t r0 = r / HMX_FP16_TILE_N_ROWS;
265
+ size_t r1 = r % HMX_FP16_TILE_N_ROWS;
266
+ __fp16 * out_base = vtcm_q_tiles + r0 * HMX_FP16_TILE_N_ROWS * DK;
267
+
268
+ if (r >= g_rows_end) {
269
+ for (uint32_t d = 0; d < 4; ++d) {
270
+ ((HVX_Vector *) (out_base + d * HMX_FP16_TILE_N_ELMS))[r1 / 2] = Q6_V_vzero();
271
+ }
272
+ continue;
273
+ }
274
+
275
+ const size_t q_idx0 = fastdiv(r + 0, div_G);
276
+ const size_t h_idx0 = fastmodulo(r + 0, G, div_G);
277
+ const size_t q_idx1 = fastdiv(r + 1, div_G);
278
+ const size_t h_idx1 = fastmodulo(r + 1, G, div_G);
279
+
280
+ const size_t offset0 = q_transposed ? (h_idx0 * n_rows_q + q_idx0) : (q_idx0 * G + h_idx0);
281
+ const size_t offset1 = q_transposed ? (h_idx1 * n_rows_q + q_idx1) : (q_idx1 * G + h_idx1);
282
+
283
+ const HVX_Vector * pv_in0 = (const HVX_Vector *) (temp_q_vtcm + offset0 * DK * sizeof(float));
284
+ const HVX_Vector * pv_in1 = (r + 1 < g_rows_end)
285
+ ? (const HVX_Vector *) (temp_q_vtcm + offset1 * DK * sizeof(float))
286
+ : NULL;
287
+
288
+ for (uint32_t d = 0; d < 4; ++d) {
289
+ HVX_Vector v0 = pv_in0[d];
290
+ HVX_Vector v1 = pv_in1 ? pv_in1[d] : Q6_V_vzero();
291
+ HVX_Vector v_hf = hvx_vec_f32_to_f16_shuff(v0, v1);
292
+ ((HVX_Vector *) (out_base + d * HMX_FP16_TILE_N_ELMS))[r1 / 2] = v_hf;
293
+ }
294
+ }
295
+ }
296
+
297
+ static inline void hmx_fa_q_prep_fp32(
298
+ __fp16 * vtcm_q_tiles, const uint8_t * temp_q_vtcm,
299
+ size_t start, size_t end, size_t g_rows_end,
300
+ size_t DK, size_t G, size_t n_rows_q,
301
+ const struct fastdiv_values * div_G, uint32_t d_limit, bool q_transposed
302
+ ) {
303
+ for (size_t r = start; r < end; r += 2) {
304
+ size_t r0 = r / HMX_FP16_TILE_N_ROWS;
305
+ size_t r1 = r % HMX_FP16_TILE_N_ROWS;
306
+ __fp16 * out_base = vtcm_q_tiles + r0 * HMX_FP16_TILE_N_ROWS * DK;
307
+
308
+ if (r >= g_rows_end) {
309
+ for (uint32_t d = 0; d < d_limit; ++d) {
310
+ ((HVX_Vector *) (out_base + d * HMX_FP16_TILE_N_ELMS))[r1 / 2] = Q6_V_vzero();
311
+ }
312
+ continue;
313
+ }
314
+
315
+ const size_t q_idx0 = fastdiv(r + 0, div_G);
316
+ const size_t h_idx0 = fastmodulo(r + 0, G, div_G);
317
+ const size_t q_idx1 = fastdiv(r + 1, div_G);
318
+ const size_t h_idx1 = fastmodulo(r + 1, G, div_G);
319
+
320
+ const size_t offset0 = q_transposed ? (h_idx0 * n_rows_q + q_idx0) : (q_idx0 * G + h_idx0);
321
+ const size_t offset1 = q_transposed ? (h_idx1 * n_rows_q + q_idx1) : (q_idx1 * G + h_idx1);
322
+
323
+ const HVX_Vector * pv_in0 = (const HVX_Vector *) (temp_q_vtcm + offset0 * DK * sizeof(float));
324
+ const HVX_Vector * pv_in1 = (r + 1 < g_rows_end)
325
+ ? (const HVX_Vector *) (temp_q_vtcm + offset1 * DK * sizeof(float))
326
+ : NULL;
327
+
328
+ for (uint32_t d = 0; d < d_limit; ++d) {
329
+ HVX_Vector v0 = pv_in0[d];
330
+ HVX_Vector v1 = pv_in1 ? pv_in1[d] : Q6_V_vzero();
331
+ HVX_Vector v_hf = hvx_vec_f32_to_f16_shuff(v0, v1);
332
+
333
+ HVX_Vector * out_tile = (HVX_Vector *) (out_base + d * HMX_FP16_TILE_N_ELMS);
334
+ out_tile[r1 / 2] = v_hf;
335
+ }
336
+ }
337
+ }
338
+
339
+ static inline void hmx_fa_q_prep_fp16_d1(
340
+ __fp16 * vtcm_q_tiles, const uint8_t * temp_q_vtcm,
341
+ size_t start, size_t end, size_t g_rows_end,
342
+ size_t DK, size_t G, size_t n_rows_q,
343
+ const struct fastdiv_values * div_G, bool q_transposed
344
+ ) {
345
+ for (size_t r = start; r < end; r += 2) {
346
+ size_t r0 = r / HMX_FP16_TILE_N_ROWS;
347
+ size_t r1 = r % HMX_FP16_TILE_N_ROWS;
348
+ __fp16 * out_base = vtcm_q_tiles + r0 * HMX_FP16_TILE_N_ROWS * DK;
349
+
350
+ if (r >= g_rows_end) {
351
+ __fp16 * out_dtile = out_base + 0 * HMX_FP16_TILE_N_ELMS * 2;
352
+ HVX_Vector * pv_out0 = ((HVX_Vector *) out_dtile) + r1 / 2;
353
+ HVX_Vector * pv_out1 = pv_out0 + 16;
354
+ *pv_out0 = Q6_V_vzero();
355
+ *pv_out1 = Q6_V_vzero();
356
+ continue;
357
+ }
358
+
359
+ const size_t q_idx0 = fastdiv(r + 0, div_G);
360
+ const size_t h_idx0 = fastmodulo(r + 0, G, div_G);
361
+ const size_t q_idx1 = fastdiv(r + 1, div_G);
362
+ const size_t h_idx1 = fastmodulo(r + 1, G, div_G);
363
+
364
+ const size_t offset0 = q_transposed ? (h_idx0 * n_rows_q + q_idx0) : (q_idx0 * G + h_idx0);
365
+ const size_t offset1 = q_transposed ? (h_idx1 * n_rows_q + q_idx1) : (q_idx1 * G + h_idx1);
366
+
367
+ const HVX_Vector * pv_in0 = (const HVX_Vector *) (temp_q_vtcm + offset0 * DK * sizeof(__fp16));
368
+ const HVX_Vector * pv_in1 = (r + 1 < g_rows_end)
369
+ ? (const HVX_Vector *) (temp_q_vtcm + offset1 * DK * sizeof(__fp16))
370
+ : NULL;
371
+
372
+ HVX_Vector v0 = pv_in0[0];
373
+ HVX_Vector v1 = pv_in1 ? pv_in1[0] : Q6_V_vzero();
374
+ HVX_VectorPair vp = Q6_W_vshuff_VVR(v1, v0, -2);
375
+
376
+ __fp16 * out_dtile = out_base + 0 * HMX_FP16_TILE_N_ELMS * 2;
377
+ HVX_Vector * pv_out0 = ((HVX_Vector *) out_dtile) + r1 / 2;
378
+ HVX_Vector * pv_out1 = pv_out0 + 16;
379
+
380
+ *pv_out0 = Q6_V_lo_W(vp);
381
+ *pv_out1 = Q6_V_hi_W(vp);
382
+ }
383
+ }
384
+
385
+ static inline void hmx_fa_q_prep_fp16_d2(
386
+ __fp16 * vtcm_q_tiles, const uint8_t * temp_q_vtcm,
387
+ size_t start, size_t end, size_t g_rows_end,
388
+ size_t DK, size_t G, size_t n_rows_q,
389
+ const struct fastdiv_values * div_G, bool q_transposed
390
+ ) {
391
+ for (size_t r = start; r < end; r += 2) {
392
+ size_t r0 = r / HMX_FP16_TILE_N_ROWS;
393
+ size_t r1 = r % HMX_FP16_TILE_N_ROWS;
394
+ __fp16 * out_base = vtcm_q_tiles + r0 * HMX_FP16_TILE_N_ROWS * DK;
395
+
396
+ if (r >= g_rows_end) {
397
+ for (uint32_t d = 0; d < 2; ++d) {
398
+ __fp16 * out_dtile = out_base + d * HMX_FP16_TILE_N_ELMS * 2;
399
+ HVX_Vector * pv_out0 = ((HVX_Vector *) out_dtile) + r1 / 2;
400
+ HVX_Vector * pv_out1 = pv_out0 + 16;
401
+ *pv_out0 = Q6_V_vzero();
402
+ *pv_out1 = Q6_V_vzero();
403
+ }
404
+ continue;
405
+ }
406
+
407
+ const size_t q_idx0 = fastdiv(r + 0, div_G);
408
+ const size_t h_idx0 = fastmodulo(r + 0, G, div_G);
409
+ const size_t q_idx1 = fastdiv(r + 1, div_G);
410
+ const size_t h_idx1 = fastmodulo(r + 1, G, div_G);
411
+
412
+ const size_t offset0 = q_transposed ? (h_idx0 * n_rows_q + q_idx0) : (q_idx0 * G + h_idx0);
413
+ const size_t offset1 = q_transposed ? (h_idx1 * n_rows_q + q_idx1) : (q_idx1 * G + h_idx1);
414
+
415
+ const HVX_Vector * pv_in0 = (const HVX_Vector *) (temp_q_vtcm + offset0 * DK * sizeof(__fp16));
416
+ const HVX_Vector * pv_in1 = (r + 1 < g_rows_end)
417
+ ? (const HVX_Vector *) (temp_q_vtcm + offset1 * DK * sizeof(__fp16))
418
+ : NULL;
419
+
420
+ {
421
+ HVX_Vector v0 = pv_in0[0];
422
+ HVX_Vector v1 = pv_in1 ? pv_in1[0] : Q6_V_vzero();
423
+ HVX_VectorPair vp = Q6_W_vshuff_VVR(v1, v0, -2);
424
+
425
+ __fp16 * out_dtile = out_base + 0 * HMX_FP16_TILE_N_ELMS * 2;
426
+ HVX_Vector * pv_out0 = ((HVX_Vector *) out_dtile) + r1 / 2;
427
+ HVX_Vector * pv_out1 = pv_out0 + 16;
428
+
429
+ *pv_out0 = Q6_V_lo_W(vp);
430
+ *pv_out1 = Q6_V_hi_W(vp);
431
+ }
432
+ {
433
+ HVX_Vector v0 = pv_in0[1];
434
+ HVX_Vector v1 = pv_in1 ? pv_in1[1] : Q6_V_vzero();
435
+ HVX_VectorPair vp = Q6_W_vshuff_VVR(v1, v0, -2);
436
+
437
+ __fp16 * out_dtile = out_base + 1 * HMX_FP16_TILE_N_ELMS * 2;
438
+ HVX_Vector * pv_out0 = ((HVX_Vector *) out_dtile) + r1 / 2;
439
+ HVX_Vector * pv_out1 = pv_out0 + 16;
440
+
441
+ *pv_out0 = Q6_V_lo_W(vp);
442
+ *pv_out1 = Q6_V_hi_W(vp);
443
+ }
444
+ }
445
+ }
446
+
447
+ static inline void hmx_fa_q_prep_fp16(
448
+ __fp16 * vtcm_q_tiles, const uint8_t * temp_q_vtcm,
449
+ size_t start, size_t end, size_t g_rows_end,
450
+ size_t DK, size_t G, size_t n_rows_q,
451
+ const struct fastdiv_values * div_G, uint32_t d_limit, bool q_transposed
452
+ ) {
453
+ for (size_t r = start; r < end; r += 2) {
454
+ size_t r0 = r / HMX_FP16_TILE_N_ROWS;
455
+ size_t r1 = r % HMX_FP16_TILE_N_ROWS;
456
+ __fp16 * out_base = vtcm_q_tiles + r0 * HMX_FP16_TILE_N_ROWS * DK;
457
+
458
+ if (r >= g_rows_end) {
459
+ for (uint32_t d = 0; d < d_limit; ++d) {
460
+ __fp16 * out_dtile = out_base + d * HMX_FP16_TILE_N_ELMS * 2;
461
+ HVX_Vector * pv_out0 = ((HVX_Vector *) out_dtile) + r1 / 2;
462
+ HVX_Vector * pv_out1 = pv_out0 + 16;
463
+ *pv_out0 = Q6_V_vzero();
464
+ *pv_out1 = Q6_V_vzero();
465
+ }
466
+ continue;
467
+ }
468
+
469
+ const size_t q_idx0 = fastdiv(r + 0, div_G);
470
+ const size_t h_idx0 = fastmodulo(r + 0, G, div_G);
471
+ const size_t q_idx1 = fastdiv(r + 1, div_G);
472
+ const size_t h_idx1 = fastmodulo(r + 1, G, div_G);
473
+
474
+ const size_t offset0 = q_transposed ? (h_idx0 * n_rows_q + q_idx0) : (q_idx0 * G + h_idx0);
475
+ const size_t offset1 = q_transposed ? (h_idx1 * n_rows_q + q_idx1) : (q_idx1 * G + h_idx1);
476
+
477
+ const HVX_Vector * pv_in0 = (const HVX_Vector *) (temp_q_vtcm + offset0 * DK * sizeof(__fp16));
478
+ const HVX_Vector * pv_in1 = (r + 1 < g_rows_end)
479
+ ? (const HVX_Vector *) (temp_q_vtcm + offset1 * DK * sizeof(__fp16))
480
+ : NULL;
481
+
482
+ for (uint32_t d = 0; d < d_limit; ++d) {
483
+ HVX_Vector v0 = pv_in0[d];
484
+ HVX_Vector v1 = pv_in1 ? pv_in1[d] : Q6_V_vzero();
485
+ HVX_VectorPair vp = Q6_W_vshuff_VVR(v1, v0, -2);
486
+
487
+ __fp16 * out_dtile = out_base + d * HMX_FP16_TILE_N_ELMS * 2;
488
+ HVX_Vector * pv_out0 = ((HVX_Vector *) out_dtile) + r1 / 2;
489
+ HVX_Vector * pv_out1 = pv_out0 + 16;
490
+
491
+ *pv_out0 = Q6_V_lo_W(vp);
492
+ *pv_out1 = Q6_V_hi_W(vp);
493
+ }
494
+ }
495
+ }
496
+
497
+
498
+ static inline void hmx_fa_q_prep_fallback(
499
+ __fp16 * vtcm_q_tiles, uintptr_t q_data,
500
+ size_t q_nb1, size_t q_nb2, size_t q_nb3,
501
+ uint32_t q_start, uint32_t kv_head, uint32_t ib3,
502
+ size_t start, size_t end, size_t n_rows_g,
503
+ size_t G, size_t DK, bool is_q_fp32,
504
+ const struct fastdiv_values * div_G
505
+ ) {
506
+ for (size_t r = start; r < end; r += 2) {
507
+ const size_t q_idx0 = fastdiv(r + 0, div_G);
508
+ const size_t h_idx0 = fastmodulo(r + 0, G, div_G);
509
+ const size_t q_idx1 = fastdiv(r + 1, div_G);
510
+ const size_t h_idx1 = fastmodulo(r + 1, G, div_G);
511
+
512
+ const uint8_t * q_ptr0 = (r + 0 < n_rows_g) ? ((const uint8_t *) q_data + (q_start + q_idx0) * q_nb1 +
513
+ (kv_head * G + h_idx0) * q_nb2 + ib3 * q_nb3) :
514
+ NULL;
515
+ const uint8_t * q_ptr1 = (r + 1 < n_rows_g) ? ((const uint8_t *) q_data + (q_start + q_idx1) * q_nb1 +
516
+ (kv_head * G + h_idx1) * q_nb2 + ib3 * q_nb3) :
517
+ NULL;
518
+
519
+ size_t r0 = r / HMX_FP16_TILE_N_ROWS;
520
+ size_t r1 = r % HMX_FP16_TILE_N_ROWS;
521
+ __fp16 * out_base = vtcm_q_tiles + r0 * HMX_FP16_TILE_N_ROWS * DK;
522
+
523
+ if (is_q_fp32) {
524
+ const HVX_UVector * pv_in0 = q_ptr0 ? (const HVX_UVector *) q_ptr0 : NULL;
525
+ const HVX_UVector * pv_in1 = q_ptr1 ? (const HVX_UVector *) q_ptr1 : NULL;
526
+
527
+ for (uint32_t d = 0; d < DK / 32; ++d) {
528
+ HVX_Vector v0 = pv_in0 ? pv_in0[d] : Q6_V_vzero();
529
+ HVX_Vector v1 = pv_in1 ? pv_in1[d] : Q6_V_vzero();
530
+ HVX_Vector v_hf = hvx_vec_f32_to_f16_shuff(v0, v1);
531
+
532
+ HVX_Vector * out_tile = (HVX_Vector *) (out_base + d * HMX_FP16_TILE_N_ELMS);
533
+ out_tile[r1 / 2] = v_hf;
534
+ }
535
+ } else {
536
+ const HVX_UVector * pv_in0 = q_ptr0 ? (const HVX_UVector *) q_ptr0 : NULL;
537
+ const HVX_UVector * pv_in1 = q_ptr1 ? (const HVX_UVector *) q_ptr1 : NULL;
538
+
539
+ for (uint32_t d = 0; d < DK / 64; ++d) {
540
+ HVX_Vector v0 = pv_in0 ? pv_in0[d] : Q6_V_vzero();
541
+ HVX_Vector v1 = pv_in1 ? pv_in1[d] : Q6_V_vzero();
542
+ HVX_VectorPair vp = Q6_W_vshuff_VVR(v1, v0, -2);
543
+
544
+ __fp16 * out_dtile = out_base + d * HMX_FP16_TILE_N_ELMS * 2;
545
+ HVX_Vector * pv_out0 = ((HVX_Vector *) out_dtile) + r1 / 2;
546
+ HVX_Vector * pv_out1 = pv_out0 + 16;
547
+
548
+ *pv_out0 = Q6_V_lo_W(vp);
549
+ *pv_out1 = Q6_V_hi_W(vp);
550
+ }
551
+ }
552
+ }
553
+ }
554
+
555
+ #endif /* HMX_FA_KERNELS_H */