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
@@ -27,6 +27,8 @@
27
27
  #define QR5_1 2
28
28
  #define QK8_0 32
29
29
  #define QR8_0 1
30
+ #define QK1_0 128
31
+ #define QR1_0 1
30
32
  #define QK_K 256
31
33
  #define K_SCALE_SIZE (3 * QK_K / 64)
32
34
  #define K_QUANTS_PER_ITERATION 2
@@ -38,6 +40,14 @@ typedef ushort uint16_t;
38
40
  typedef int int32_t;
39
41
  typedef uint uint32_t;
40
42
 
43
+ //------------------------------------------------------------------------------
44
+ // block_q1_0
45
+ //------------------------------------------------------------------------------
46
+ typedef struct {
47
+ half d; // delta
48
+ uchar qs[QK1_0/8]; // 1-bit signs (16 bytes)
49
+ } block_q1_0;
50
+
41
51
  //------------------------------------------------------------------------------
42
52
  // block_q4_0
43
53
  //------------------------------------------------------------------------------
@@ -159,6 +169,42 @@ kernel void kernel_convert_f16_to_bf16(
159
169
  }
160
170
  }
161
171
 
172
+ //------------------------------------------------------------------------------
173
+ // kernel_convert_block_q1_0
174
+ // Convert block_q1_0 (AOS) to 2 separate arrays (SOA): quant bytes + scales.
175
+ // q1_0 bits are stored in natural order (bit j of byte i -> weight 8*i + j)
176
+ //------------------------------------------------------------------------------
177
+ kernel void kernel_convert_block_q1_0(
178
+ global block_q1_0 * src0,
179
+ global uchar * dst_q,
180
+ global half * dst_d
181
+ ) {
182
+ global block_q1_0 * b = (global block_q1_0 *) src0 + get_global_id(0);
183
+ global uchar * q = (global uchar *) dst_q + (QK1_0/8)*get_global_id(0);
184
+ global half * d = (global half *) dst_d + get_global_id(0);
185
+
186
+ *d = b->d;
187
+
188
+ for (int i = 0; i < QK1_0/8; ++i) {
189
+ q[i] = b->qs[i];
190
+ }
191
+ }
192
+
193
+ kernel void kernel_restore_block_q1_0(
194
+ global uchar * src_q,
195
+ global half * src_d,
196
+ global block_q1_0 * dst
197
+ ) {
198
+ global block_q1_0 * b = (global block_q1_0 *) dst + get_global_id(0);
199
+ global uchar * q = (global uchar *) src_q + (QK1_0/8)*get_global_id(0);
200
+ global half * d = (global half *) src_d + get_global_id(0);
201
+
202
+ b->d = *d;
203
+ for (int i = 0; i < QK1_0/8; ++i) {
204
+ b->qs[i] = q[i];
205
+ }
206
+ }
207
+
162
208
  //------------------------------------------------------------------------------
163
209
  // kernel_convert_block_q4_0
164
210
  // Convert the block_q4_0 format to 2 separate arrays (AOS -> SOA).
@@ -1582,6 +1628,158 @@ kernel void kernel_restore_block_q8_0(
1582
1628
  }
1583
1629
  }
1584
1630
 
