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
@@ -5,35 +5,31 @@ void main() {
5
5
  return;
6
6
  }
7
7
 
8
- const uint row = i / p.ne20;
9
- const uint col = i - row * p.ne20;
8
+ const uint i23 = fastdiv(i, p.ne2_012mp, p.ne2_012L);
9
+ const uint i23_offset = i23 * p.ne22*p.ne21*p.ne20;
10
+ const uint i22 = fastdiv(i - i23_offset, p.ne2_01mp, p.ne2_01L);
11
+ const uint i22_offset = i22*p.ne21*p.ne20;
12
+ const uint i21 = fastdiv(i - i23_offset - i22_offset, p.ne2_0mp, p.ne2_0L);
13
+ const uint i20 = i - i23_offset - i22_offset - i21*p.ne20;
10
14
 
11
- const uint i3 = row / (p.ne01 * p.ne02);
12
- const uint i2 = (row % (p.ne01 * p.ne02)) / p.ne01;
13
- const uint i1 = row % p.ne01;
14
- const uint src_idx = i3 * p.nb03 + i2 * p.nb02 + i1 * p.nb01 + col;
15
-
16
- const uint dst_i3 = row / (p.ne11 * p.ne12);
17
- const uint dst_i2 = (row % (p.ne11 * p.ne12)) / p.ne11;
18
- const uint dst_i1 = row % p.ne11;
19
- const uint dst_idx = dst_i3 * p.nb13 + dst_i2 * p.nb12 + dst_i1 * p.nb11 + col;
15
+ const uint src_idx_a = get_aoffset() + i23 * p.nb03 + i22 * p.nb02 + i21 * p.nb01 + i20 * p.nb00;
16
+ const uint src_idx_b = get_boffset() + i23 * p.nb13 + i22 * p.nb12 + i21 * p.nb11 + i20 * p.nb10;
17
+ const uint dst_idx = get_doffset() + i23 * p.nb23 + i22 * p.nb22 + i21 * p.nb21 + i20 * p.nb20;
20
18
 
21
19
  if (p.mode == 0) {
22
20
  // Default
23
- const uint offset = p.ne00 / 2;
24
- const uint idx = src_idx;
21
+ const uint offset = (p.ne00 / 2) * p.nb00;
22
+ const uint idx = src_idx_a;
25
23
 
26
24
  data_d[dst_idx] = D_TYPE(op(float(data_a[idx]), float(data_a[idx + offset])));
27
25
  } else if (p.mode == 1) {
28
26
  // Swapped
29
- const uint offset = p.ne00 / 2;
30
- const uint idx = src_idx;
27
+ const uint offset = (p.ne00 / 2) * p.nb00;
28
+ const uint idx = src_idx_a;
31
29
 
32
30
  data_d[dst_idx] = D_TYPE(op(float(data_a[idx + offset]), float(data_a[idx])));
33
31
  } else {
34
32
  // Split
35
- const uint idx = src_idx;
36
-
37
- data_d[dst_idx] = D_TYPE(op(float(data_a[idx]), float(data_b[idx])));
33
+ data_d[dst_idx] = D_TYPE(op(float(data_a[src_idx_a]), float(data_b[src_idx_b])));
38
34
  }
39
35
  }