1631
+ // View-aware AoS q8_0 -> f32 dequant (f32/f32 FA path).
1632
+ kernel void kernel_dequant_q8_0_f32_view_aos(
1633
+ global char * src,
1634
+ ulong src_offset,
1635
+ ulong src_nb1,
1636
+ ulong src_nb2,
1637
+ ulong src_nb3,
1638
+ int nblk0,
1639
+ int ne1,
1640
+ int ne2,
1641
+ int ne3,
1642
+ global float * dst
1643
+ ) {
1644
+ int blk_i0 = get_global_id(0);
1645
+ int i1 = get_global_id(1);
1646
+ int batch = get_global_id(2);
1647
+
1648
+ if (blk_i0 >= nblk0) return;
1649
+ if (i1 >= ne1) return;
1650
+
1651
+ int i2 = batch % ne2;
1652
+ int i3 = batch / ne2;
1653
+ if (i3 >= ne3) return;
1654
+
1655
+ global char * block = src + src_offset + (ulong)i3*src_nb3 + (ulong)i2*src_nb2 + (ulong)i1*src_nb1 + (ulong)blk_i0 * (2 + QK8_0);
1656
+ float d = vload_half(0, (global half *)block);
1657
+ global char * qs = block + 2;
1658
+
1659
+ ulong dst_row_base = ((ulong)i3 * ne2 * ne1 + (ulong)i2 * ne1 + (ulong)i1) * nblk0;
1660
+ global float * out = dst + (dst_row_base + blk_i0) * QK8_0;
1661
+
1662
+ for (int i = 0; i < QK8_0; ++i) {
1663
+ out[i] = d * (float)qs[i];
1664
+ }
1665
+ }
1666
+
1667
+ // View-aware AoS q8_0 -> f16 dequant. Rows tight, batch strides may be gapped.
1668
+ kernel void kernel_dequant_q8_0_f16_view_aos(
1669
+ global char * src,
1670
+ ulong src_offset,
1671
+ ulong src_nb1,
1672
+ ulong src_nb2,
1673
+ ulong src_nb3,
1674
+ int nblk0,
1675
+ int ne1,
1676
+ int ne2,
1677
+ int ne3,
1678
+ global half * dst
1679
+ ) {
1680
+ int blk_i0 = get_global_id(0);
1681
+ int i1 = get_global_id(1);
1682
+ int batch = get_global_id(2);
1683
+
1684
+ if (blk_i0 >= nblk0) return;
1685
+ if (i1 >= ne1) return;
1686
+
1687
+ int i2 = batch % ne2;
1688
+ int i3 = batch / ne2;
1689
+ if (i3 >= ne3) return;
1690
+
1691
+ global char * block = src + src_offset + (ulong)i3*src_nb3 + (ulong)i2*src_nb2 + (ulong)i1*src_nb1 + (ulong)blk_i0 * (2 + QK8_0);
1692
+ float d = vload_half(0, (global half *)block);
1693
+ global char * qs = block + 2;
1694
+
1695
+ ulong dst_row_base = ((ulong)i3 * ne2 * ne1 + (ulong)i2 * ne1 + (ulong)i1) * nblk0;
1696
+ global half * out = dst + (dst_row_base + blk_i0) * QK8_0;
1697
+
1698
+ for (int i = 0; i < QK8_0; ++i) {
1699
+ out[i] = (half)(d * (float)qs[i]);
1700
+ }
1701
+ }
1702
+
1703
+ // View-aware AoS q4_0 -> f32 dequant (mirrors the q8_0 view variant).
1704
+ kernel void kernel_dequant_q4_0_f32_view_aos(
1705
+ global char * src,
1706
+ ulong src_offset,
1707
+ ulong src_nb1,
1708
+ ulong src_nb2,
1709
+ ulong src_nb3,
1710
+ int nblk0,
1711
+ int ne1,
1712
+ int ne2,
1713
+ int ne3,
1714
+ global float * dst
1715
+ ) {
1716
+ int blk_i0 = get_global_id(0);
1717
+ int i1 = get_global_id(1);
1718
+ int batch = get_global_id(2);
1719
+
1720
+ if (blk_i0 >= nblk0) return;
1721
+ if (i1 >= ne1) return;
1722
+
1723
+ int i2 = batch % ne2;
1724
+ int i3 = batch / ne2;
1725
+ if (i3 >= ne3) return;
1726
+
1727
+ global char * block = src + src_offset + (ulong)i3*src_nb3 + (ulong)i2*src_nb2 + (ulong)i1*src_nb1 + (ulong)blk_i0 * (2 + QK4_0/2);
1728
+ float d = vload_half(0, (global half *)block);
1729
+ global uchar * qs = (global uchar *)(block + 2);
1730
+
1731
+ ulong dst_row_base = ((ulong)i3 * ne2 * ne1 + (ulong)i2 * ne1 + (ulong)i1) * nblk0;
1732
+ global float * out = dst + (dst_row_base + blk_i0) * QK4_0;
1733
+
1734
+ for (int i = 0; i < QK4_0/2; ++i) {
1735
+ uchar byte = qs[i];
1736
+ int q0 = (int)(byte & 0x0F) - 8;
1737
+ int q1 = (int)(byte >> 4) - 8;
1738
+ out[i] = d * (float)q0;
1739
+ out[i + QK4_0/2] = d * (float)q1;
1740
+ }
1741
+ }
1742
+
1743
+ // View-aware AoS q4_0 -> f16 dequant (mirrors the q8_0 view variant).
1744
+ kernel void kernel_dequant_q4_0_f16_view_aos(
1745
+ global char * src,
1746
+ ulong src_offset,
1747
+ ulong src_nb1,
1748
+ ulong src_nb2,
1749
+ ulong src_nb3,
1750
+ int nblk0,
1751
+ int ne1,
1752
+ int ne2,
1753
+ int ne3,
1754
+ global half * dst
1755
+ ) {
1756
+ int blk_i0 = get_global_id(0);
1757
+ int i1 = get_global_id(1);
1758
+ int batch = get_global_id(2);
1759
+
1760
+ if (blk_i0 >= nblk0) return;
1761
+ if (i1 >= ne1) return;
1762
+
1763
+ int i2 = batch % ne2;
1764
+ int i3 = batch / ne2;
1765
+ if (i3 >= ne3) return;
1766
+
1767
+ global char * block = src + src_offset + (ulong)i3*src_nb3 + (ulong)i2*src_nb2 + (ulong)i1*src_nb1 + (ulong)blk_i0 * (2 + QK4_0/2);
1768
+ float d = vload_half(0, (global half *)block);
1769
+ global uchar * qs = (global uchar *)(block + 2);
1770
+
1771
+ ulong dst_row_base = ((ulong)i3 * ne2 * ne1 + (ulong)i2 * ne1 + (ulong)i1) * nblk0;
1772
+ global half * out = dst + (dst_row_base + blk_i0) * QK4_0;
1773
+
1774
+ for (int i = 0; i < QK4_0/2; ++i) {
1775
+ uchar byte = qs[i];
1776
+ int q0 = (int)(byte & 0x0F) - 8;
1777
+ int q1 = (int)(byte >> 4) - 8;
1778
+ out[i] = (half)(d * (float)q0);
1779
+ out[i + QK4_0/2] = (half)(d * (float)q1);
1780
+ }
1781
+ }
1782
+
1585
1783
  kernel void kernel_restore_block_q8_0_trans(
1586
1784
  global uchar * src_q,
1587
1785
  global half * src_d,
@@ -4,13 +4,30 @@
4
4
  #define ACC_TYPE4 float4
5
5
  #define DATA_TYPE half
6
6
  #define DATA_TYPE4 half4
7
- #define CONVERT_ACC4(x) convert_float4(x)
8
- #define CONVERT_DATA4(x) convert_half4(x)
7
+ #define CONVERT_ACC4(x) ((float4)((float)(x).s0, (float)(x).s1, (float)(x).s2, (float)(x).s3))
8
+ #define CONVERT_DATA4(x) ((half4)((half)(x).s0, (half)(x).s1, (half)(x).s2, (half)(x).s3))
9
9
 
10
10
  #define DK_VEC (DK/4)
11
11
  #define DV_VEC (DV/4)
12
12
  #define WG_SIZE (BLOCK_M)
13
- #define Q1_WG_SIZE 64
13
+ // q1 reduces over a Q1_WG_SIZE-wide WG via work-group barriers; the launch WG
14
+ // must match. Defaults to the Adreno sg (64); host passes -D FA_SG=32 on Intel.
15
+ #ifndef FA_SG
16
+ #define FA_SG 64
17
+ #endif
18
+ #define Q1_WG_SIZE FA_SG
19
+
20
+ // The kernels are built with -cl-finite-math-only. On some older Adreno GPUs,
21
+ // infinite operand can cause undefined behavior and miscompilation for exp.
22
+ // Therefore, a large negative value is used instead.
23
+ #define FA_M_INIT (-3.0e38f)
24
+
25
+ // Drop full unroll at DK>=192 — Adreno compiler host-memory budget.
26
+ #if DK >= 192
27
+ #define FA_UNROLL
28
+ #else
29
+ #define FA_UNROLL _Pragma("unroll")
30
+ #endif
14
31
 
15
32
  inline float get_alibi_slope(
16
33
  const float max_bias, const uint h, const uint n_head_log2, const float m0, const float m1
@@ -81,18 +98,18 @@ __kernel void flash_attn_f16(
81
98
  if (my_query_row < n_q) {
82
99
  const ulong q_row_offset = batch_idx * q_nb3 + head_idx * q_nb2 + my_query_row * q_nb1;
83
100
  const global DATA_TYPE4* q_ptr = (const global DATA_TYPE4*)(q_base + q_row_offset);
84
- #pragma unroll
101
+ FA_UNROLL
85
102
  for (int i = 0; i < DK_VEC; ++i) {
86
103
  q_priv[i] = CONVERT_ACC4(q_ptr[i]);
87
104
  }
88
105
  }
89
106
 
90
107
  ACC_TYPE4 o_acc[DV_VEC];
91
- #pragma unroll
108
+ FA_UNROLL
92
109
  for (int i = 0; i < DV_VEC; ++i) {
93
110
  o_acc[i] = (ACC_TYPE4)(0.0f);
94
111
  }
95
- ACC_TYPE m_i = -INFINITY;
112
+ ACC_TYPE m_i = FA_M_INIT;
96
113
  ACC_TYPE l_i = 0.0f;
97
114
 
98
115
  float slope = get_alibi_slope(max_bias, head_idx, n_head_log2, m0, m1);
@@ -125,49 +142,72 @@ __kernel void flash_attn_f16(
125
142
  continue;
126
143
  }
127
144
 
128
- for (int j = 0; j < BLOCK_N; j += 2) {
145
+ for (int j = 0; j < BLOCK_N; j += 4) {
129
146
  const int k_row0 = k_start + j;
130
147
  const int k_row1 = k_start + j + 1;
148
+ const int k_row2 = k_start + j + 2;
149
+ const int k_row3 = k_start + j + 3;
131
150
 
132
151
  ACC_TYPE4 dot_acc0 = (ACC_TYPE4)(0.0f);
133
152
  ACC_TYPE4 dot_acc1 = (ACC_TYPE4)(0.0f);
134
- #pragma unroll
153
+ ACC_TYPE4 dot_acc2 = (ACC_TYPE4)(0.0f);
154
+ ACC_TYPE4 dot_acc3 = (ACC_TYPE4)(0.0f);
155
+ FA_UNROLL
135
156
  for (int k = 0; k < DK_VEC; k++) {
136
- dot_acc0 = mad(q_priv[k], CONVERT_ACC4(l_k[j][k]), dot_acc0);
137
- dot_acc1 = mad(q_priv[k], CONVERT_ACC4(l_k[j+1][k]), dot_acc1);
157
+ const ACC_TYPE4 qk = q_priv[k];
158
+ dot_acc0 = mad(qk, CONVERT_ACC4(l_k[j][k]), dot_acc0);
159
+ dot_acc1 = mad(qk, CONVERT_ACC4(l_k[j+1][k]), dot_acc1);
160
+ dot_acc2 = mad(qk, CONVERT_ACC4(l_k[j+2][k]), dot_acc2);
161
+ dot_acc3 = mad(qk, CONVERT_ACC4(l_k[j+3][k]), dot_acc3);
138
162
  }
139
- ACC_TYPE score0 = (dot_acc0.s0 + dot_acc0.s1 + dot_acc0.s2 + dot_acc0.s3) * scale;
140
- ACC_TYPE score1 = (dot_acc1.s0 + dot_acc1.s1 + dot_acc1.s2 + dot_acc1.s3) * scale;
163
+ ACC_TYPE s0 = (dot_acc0.s0 + dot_acc0.s1 + dot_acc0.s2 + dot_acc0.s3) * scale;
164
+ ACC_TYPE s1 = (dot_acc1.s0 + dot_acc1.s1 + dot_acc1.s2 + dot_acc1.s3) * scale;
165
+ ACC_TYPE s2 = (dot_acc2.s0 + dot_acc2.s1 + dot_acc2.s2 + dot_acc2.s3) * scale;
166
+ ACC_TYPE s3 = (dot_acc3.s0 + dot_acc3.s1 + dot_acc3.s2 + dot_acc3.s3) * scale;
141
167
 
142
168
  if (is_causal) {
143
- if (k_row0 > (n_kv - n_q + my_query_row)) score0 = -INFINITY;
144
- if (k_row1 > (n_kv - n_q + my_query_row)) score1 = -INFINITY;
169
+ const int causal_limit = n_kv - n_q + my_query_row;
170
+ if (k_row0 > causal_limit) s0 = FA_M_INIT;
171
+ if (k_row1 > causal_limit) s1 = FA_M_INIT;
172
+ if (k_row2 > causal_limit) s2 = FA_M_INIT;
173
+ if (k_row3 > causal_limit) s3 = FA_M_INIT;
145
174
  }
146
-
147
- if (k_row0 >= n_kv) score0 = -INFINITY;
148
- if (k_row1 >= n_kv) score1 = -INFINITY;
175
+ if (k_row0 >= n_kv) s0 = FA_M_INIT;
176
+ if (k_row1 >= n_kv) s1 = FA_M_INIT;
177
+ if (k_row2 >= n_kv) s2 = FA_M_INIT;
178
+ if (k_row3 >= n_kv) s3 = FA_M_INIT;
149
179
 
150
180
  if (mask_base != NULL) {
151
181
  const global DATA_TYPE* mask_ptr = (const global DATA_TYPE*)(mask_base + my_query_row * mask_nb1);
152
- if (k_row0 < n_kv) score0 += slope * (ACC_TYPE)mask_ptr[k_row0];
153
- if (k_row1 < n_kv) score1 += slope * (ACC_TYPE)mask_ptr[k_row1];
182
+ if (k_row0 < n_kv) s0 += slope * (ACC_TYPE)mask_ptr[k_row0];
183
+ if (k_row1 < n_kv) s1 += slope * (ACC_TYPE)mask_ptr[k_row1];
184
+ if (k_row2 < n_kv) s2 += slope * (ACC_TYPE)mask_ptr[k_row2];
185
+ if (k_row3 < n_kv) s3 += slope * (ACC_TYPE)mask_ptr[k_row3];
154
186
  }
155
187
 
156
188
  if (logit_softcap > 0.0f) {
157
- score0 = logit_softcap * tanh(score0 / logit_softcap);
158
- score1 = logit_softcap * tanh(score1 / logit_softcap);
189
+ s0 = logit_softcap * tanh(s0 / logit_softcap);
190
+ s1 = logit_softcap * tanh(s1 / logit_softcap);
191
+ s2 = logit_softcap * tanh(s2 / logit_softcap);
192
+ s3 = logit_softcap * tanh(s3 / logit_softcap);
159
193
  }
160
194
 
161
- const ACC_TYPE m_new = max(m_i, max(score0, score1));
162
- const ACC_TYPE p0 = exp(score0 - m_new);
163
- const ACC_TYPE p1 = exp(score1 - m_new);
164
- const ACC_TYPE scale_prev = exp(m_i - m_new);
195
+ const ACC_TYPE m_new = max(m_i, max(max(s0, s1), max(s2, s3)));
196
+ const ACC_TYPE scale_prev = native_exp(m_i - m_new);
197
+ const ACC_TYPE p0 = native_exp(s0 - m_new);
198
+ const ACC_TYPE p1 = native_exp(s1 - m_new);
199
+ const ACC_TYPE p2 = native_exp(s2 - m_new);
200
+ const ACC_TYPE p3 = native_exp(s3 - m_new);
165
201
 
166
- #pragma unroll
202
+ FA_UNROLL
167
203
  for (int i = 0; i < DV_VEC; ++i) {
168
- o_acc[i] = o_acc[i] * scale_prev + p0 * CONVERT_ACC4(l_v[j][i]) + p1 * CONVERT_ACC4(l_v[j+1][i]);
204
+ o_acc[i] = mad(p3, CONVERT_ACC4(l_v[j+3][i]),
205
+ mad(p2, CONVERT_ACC4(l_v[j+2][i]),
206
+ mad(p1, CONVERT_ACC4(l_v[j+1][i]),
207
+ mad(p0, CONVERT_ACC4(l_v[j][i]),
208
+ o_acc[i] * scale_prev))));
169
209
  }
170
- l_i = l_i * scale_prev + p0 + p1;
210
+ l_i = l_i * scale_prev + p0 + p1 + p2 + p3;
171
211
  m_i = m_new;
172
212
  }
173
213
  }
@@ -179,7 +219,7 @@ __kernel void flash_attn_f16(
179
219
  const ACC_TYPE m_final = max(m_i, m_sink);
180
220
 
181
221
  const ACC_TYPE scale_o = exp(m_i - m_final);
182
- #pragma unroll
222
+ FA_UNROLL
183
223
  for (int i = 0; i < DV_VEC; ++i) {
184
224
  o_acc[i] *= scale_o;
185
225
  }
@@ -191,12 +231,12 @@ __kernel void flash_attn_f16(
191
231
  global DATA_TYPE4 *o_row = (global DATA_TYPE4 *)(o_base + o_row_offset);
192
232
  if (l_i > 0.0f) {
193
233
  const ACC_TYPE l_inv = 1.0f / l_i;
194
- #pragma unroll
234
+ FA_UNROLL
195
235
  for (int i = 0; i < DV_VEC; ++i) {
196
236
  o_row[i] = CONVERT_DATA4(o_acc[i] * l_inv);
197
237
  }
198
238
  } else {
199
- #pragma unroll
239
+ FA_UNROLL
200
240
  for (int i = 0; i < DV_VEC; ++i) {
201
241
  o_row[i] = (DATA_TYPE4)(0.0f);
202
242
  }
@@ -258,7 +298,7 @@ __kernel void flash_attn_f16_q1(
258
298
  ACC_TYPE4 q_priv[DK_VEC];
259
299
  const ulong q_row_offset = batch_idx * q_nb3 + head_idx * q_nb2;
260
300
  const global DATA_TYPE4* q_ptr = (const global DATA_TYPE4*)(q_base + q_row_offset);
261
- #pragma unroll
301
+ FA_UNROLL
262
302
  for (int i = 0; i < DK_VEC; ++i) {
263
303
  q_priv[i] = CONVERT_ACC4(q_ptr[i]);
264
304
  }
@@ -270,12 +310,12 @@ __kernel void flash_attn_f16_q1(
270
310
  sinks_ptr = (const global ACC_TYPE*)((const global char*)sinks_void + sinks_offset);
271
311
  }
272
312
 
273
- ACC_TYPE m_i = (sinks_ptr != NULL) ? sinks_ptr[head_idx] : -INFINITY;
313
+ ACC_TYPE m_i = (sinks_ptr != NULL) ? sinks_ptr[head_idx] : FA_M_INIT;
274
314
  for (int k_idx = tid; k_idx < n_kv; k_idx += Q1_WG_SIZE) {
275
315
  const ulong k_row_offset = batch_idx * k_nb3 + head_kv_idx * k_nb2 + k_idx * k_nb1;
276
316
  const global DATA_TYPE4* k_ptr = (const global DATA_TYPE4*)(k_base + k_row_offset);
277
317
  ACC_TYPE4 dot_acc = (ACC_TYPE4)(0.0f);
278
- #pragma unroll
318
+ FA_UNROLL
279
319
  for (int k = 0; k < DK_VEC; k++) {
280
320
  dot_acc = mad(q_priv[k], CONVERT_ACC4(k_ptr[k]), dot_acc);
281
321
  }
@@ -293,7 +333,7 @@ __kernel void flash_attn_f16_q1(
293
333
  __local ACC_TYPE local_m[Q1_WG_SIZE];
294
334
  local_m[tid] = m_i;
295
335
  barrier(CLK_LOCAL_MEM_FENCE);
296
- #pragma unroll
336
+ FA_UNROLL
297
337
  for (int s = Q1_WG_SIZE / 2; s > 0; s >>= 1) {
298
338
  if (tid < s) local_m[tid] = max(local_m[tid], local_m[tid + s]);
299
339
  barrier(CLK_LOCAL_MEM_FENCE);
@@ -301,7 +341,7 @@ __kernel void flash_attn_f16_q1(
301
341
  const ACC_TYPE m_final = local_m[0];
302
342
 
303
343
  ACC_TYPE4 o_acc[DV_VEC];
304
- #pragma unroll
344
+ FA_UNROLL
305
345
  for (int i = 0; i < DV_VEC; ++i) o_acc[i] = (ACC_TYPE4)(0.0f);
306
346
  ACC_TYPE l_i = 0.0f;
307
347
 
@@ -311,7 +351,7 @@ __kernel void flash_attn_f16_q1(
311
351
  const global DATA_TYPE4* k_ptr = (const global DATA_TYPE4*)(k_base + k_row_offset);
312
352
  const global DATA_TYPE4* v_ptr = (const global DATA_TYPE4*)(v_base + v_row_offset);
313
353
  ACC_TYPE4 dot_acc = (ACC_TYPE4)(0.0f);
314
- #pragma unroll
354
+ FA_UNROLL
315
355
  for (int k = 0; k < DK_VEC; k++) {
316
356
  dot_acc = mad(q_priv[k], CONVERT_ACC4(k_ptr[k]), dot_acc);
317
357
  }
@@ -325,7 +365,7 @@ __kernel void flash_attn_f16_q1(
325
365
  }
326
366
  const ACC_TYPE p = exp(score - m_final);
327
367
  l_i += p;
328
- #pragma unroll
368
+ FA_UNROLL
329
369
  for (int i = 0; i < DV_VEC; i++) {
330
370
  o_acc[i] = mad(p, CONVERT_ACC4(v_ptr[i]), o_acc[i]);
331
371
  }
@@ -335,7 +375,7 @@ __kernel void flash_attn_f16_q1(
335
375
  __local ACC_TYPE4 local_o_comp[Q1_WG_SIZE];
336
376
  local_l[tid] = l_i;
337
377
  barrier(CLK_LOCAL_MEM_FENCE);
338
- #pragma unroll
378
+ FA_UNROLL
339
379
  for (int s = Q1_WG_SIZE / 2; s > 0; s >>= 1) {
340
380
  if (tid < s) local_l[tid] += local_l[tid + s];
341
381
  barrier(CLK_LOCAL_MEM_FENCE);
@@ -354,7 +394,7 @@ __kernel void flash_attn_f16_q1(
354
394
  for (int i = 0; i < DV_VEC; i++) {
355
395
  local_o_comp[tid] = o_acc[i];
356
396
  barrier(CLK_LOCAL_MEM_FENCE);
357
- #pragma unroll
397
+ FA_UNROLL
358
398
  for (int s = Q1_WG_SIZE / 2; s > 0; s >>= 1) {
359
399
  if (tid < s) local_o_comp[tid] += local_o_comp[tid + s];
360
400
  barrier(CLK_LOCAL_MEM_FENCE);
@@ -364,7 +404,7 @@ __kernel void flash_attn_f16_q1(
364
404
  }
365
405
  }
366
406
  } else if (tid == 0) {
367
- #pragma unroll
407
+ FA_UNROLL
368
408
  for (int i = 0; i < DV_VEC; ++i) o_row[i] = (DATA_TYPE4)(0.0f);
369
409
  }
370
410
  }