@@ -14,16 +14,13 @@ void main() {
14
14
  const uint row = gl_WorkGroupID.z * 262144 + gl_WorkGroupID.y * 512 + gl_WorkGroupID.x;
15
15
  const uint tid = gl_LocalInvocationID.x;
16
16
 
17
- const uint i3 = row / (p.ne11 * p.ne12);
18
- const uint i3_offset = i3 * p.ne12 * p.ne11;
19
- const uint i2 = (row - i3_offset) / p.ne11;
20
- const uint i2_offset = i2 * p.ne11;
21
- const uint i1 = row - i3_offset - i2_offset;
17
+ const uint a_base = get_aoffset() + src0_idx(row * p.ne00);
18
+ const uint d_base = get_doffset() + dst_idx(row * p.ne10);
22
19
 
23
20
  sum[tid] = FLOAT_TYPE(0.0f); // partial sum for thread in warp
24
21
 
25
22
  [[unroll]] for (uint i0 = tid; i0 < p.ne00; i0 += BLOCK_SIZE) {
26
- const FLOAT_TYPE xi = FLOAT_TYPE(data_a[i3*p.nb03 + i2*p.nb02 + i1*p.nb01 + i0]);
23
+ const FLOAT_TYPE xi = FLOAT_TYPE(data_a[a_base + i0*p.nb00]);
27
24
  sum[tid] += xi * xi;
28
25
  }
29
26
 
@@ -39,6 +36,6 @@ void main() {
39
36
  const FLOAT_TYPE scale = 1.0f / max(sqrt(sum[0]), FLOAT_TYPE(p.param1));
40
37
 
41
38
  [[unroll]] for (uint i0 = tid; i0 < p.ne00; i0 += BLOCK_SIZE) {
42
- data_d[i3*p.nb13 + i2*p.nb12 + i1*p.nb11 + i0] = D_TYPE(scale * FLOAT_TYPE(data_a[i3*p.nb03 + i2*p.nb02 + i1*p.nb01 + i0]));
39
+ data_d[d_base + i0*p.nb10] = D_TYPE(scale * FLOAT_TYPE(data_a[a_base + i0*p.nb00]));
43
40
  }
44
41
  }
@@ -28,13 +28,10 @@ vec2 cache_b_ds;
28
28
 
29
29
  #include "mul_mat_vecq_funcs.glsl"
30
30
 
31
- void iter(inout FLOAT_TYPE temp[NUM_COLS][NUM_ROWS], const uint first_row, const uint num_rows, const uint tid, const uint i) {
31
+ void iter(inout FLOAT_TYPE temp[NUM_COLS][NUM_ROWS], const uint first_row, const uint num_rows, const uint col, const uint b_qs_idx) {
32
32
  [[unroll]] for (uint j = 0; j < NUM_COLS; ++j) {
33
- const uint col = i*BLOCK_SIZE + tid*K_PER_ITER;
34
-
35
33
  // Preload data_b block
36
34
  const uint b_block_idx = (j*p.batch_stride_b + col) / QUANT_K_Q8_1 + b_offset;
37
- const uint b_qs_idx = tid % (32 / K_PER_ITER);
38
35
  const uint b_block_idx_outer = b_block_idx / 4;
39
36
  const uint b_block_idx_inner = b_block_idx % 4;
40
37
  cache_b_ds = vec2(data_b[b_block_idx_outer].ds[b_block_idx_inner]);
@@ -91,35 +88,35 @@ void compute_outputs(const uint32_t first_row, const uint32_t num_rows) {
91
88
  }
92
89
  }
93
90
 
94
- uint num_iters = p.ncols / (K_PER_ITER * BLOCK_SIZE);
95
- if (num_iters * K_PER_ITER * BLOCK_SIZE + K_PER_ITER*tid < p.ncols) {
91
+ const uint col_stride = K_PER_ITER * BLOCK_SIZE;
92
+ uint num_iters = p.ncols / col_stride;
93
+ if (num_iters * col_stride + K_PER_ITER * tid < p.ncols) {
96
94
  num_iters++;
97
95
  }
98
- int unroll_count = 4;
99
- uint unrolled_iters = num_iters & ~(unroll_count - 1);
100
96
 
101
- uint i = 0;
102
- while (i < unrolled_iters) {
97
+ const uint b_qs_idx = tid % (32 / K_PER_ITER);
98
+ uint col = tid * K_PER_ITER;
99
+ while (num_iters >= 4) {
103
100
  // Manually partially unroll the loop
104
- [[unroll]] for (uint k = 0; k < unroll_count; ++k) {
105
- iter(temp, first_row, num_rows, tid, i*K_PER_ITER);
106
- i++;
101
+ [[unroll]] for (uint k = 0; k < 4; ++k) {
102
+ iter(temp, first_row, num_rows, col, b_qs_idx);
103
+ col += col_stride;
107
104
  }
108
- }
109
105
 
110
- unroll_count = 2;
111
- unrolled_iters = num_iters & ~(unroll_count - 1);
106
+ num_iters -= 4;
107
+ }
112
108
 
113
- while (i < unrolled_iters) {
109
+ if (num_iters >= 2) {
114
110
  // Manually partially unroll the loop
115
- [[unroll]] for (uint k = 0; k < unroll_count; ++k) {
116
- iter(temp, first_row, num_rows, tid, i*K_PER_ITER);
117
- i++;
118
- }
111
+ iter(temp, first_row, num_rows, col, b_qs_idx);
112
+ col += col_stride;
113
+ iter(temp, first_row, num_rows, col, b_qs_idx);
114
+ col += col_stride;
115
+ num_iters -= 2;
119
116
  }
120
- while (i < num_iters) {
121
- iter(temp, first_row, num_rows, tid, i*K_PER_ITER);
122
- i++;
117
+
118
+ if (num_iters > 0) {
119
+ iter(temp, first_row, num_rows, col, b_qs_idx);
123
120
  }
124
121
 
125
122
  reduce_result(temp, d_offset, first_row, num_rows, tid);
@@ -38,17 +38,7 @@
38
38
  #define LOAD_VEC_B 1
39
39
  #endif
40
40
 
41
- // Load 2 values at once without affecting index calculations through LOAD_VEC
42
- #if (defined(DATA_A_F32) || defined(DATA_A_F16) || defined(DATA_A_BF16)) && !defined(ALIGNED)
43
- #define LOAD_VEC_BATCH_A 2
44
- #else
45
- #define LOAD_VEC_BATCH_A 1
46
- #endif
47
- #if !defined(ALIGNED)
48
- #define LOAD_VEC_BATCH_B 2
49
- #else
50
- #define LOAD_VEC_BATCH_B 1
51
- #endif
41
+ layout (constant_id = 11) const uint ALIGNED = 0;
52
42
 
53
43
  #if !defined(TO_FLOAT_TYPE)
54
44
  #define TO_FLOAT_TYPE FLOAT_TYPE
@@ -57,6 +47,13 @@
57
47
  layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in;
58
48
 
59
49
  layout (binding = 0) readonly buffer A {A_TYPE data_a[];};
50
+ #if defined(DATA_A_F32)
51
+ layout (binding = 0) readonly buffer A_SCALAR {float data_a_scalar[];};
52
+ #elif defined(DATA_A_F16)
53
+ layout (binding = 0) readonly buffer A_SCALAR {float16_t data_a_scalar[];};
54
+ #elif defined(DATA_A_BF16)
55
+ layout (binding = 0) readonly buffer A_SCALAR {uint16_t data_a_scalar[];};
56
+ #endif
60
57
  #if defined(A_TYPE_PACKED16)
61
58
  layout (binding = 0) readonly buffer A_PACKED16 {A_TYPE_PACKED16 data_a_packed16[];};
62
59
  #endif
@@ -65,6 +62,7 @@ layout (binding = 0) readonly buffer A_PACKED32 {A_TYPE_PACKED32 data_a_packed32
65
62
  #endif
66
63
 
67
64
  layout (binding = 1) readonly buffer B {B_TYPE data_b[];};
65
+ layout (binding = 1) readonly buffer B_SCALAR {B_TYPE_SCALAR data_b_scalar[];};
68
66
  layout (binding = 2) writeonly buffer D {D_TYPE data_d[];};
69
67
 
70
68
  #ifdef MUL_MAT_ID
@@ -194,13 +192,23 @@ void main() {
194
192
  const uint warp_r = warp_i % (BM / WM);
195
193
  const uint warp_c = warp_i / (BM / WM);
196
194
 
197
- const uint loadr_a = gl_LocalInvocationID.x % (BK / LOAD_VEC_A / LOAD_VEC_BATCH_A);
198
- const uint loadc_a = gl_LocalInvocationID.x / (BK / LOAD_VEC_A / LOAD_VEC_BATCH_A);
199
- const uint loadr_b = gl_LocalInvocationID.x % (BK / LOAD_VEC_B / LOAD_VEC_BATCH_B);
200
- const uint loadc_b = gl_LocalInvocationID.x / (BK / LOAD_VEC_B / LOAD_VEC_BATCH_B);
195
+ #if defined(DATA_A_F32) || defined(DATA_A_F16) || defined(DATA_A_BF16)
196
+ const uint LOAD_VEC_A_EFF = (ALIGNED != 0) ? LOAD_VEC_A : 1;
197
+ const uint LOAD_VEC_BATCH_A = (ALIGNED != 0) ? 1 : 2;
198
+ #else
199
+ const uint LOAD_VEC_A_EFF = LOAD_VEC_A;
200
+ const uint LOAD_VEC_BATCH_A = 1;
201
+ #endif
202
+ const uint LOAD_VEC_B_EFF = (ALIGNED != 0) ? LOAD_VEC_B : 1;
203
+ const uint LOAD_VEC_BATCH_B = (ALIGNED != 0) ? 1 : 2;
204
+
205
+ const uint loadr_a = gl_LocalInvocationID.x % (BK / LOAD_VEC_A_EFF / LOAD_VEC_BATCH_A);
206
+ const uint loadc_a = gl_LocalInvocationID.x / (BK / LOAD_VEC_A_EFF / LOAD_VEC_BATCH_A);
207
+ const uint loadr_b = gl_LocalInvocationID.x % (BK / LOAD_VEC_B_EFF / LOAD_VEC_BATCH_B);
208
+ const uint loadc_b = gl_LocalInvocationID.x / (BK / LOAD_VEC_B_EFF / LOAD_VEC_BATCH_B);
201
209
 
202
- const uint loadstride_a = gl_WorkGroupSize.x * LOAD_VEC_A * LOAD_VEC_BATCH_A / BK;
203
- const uint loadstride_b = gl_WorkGroupSize.x * LOAD_VEC_B * LOAD_VEC_BATCH_B / BK;
210
+ const uint loadstride_a = gl_WorkGroupSize.x * LOAD_VEC_A_EFF * LOAD_VEC_BATCH_A / BK;
211
+ const uint loadstride_b = gl_WorkGroupSize.x * LOAD_VEC_B_EFF * LOAD_VEC_BATCH_B / BK;
204
212
 
205
213
  #ifdef MUL_MAT_ID
206
214
  #ifdef MUL_MAT_ID_USE_SUBGROUPS
@@ -239,15 +247,15 @@ void main() {
239
247
 
240
248
  uint pos_a =
241
249
  #ifdef MUL_MAT_ID
242
- expert_idx * (p.batch_stride_a / LOAD_VEC_A) +
250
+ expert_idx * (p.batch_stride_a / LOAD_VEC_A_EFF) +
243
251
  #else
244
- batch_idx_a * (p.batch_stride_a / LOAD_VEC_A) +
252
+ batch_idx_a * (p.batch_stride_a / LOAD_VEC_A_EFF) +
245
253
  #endif
246
- (ir * BM * p.stride_a + start_k) / LOAD_VEC_A;
254
+ (ir * BM * p.stride_a + start_k) / LOAD_VEC_A_EFF;
247
255
  #ifdef MUL_MAT_ID
248
256
  uint pos_b = 0;
249
257
  #else
250
- uint pos_b = (batch_idx * p.batch_stride_b + ic * BN * p.stride_b + start_k) / LOAD_VEC_B;
258
+ uint pos_b = (batch_idx * p.batch_stride_b + ic * BN * p.stride_b + start_k) / LOAD_VEC_B_EFF;
251
259
  #endif
252
260
 
253
261
  #ifdef COOPMAT
@@ -287,8 +295,8 @@ void main() {
287
295
 
288
296
  barrier();
289
297
 
290
- pos_a += BK / LOAD_VEC_A;
291
- pos_b += BK / LOAD_VEC_B;
298
+ pos_a += BK / LOAD_VEC_A_EFF;
299
+ pos_b += BK / LOAD_VEC_B_EFF;
292
300
 
293
301
  #ifdef COOPMAT
294
302
  [[unroll]] for (uint i = 0; i < BK; i += TK) {
@@ -36,6 +36,7 @@ layout (constant_id = 3) const uint BK = 16; // Assumed to be 32 if working wit
36
36
  layout (constant_id = 4) const bool enable_smaller_matrices = false;
37
37
  const uint BNover2 = enable_smaller_matrices ? (BN / 2) : BN;
38
38
  const uint BNover4 = enable_smaller_matrices ? (BN / 4) : BN;
39
+ layout (constant_id = 5) const uint ALIGNED = 0;
39
40
 
40
41
  layout (push_constant) uniform parameter
41
42
  {
@@ -111,7 +112,7 @@ layout(buffer_reference, std430, buffer_reference_align = 2) buffer decodeBufB {
111
112
  };
112
113
 
113
114
  uint _ne1;
114
- layout (constant_id = 5) const uint subgroup_size = 32;
115
+ layout (constant_id = 6) const uint subgroup_size = 32;
115
116
  shared uvec4 ballots_sh[BLOCK_SIZE / subgroup_size];
116
117
 
117
118
  B_TYPE decodeFuncB(const in decodeBufB bl, const in uint blockCoords[2], const in uint coordInBlock[2])
@@ -297,12 +298,12 @@ void main() {
297
298
 
298
299
  // Hint to the compiler that values are aligned (want 16B alignment).
299
300
  // Quants are always block-aligned, no alignment needed.
300
- #if ALIGNED
301
+ if (ALIGNED != 0) {
301
302
  #if QUANT_K == 1
302
- stride_a &= ~7;
303
- #endif
304
- stride_b &= ~7;
303
+ stride_a &= ~7;
305
304
  #endif
305
+ stride_b &= ~7;
306
+ }
306
307
 
307
308
  // Create layouts for both clamped and unclamped accesses
308
309
  tensorLayoutNV<2> tensorLayoutA = createTensorLayoutNV(2);
@@ -1,50 +1,57 @@
1
1
  void load_a_to_shmem(const uint pos_a, const uint row, const uint col, const uint idx_m, const uint block, const uint end_k) {
2
2
  #if defined(DATA_A_F32) || defined(DATA_A_F16)
3
3
  #if LOAD_VEC_A == 8
4
- const uint idx = pos_a + col * p.stride_a / LOAD_VEC_A + row;
5
- const uint buf_idx = col * SHMEM_STRIDE + row * LOAD_VEC_A / 2;
6
- FLOAT_TYPEV8 aa = FLOAT_TYPEV8(data_a[idx]);
7
- buf_a[buf_idx ] = aa[0].xy;
8
- buf_a[buf_idx + 1] = aa[0].zw;
9
- buf_a[buf_idx + 2] = aa[1].xy;
10
- buf_a[buf_idx + 3] = aa[1].zw;
4
+ if (ALIGNED != 0) {
5
+ const uint idx = pos_a + col * p.stride_a / LOAD_VEC_A + row;
6
+ const uint buf_idx = col * SHMEM_STRIDE + row * LOAD_VEC_A / 2;
7
+ FLOAT_TYPEV8 aa = FLOAT_TYPEV8(data_a[idx]);
8
+ buf_a[buf_idx ] = aa[0].xy;
9
+ buf_a[buf_idx + 1] = aa[0].zw;
10
+ buf_a[buf_idx + 2] = aa[1].xy;
11
+ buf_a[buf_idx + 3] = aa[1].zw;
12
+ return;
13
+ }
11
14
  #elif LOAD_VEC_A == 4
12
- const uint idx = pos_a + col * p.stride_a / LOAD_VEC_A + row;
13
- const uint buf_idx = col * SHMEM_STRIDE + row * LOAD_VEC_A / 2;
14
- FLOAT_TYPEV4 aa = FLOAT_TYPEV4(data_a[idx]);
15
- buf_a[buf_idx ] = aa.xy;
16
- buf_a[buf_idx + 1] = aa.zw;
17
- #else // LOAD_VEC_BATCH_A == 2
15
+ if (ALIGNED != 0) {
16
+ const uint idx = pos_a + col * p.stride_a / LOAD_VEC_A + row;
17
+ const uint buf_idx = col * SHMEM_STRIDE + row * LOAD_VEC_A / 2;
18
+ FLOAT_TYPEV4 aa = FLOAT_TYPEV4(data_a[idx]);
19
+ buf_a[buf_idx ] = aa.xy;
20
+ buf_a[buf_idx + 1] = aa.zw;
21
+ return;
22
+ }
23
+ #endif
18
24
  const uint idx = pos_a + col * p.stride_a + row * 2;
19
25
  const uint buf_idx = col * SHMEM_STRIDE + row;
20
26
  if (idx_m < p.M && block + row * 2 + 1 < end_k) {
21
- buf_a[buf_idx] = FLOAT_TYPEV2(data_a[idx],
22
- data_a[idx + 1]);
27
+ buf_a[buf_idx] = FLOAT_TYPEV2(data_a_scalar[idx],
28
+ data_a_scalar[idx + 1]);
23
29
  } else if (idx_m < p.M && block + row * 2 < end_k) {
24
- buf_a[buf_idx] = FLOAT_TYPEV2(data_a[idx], 0.0f);
30
+ buf_a[buf_idx] = FLOAT_TYPEV2(data_a_scalar[idx], 0.0f);
25
31
  } else {
26
32
  buf_a[buf_idx] = FLOAT_TYPEV2(0.0f);
27
33
  }
28
- #endif
29
34
  #elif defined(DATA_A_BF16)
30
35
  #if LOAD_VEC_A == 4
31
- const uint idx = pos_a + col * p.stride_a / LOAD_VEC_A + row;
32
- const uint buf_idx = col * SHMEM_STRIDE + row * LOAD_VEC_A / 2;
33
- FLOAT_TYPEV4 aa = FLOAT_TYPEV4(TO_FLOAT_TYPE(data_a[idx]));
34
- buf_a[buf_idx ] = aa.xy;
35
- buf_a[buf_idx + 1] = aa.zw;
36
- #else // LOAD_VEC_BATCH_A == 2
36
+ if (ALIGNED != 0) {
37
+ const uint idx = pos_a + col * p.stride_a / LOAD_VEC_A + row;
38
+ const uint buf_idx = col * SHMEM_STRIDE + row * LOAD_VEC_A / 2;
39
+ FLOAT_TYPEV4 aa = FLOAT_TYPEV4(TO_FLOAT_TYPE(data_a[idx]));
40
+ buf_a[buf_idx ] = aa.xy;
41
+ buf_a[buf_idx + 1] = aa.zw;
42
+ return;
43
+ }
44
+ #endif
37
45
  const uint idx = pos_a + col * p.stride_a + row * 2;
38
46
  const uint buf_idx = col * SHMEM_STRIDE + row;
39
47
  if (idx_m < p.M && block + row * 2 + 1 < end_k) {
40
- buf_a[buf_idx] = FLOAT_TYPEV2(TO_FLOAT_TYPE(data_a[idx]),
41
- TO_FLOAT_TYPE(data_a[idx + 1]));
48
+ buf_a[buf_idx] = FLOAT_TYPEV2(TO_FLOAT_TYPE(data_a_scalar[idx]),
49
+ TO_FLOAT_TYPE(data_a_scalar[idx + 1]));
42
50
  } else if (idx_m < p.M && block + row * 2 < end_k) {
43
- buf_a[buf_idx] = FLOAT_TYPEV2(TO_FLOAT_TYPE(data_a[idx]), 0.0f);
51
+ buf_a[buf_idx] = FLOAT_TYPEV2(TO_FLOAT_TYPE(data_a_scalar[idx]), 0.0f);
44
52
  } else {
45
53
  buf_a[buf_idx] = FLOAT_TYPEV2(0.0f);
46
54
  }
47
- #endif
48
55
  #elif defined(DATA_A_Q4_0)
49
56
  const uint idx = pos_a + col * p.stride_a / LOAD_VEC_A + row;
50
57
  const uint buf_idx = col * SHMEM_STRIDE + row * LOAD_VEC_A / 4;
@@ -526,75 +533,85 @@ void load_a_to_shmem(const uint pos_a, const uint row, const uint col, const uin
526
533
  #if !defined(MUL_MAT_ID)
527
534
  void load_b_to_shmem(const uint pos_b, const uint row, const uint col, const uint idx_n, const uint block, const uint end_k) {
528
535
  #if LOAD_VEC_B == 8
529
- // Not supported for b_type bf16 because bf16mat2x4 does not exist
530
- const uint idx = pos_b + col * p.stride_b / LOAD_VEC_B + row;
531
- const uint buf_idx = col * SHMEM_STRIDE + row * LOAD_VEC_B / 2;
532
- FLOAT_TYPEV8 bb = FLOAT_TYPEV8(data_b[idx]);
533
- buf_b[buf_idx + 0] = bb[0].xy;
534
- buf_b[buf_idx + 1] = bb[0].zw;
535
- buf_b[buf_idx + 2] = bb[1].xy;
536
- buf_b[buf_idx + 3] = bb[1].zw;
536
+ if (ALIGNED != 0) {
537
+ // Not supported for b_type bf16 because bf16mat2x4 does not exist
538
+ const uint idx = pos_b + col * p.stride_b / LOAD_VEC_B + row;
539
+ const uint buf_idx = col * SHMEM_STRIDE + row * LOAD_VEC_B / 2;
540
+ FLOAT_TYPEV8 bb = FLOAT_TYPEV8(data_b[idx]);
541
+ buf_b[buf_idx + 0] = bb[0].xy;
542
+ buf_b[buf_idx + 1] = bb[0].zw;
543
+ buf_b[buf_idx + 2] = bb[1].xy;
544
+ buf_b[buf_idx + 3] = bb[1].zw;
545
+ return;
546
+ }
537
547
  #elif LOAD_VEC_B == 4
538
- const uint idx = pos_b + col * p.stride_b / LOAD_VEC_B + row;
539
- const uint buf_idx = col * SHMEM_STRIDE + row * LOAD_VEC_B / 2;
548
+ if (ALIGNED != 0) {
549
+ const uint idx = pos_b + col * p.stride_b / LOAD_VEC_B + row;
550
+ const uint buf_idx = col * SHMEM_STRIDE + row * LOAD_VEC_B / 2;
540
551
  #if defined(DATA_B_BF16)
541
- FLOAT_TYPEV4 bb = FLOAT_TYPEV4(TO_FLOAT_TYPE(data_b[idx]));
552
+ FLOAT_TYPEV4 bb = FLOAT_TYPEV4(TO_FLOAT_TYPE(data_b[idx]));
542
553
  #else
543
- FLOAT_TYPEV4 bb = FLOAT_TYPEV4(data_b[idx]);
554
+ FLOAT_TYPEV4 bb = FLOAT_TYPEV4(data_b[idx]);
555
+ #endif
556
+ buf_b[buf_idx + 0] = bb.xy;
557
+ buf_b[buf_idx + 1] = bb.zw;
558
+ return;
559
+ }
544
560
  #endif
545
- buf_b[buf_idx + 0] = bb.xy;
546
- buf_b[buf_idx + 1] = bb.zw;
547
- #else // LOAD_VEC_BATCH_B == 2
548
561
  const uint idx = pos_b + col * p.stride_b + row * 2;
549
562
  const uint buf_idx = col * SHMEM_STRIDE + row;
550
563
  if (idx_n < p.N && block + row * 2 + 1 < end_k) {
551
- buf_b[buf_idx] = FLOAT_TYPEV2(TO_FLOAT_TYPE(data_b[idx]),
552
- TO_FLOAT_TYPE(data_b[idx + 1]));
564
+ buf_b[buf_idx] = FLOAT_TYPEV2(TO_FLOAT_TYPE(data_b_scalar[idx]),
565
+ TO_FLOAT_TYPE(data_b_scalar[idx + 1]));
553
566
  } else if (idx_n < p.N && block + row * 2 < end_k) {
554
- buf_b[buf_idx] = FLOAT_TYPEV2(TO_FLOAT_TYPE(data_b[idx]), 0.0f);
567
+ buf_b[buf_idx] = FLOAT_TYPEV2(TO_FLOAT_TYPE(data_b_scalar[idx]), 0.0f);
555
568
  } else {
556
569
  buf_b[buf_idx] = FLOAT_TYPEV2(0.0f);
557
570
  }
558
- #endif
559
571
  }
560
572
  #else
561
573
  void load_b_to_shmem(const uint pos_b, const uint row, const uint col, const uint ic, const uint _ne1, const uint block, const uint end_k) {
562
574
  #if LOAD_VEC_B == 8
563
- // Not supported for b_type bf16 because bf16mat2x4 does not exist
564
- const u16vec2 row_idx = row_ids[col];
565
- const uint idx = pos_b + row_idx.y * p.batch_stride_b / LOAD_VEC_B + (row_idx.x % p.ne11) * p.stride_b / LOAD_VEC_B + row;
566
- const uint buf_idx = col * SHMEM_STRIDE + row * LOAD_VEC_B / 2;
567
- FLOAT_TYPEV8 bb = FLOAT_TYPEV8(data_b[idx]);
568
- buf_b[buf_idx + 0] = bb[0].xy;
569
- buf_b[buf_idx + 1] = bb[0].zw;
570
- buf_b[buf_idx + 2] = bb[1].xy;
571
- buf_b[buf_idx + 3] = bb[1].zw;
575
+ if (ALIGNED != 0) {
576
+ // Not supported for b_type bf16 because bf16mat2x4 does not exist
577
+ const u16vec2 row_idx = row_ids[col];
578
+ const uint idx = pos_b + row_idx.y * p.batch_stride_b / LOAD_VEC_B + (row_idx.x % p.ne11) * p.stride_b / LOAD_VEC_B + row;
579
+ const uint buf_idx = col * SHMEM_STRIDE + row * LOAD_VEC_B / 2;
580
+ FLOAT_TYPEV8 bb = FLOAT_TYPEV8(data_b[idx]);
581
+ buf_b[buf_idx + 0] = bb[0].xy;
582
+ buf_b[buf_idx + 1] = bb[0].zw;
583
+ buf_b[buf_idx + 2] = bb[1].xy;
584
+ buf_b[buf_idx + 3] = bb[1].zw;
585
+ return;
586
+ }
572
587
  #elif LOAD_VEC_B == 4
573
- const u16vec2 row_idx = row_ids[col];
574
- const uint idx = pos_b + row_idx.y * p.batch_stride_b / LOAD_VEC_B + (row_idx.x % p.ne11) * p.stride_b / LOAD_VEC_B + row;
575
- const uint buf_idx = col * SHMEM_STRIDE + row * LOAD_VEC_B / 2;
588
+ if (ALIGNED != 0) {
589
+ const u16vec2 row_idx = row_ids[col];
590
+ const uint idx = pos_b + row_idx.y * p.batch_stride_b / LOAD_VEC_B + (row_idx.x % p.ne11) * p.stride_b / LOAD_VEC_B + row;
591
+ const uint buf_idx = col * SHMEM_STRIDE + row * LOAD_VEC_B / 2;
576
592
  #if defined(DATA_B_BF16)
577
- FLOAT_TYPEV4 bb = FLOAT_TYPEV4(TO_FLOAT_TYPE(data_b[idx]));
593
+ FLOAT_TYPEV4 bb = FLOAT_TYPEV4(TO_FLOAT_TYPE(data_b[idx]));
578
594
  #else
579
- FLOAT_TYPEV4 bb = FLOAT_TYPEV4(data_b[idx]);
595
+ FLOAT_TYPEV4 bb = FLOAT_TYPEV4(data_b[idx]);
596
+ #endif
597
+ buf_b[buf_idx + 0] = bb.xy;
598
+ buf_b[buf_idx + 1] = bb.zw;
599
+ return;
600
+ }
580
601
  #endif
581
- buf_b[buf_idx + 0] = bb.xy;
582
- buf_b[buf_idx + 1] = bb.zw;
583
- #else // LOAD_VEC_BATCH_B == 2
584
602
  const uint row_i = ic * BN + col;
585
603
  const uint buf_idx = col * SHMEM_STRIDE + row;
586
604
  if (row_i < _ne1 && block + row * 2 + 1 < end_k) {
587
605
  const u16vec2 row_idx = row_ids[col];
588
606
  const uint idx = pos_b + row_idx.y * p.batch_stride_b + (row_idx.x % p.ne11) * p.stride_b + row * 2;
589
- buf_b[buf_idx] = FLOAT_TYPEV2(TO_FLOAT_TYPE(data_b[idx]),
590
- TO_FLOAT_TYPE(data_b[idx + 1]));
607
+ buf_b[buf_idx] = FLOAT_TYPEV2(TO_FLOAT_TYPE(data_b_scalar[idx]),
608
+ TO_FLOAT_TYPE(data_b_scalar[idx + 1]));
591
609
  } else if (row_i < _ne1 && block + row * 2 < end_k) {
592
610
  const u16vec2 row_idx = row_ids[col];
593
611
  const uint idx = pos_b + row_idx.y * p.batch_stride_b + (row_idx.x % p.ne11) * p.stride_b + row * 2;
594
- buf_b[buf_idx] = FLOAT_TYPEV2(TO_FLOAT_TYPE(data_b[idx]), 0.0f);
612
+ buf_b[buf_idx] = FLOAT_TYPEV2(TO_FLOAT_TYPE(data_b_scalar[idx]), 0.0f);
595
613
  } else {
596
614
  buf_b[buf_idx] = FLOAT_TYPEV2(0.0f);
597
615
  }
598
- #endif
599
616
  }
600
617
  #endif
@@ -1,26 +1,26 @@
1
1
  #version 450
2
2
 
3
- #include "generic_head.glsl"
4
3
  #include "types.glsl"
4
+ #include "generic_unary_head.glsl"
5
5
 
6
6
  #extension GL_EXT_control_flow_attributes : enable
7
7
  #define BLOCK_SIZE 512
8
8
 
9
9
  layout(local_size_x = BLOCK_SIZE, local_size_y = 1, local_size_z = 1) in;
10
10
 
11
- layout (binding = 0) readonly buffer X {A_TYPE data_a[];};
12
- layout (binding = 1) writeonly buffer D {D_TYPE data_d[];};
13
-
14
11
  shared vec2 sum[BLOCK_SIZE];
15
12
 
16
13
  void main() {
17
14
  const uint row = gl_WorkGroupID.z * 262144 + gl_WorkGroupID.y * 512 + gl_WorkGroupID.x;
18
15
  const uint tid = gl_LocalInvocationID.x;
19
16
 
17
+ const uint a_base = get_aoffset() + src0_idx(row * p.ne00);
18
+ const uint d_base = get_doffset() + dst_idx(row * p.ne10);
19
+
20
20
  sum[tid] = vec2(0.0f, 0.0f);
21
21
 
22
- [[unroll]] for (uint col = tid; col < p.KX; col += BLOCK_SIZE) {
23
- const float xi = float(data_a[row*p.KX + col]);
22
+ [[unroll]] for (uint i0 = tid; i0 < p.ne00; i0 += BLOCK_SIZE) {
23
+ const float xi = float(data_a[a_base + i0*p.nb00]);
24
24
  sum[tid].x += xi;
25
25
  sum[tid].y += xi * xi;
26
26
  }
@@ -34,11 +34,11 @@ void main() {
34
34
  barrier();
35
35
  }
36
36
 
37
- const float mean = sum[0].x / p.KX;
38
- const float var = sum[0].y / p.KX - mean * mean;
37
+ const float mean = sum[0].x / p.ne00;
38
+ const float var = sum[0].y / p.ne00 - mean * mean;
39
39
  const float inv_std = inversesqrt(var + p.param1);
40
40
 
41
- [[unroll]] for (uint col = tid; col < p.KX; col += BLOCK_SIZE) {
42
- data_d[row*p.KX + col] = D_TYPE((float(data_a[row*p.KX + col]) - mean) * inv_std);
41
+ [[unroll]] for (uint i0 = tid; i0 < p.ne00; i0 += BLOCK_SIZE) {
42
+ data_d[d_base + i0*p.nb10] = D_TYPE((float(data_a[a_base + i0*p.nb00]) - mean) * inv_std);
43
43
  }
44
44
  }
@@ -13,11 +13,11 @@ void main() {
13
13
  }
14
14
 
15
15
  // Destination multi-index (inlined dst_idx)
16
- const uint i13 = fastdiv(idx, p.ne1_012mp, p.ne1_012L);
16
+ const uint i13 = fastdiv(idx, p.ne1_012mp, fastdiv_L(p.ne1_Ls, 0));
17
17
  const uint i13_offset = i13 * p.ne12*p.ne11*p.ne10;
18
- const uint i12 = fastdiv(idx - i13_offset, p.ne1_01mp, p.ne1_01L);
18
+ const uint i12 = fastdiv(idx - i13_offset, p.ne1_01mp, fastdiv_L(p.ne1_Ls, 1));
19
19
  const uint i12_offset = i12*p.ne11*p.ne10;
20
- const uint i11 = fastdiv(idx - i13_offset - i12_offset, p.ne1_0mp, p.ne1_0L);
20
+ const uint i11 = fastdiv(idx - i13_offset - i12_offset, p.ne1_0mp, fastdiv_L(p.ne1_Ls, 2));
21
21
  const uint i10 = idx - i13_offset - i12_offset - i11*p.ne10;
22
22
  const uint d_idx = i13*p.nb13 + i12*p.nb12 + i11*p.nb11 + i10*p.nb10;
23
23
 
@@ -20,11 +20,11 @@ void main() {
20
20
  return;
21
21
  }
22
22
 
23
- const uint i3 = fastdiv(idx, p.ne1_012mp, p.ne1_012L);
23
+ const uint i3 = fastdiv(idx, p.ne1_012mp, fastdiv_L(p.ne1_Ls, 0));
24
24
  const uint i3_offset = i3 * p.ne12*p.ne11*p.ne10;
25
- const uint i2 = fastdiv(idx - i3_offset, p.ne1_01mp, p.ne1_01L);
25
+ const uint i2 = fastdiv(idx - i3_offset, p.ne1_01mp, fastdiv_L(p.ne1_Ls, 1));
26
26
  const uint i2_offset = i2*p.ne11*p.ne10;
27
- const uint i1 = fastdiv(idx - i3_offset - i2_offset, p.ne1_0mp, p.ne1_0L);
27
+ const uint i1 = fastdiv(idx - i3_offset - i2_offset, p.ne1_0mp, fastdiv_L(p.ne1_Ls, 2));
28
28
  const uint i0 = idx - i3_offset - i2_offset - i1*p.ne10;
29
29
 
30
30
  const uint p1 = floatBitsToUint(p.param1);
@@ -17,11 +17,11 @@ void main() {
17
17
  return;
18
18
  }
19
19
 
20
- const uint i03 = fastdiv(idx, p.ne0_012mp, p.ne0_012L);
20
+ const uint i03 = fastdiv(idx, p.ne0_012mp, fastdiv_L(p.ne0_Ls, 0));
21
21
  const uint i03_offset = i03 * p.ne02*p.ne01*p.ne00;
22
- const uint i02 = fastdiv(idx - i03_offset, p.ne0_01mp, p.ne0_01L);
22
+ const uint i02 = fastdiv(idx - i03_offset, p.ne0_01mp, fastdiv_L(p.ne0_Ls, 1));
23
23
  const uint i02_offset = i02*p.ne01*p.ne00;
24
- const uint i01 = fastdiv(idx - i03_offset - i02_offset, p.ne0_0mp, p.ne0_0L);
24
+ const uint i01 = fastdiv(idx - i03_offset - i02_offset, p.ne0_0mp, fastdiv_L(p.ne0_Ls, 2));
25
25
  const uint i00 = idx - i03_offset - i02_offset - i01*p.ne00;
26
26
 
27
27
  int param = floatBitsToInt(p.param1);