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
@@ -98,6 +98,46 @@
98
98
  c_reg.lo += convert_float8(acc.lo); \
99
99
  c_reg.hi += convert_float8(acc.hi); \
100
100
 
101
+ // Quarter-tile variant: computes 8 output columns (one skip-group) into a float8
102
+ // accumulator. Same reduction order / flush cadence as dotx16_reduce8, so the
103
+ // non-skipped path is byte-identical; it just lets the caller skip empty
104
+ // 8-column groups at finer granularity. Uses a private half8 `acc8`.
105
+ #define dotx8_reduce4(a_reg, b_lm, c_reg, lm_offset) \
106
+ acc8.s0 = dot(a_reg.s0123, b_lm[lm_offset + 0]); \
107
+ acc8.s1 = dot(a_reg.s0123, b_lm[lm_offset + 1]); \
108
+ acc8.s2 = dot(a_reg.s0123, b_lm[lm_offset + 2]); \
109
+ acc8.s3 = dot(a_reg.s0123, b_lm[lm_offset + 3]); \
110
+ acc8.s4 = dot(a_reg.s0123, b_lm[lm_offset + 4]); \
111
+ acc8.s5 = dot(a_reg.s0123, b_lm[lm_offset + 5]); \
112
+ acc8.s6 = dot(a_reg.s0123, b_lm[lm_offset + 6]); \
113
+ acc8.s7 = dot(a_reg.s0123, b_lm[lm_offset + 7]); \
114
+ acc8.s0 += dot(a_reg.s4567, b_lm[lm_offset + 32]); \
115
+ acc8.s1 += dot(a_reg.s4567, b_lm[lm_offset + 33]); \
116
+ acc8.s2 += dot(a_reg.s4567, b_lm[lm_offset + 34]); \
117
+ acc8.s3 += dot(a_reg.s4567, b_lm[lm_offset + 35]); \
118
+ acc8.s4 += dot(a_reg.s4567, b_lm[lm_offset + 36]); \
119
+ acc8.s5 += dot(a_reg.s4567, b_lm[lm_offset + 37]); \
120
+ acc8.s6 += dot(a_reg.s4567, b_lm[lm_offset + 38]); \
121
+ acc8.s7 += dot(a_reg.s4567, b_lm[lm_offset + 39]); \
122
+ c_reg += convert_float8(acc8); \
123
+ acc8.s0 = dot(a_reg.s89ab, b_lm[lm_offset + 64]); \
124
+ acc8.s1 = dot(a_reg.s89ab, b_lm[lm_offset + 65]); \
125
+ acc8.s2 = dot(a_reg.s89ab, b_lm[lm_offset + 66]); \
126
+ acc8.s3 = dot(a_reg.s89ab, b_lm[lm_offset + 67]); \
127
+ acc8.s4 = dot(a_reg.s89ab, b_lm[lm_offset + 68]); \
128
+ acc8.s5 = dot(a_reg.s89ab, b_lm[lm_offset + 69]); \
129
+ acc8.s6 = dot(a_reg.s89ab, b_lm[lm_offset + 70]); \
130
+ acc8.s7 = dot(a_reg.s89ab, b_lm[lm_offset + 71]); \
131
+ acc8.s0 += dot(a_reg.scdef, b_lm[lm_offset + 96]); \
132
+ acc8.s1 += dot(a_reg.scdef, b_lm[lm_offset + 97]); \
133
+ acc8.s2 += dot(a_reg.scdef, b_lm[lm_offset + 98]); \
134
+ acc8.s3 += dot(a_reg.scdef, b_lm[lm_offset + 99]); \
135
+ acc8.s4 += dot(a_reg.scdef, b_lm[lm_offset + 100]); \
136
+ acc8.s5 += dot(a_reg.scdef, b_lm[lm_offset + 101]); \
137
+ acc8.s6 += dot(a_reg.scdef, b_lm[lm_offset + 102]); \
138
+ acc8.s7 += dot(a_reg.scdef, b_lm[lm_offset + 103]); \
139
+ c_reg += convert_float8(acc8); \
140
+
101
141
 
102
142
  __attribute__((qcom_wave_pair_mode(1))) // 1=force single 2=force pair
103
143
  kernel void kernel_gemm_moe_q5_0_f32_ns(
@@ -110,7 +150,9 @@ kernel void kernel_gemm_moe_q5_0_f32_ns(
110
150
  __write_only image1d_buffer_t dst,
111
151
  __global int * total_tiles,
112
152
  uint ne00,
113
- uint ne01
153
+ uint ne01,
154
+ uint is_ragged,
155
+ uint skip_gran
114
156
  ) {
115
157
  uint block_id_m = get_global_id(1); // m_tile
116
158
  uint block_id_n = get_global_id(2); // n_tile
@@ -120,6 +162,28 @@ kernel void kernel_gemm_moe_q5_0_f32_ns(
120
162
  return;
121
163
  }
122
164
 
165
+ // Ragged tile-skip: when is_ragged and the upper 16 token-slots of this tile are all
166
+ // padding (router 0xFFFFFFFF), skip the second (reg_c.hi) dotx16_reduce8 half -> ~half
167
+ // the GEMM dot for sparse tiles. Numerically identical (the skipped lanes are padding).
168
+ // Ragged tile-skip: tokens are packed contiguously per expert (moe_scatter fills
169
+ // lanes 0..V-1, moe_fill pre-pads the rest), so router padding (0xFFFFFFFF) is always
170
+ // trailing. Find the valid-token count V and round it UP to the skip granularity
171
+ // skip_gran (columns per skip-group: 8 = quarter, 16 = half/legacy, 32 = disabled).
172
+ // A 8-column group g is all-padding iff its first column (8*g) >= n_active, so its
173
+ // dotx8_reduce4 is skipped. Numerically identical (skipped lanes are padding).
174
+ uint n_active = TILESIZE_N;
175
+ if (is_ragged && skip_gran < TILESIZE_N) {
176
+ uint n_valid = TILESIZE_N;
177
+ for (uint _t = 0; _t < TILESIZE_N; ++_t) {
178
+ if (src2[block_id_n * TILESIZE_N + _t] == 0xFFFFFFFFu) { n_valid = _t; break; }
179
+ }
180
+ n_active = min((uint)TILESIZE_N, ((n_valid + skip_gran - 1) / skip_gran) * skip_gran);
181
+ }
182
+ // Group 0 (cols 0-7) always runs; groups 1-3 skip when fully padding.
183
+ bool skip_g1 = (8u >= n_active);
184
+ bool skip_g2 = (16u >= n_active);
185
+ bool skip_g3 = (24u >= n_active);
186
+
123
187
  __private half16 reg_a;
124
188
  __private float32 reg_c = (float32)(0);
125
189
  __local half4 shared_b[128];
@@ -171,9 +235,11 @@ kernel void kernel_gemm_moe_q5_0_f32_ns(
171
235
  sub_group_barrier(CLK_LOCAL_MEM_FENCE);
172
236
 
173
237
  // 32 16x16 fp16 dot product with 8 elements reduction for better precision
174
- half16 acc;
175
- dotx16_reduce8(reg_a, shared_b, reg_c.lo, 0);
176
- dotx16_reduce8(reg_a, shared_b, reg_c.hi, 16);
238
+ half8 acc8;
239
+ dotx8_reduce4(reg_a, shared_b, reg_c.lo.lo, 0);
240
+ if (!skip_g1) { dotx8_reduce4(reg_a, shared_b, reg_c.lo.hi, 8); }
241
+ if (!skip_g2) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.lo, 16); }
242
+ if (!skip_g3) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.hi, 24); }
177
243
 
178
244
  // Repeat for second sub-block
179
245
  uint half_step = step + TILESIZE_K;
@@ -198,8 +264,10 @@ kernel void kernel_gemm_moe_q5_0_f32_ns(
198
264
  sub_group_barrier(CLK_LOCAL_MEM_FENCE);
199
265
 
200
266
  // 32 16x16 fp16 dot product with 3-levels reduction for better precision
201
- dotx16_reduce8(reg_a, shared_b, reg_c.lo, 0);
202
- dotx16_reduce8(reg_a, shared_b, reg_c.hi, 16);
267
+ dotx8_reduce4(reg_a, shared_b, reg_c.lo.lo, 0);
268
+ if (!skip_g1) { dotx8_reduce4(reg_a, shared_b, reg_c.lo.hi, 8); }
269
+ if (!skip_g2) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.lo, 16); }
270
+ if (!skip_g3) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.hi, 24); }
203
271
  }
204
272
 
205
273
  if ((get_global_id(0) + block_id_m * TILESIZE_M) >= ne01) {
@@ -98,6 +98,46 @@
98
98
  c_reg.lo += convert_float8(acc.lo); \
99
99
  c_reg.hi += convert_float8(acc.hi); \
100
100
 
101
+ // Quarter-tile variant: computes 8 output columns (one skip-group) into a float8
102
+ // accumulator. Same reduction order / flush cadence as dotx16_reduce8, so the
103
+ // non-skipped path is byte-identical; it just lets the caller skip empty
104
+ // 8-column groups at finer granularity. Uses a private half8 `acc8`.
105
+ #define dotx8_reduce4(a_reg, b_lm, c_reg, lm_offset) \
106
+ acc8.s0 = dot(a_reg.s0123, b_lm[lm_offset + 0]); \
107
+ acc8.s1 = dot(a_reg.s0123, b_lm[lm_offset + 1]); \
108
+ acc8.s2 = dot(a_reg.s0123, b_lm[lm_offset + 2]); \
109
+ acc8.s3 = dot(a_reg.s0123, b_lm[lm_offset + 3]); \
110
+ acc8.s4 = dot(a_reg.s0123, b_lm[lm_offset + 4]); \
111
+ acc8.s5 = dot(a_reg.s0123, b_lm[lm_offset + 5]); \
112
+ acc8.s6 = dot(a_reg.s0123, b_lm[lm_offset + 6]); \
113
+ acc8.s7 = dot(a_reg.s0123, b_lm[lm_offset + 7]); \
114
+ acc8.s0 += dot(a_reg.s4567, b_lm[lm_offset + 32]); \
115
+ acc8.s1 += dot(a_reg.s4567, b_lm[lm_offset + 33]); \
116
+ acc8.s2 += dot(a_reg.s4567, b_lm[lm_offset + 34]); \
117
+ acc8.s3 += dot(a_reg.s4567, b_lm[lm_offset + 35]); \
118
+ acc8.s4 += dot(a_reg.s4567, b_lm[lm_offset + 36]); \
119
+ acc8.s5 += dot(a_reg.s4567, b_lm[lm_offset + 37]); \
120
+ acc8.s6 += dot(a_reg.s4567, b_lm[lm_offset + 38]); \
121
+ acc8.s7 += dot(a_reg.s4567, b_lm[lm_offset + 39]); \
122
+ c_reg += convert_float8(acc8); \
123
+ acc8.s0 = dot(a_reg.s89ab, b_lm[lm_offset + 64]); \
124
+ acc8.s1 = dot(a_reg.s89ab, b_lm[lm_offset + 65]); \
125
+ acc8.s2 = dot(a_reg.s89ab, b_lm[lm_offset + 66]); \
126
+ acc8.s3 = dot(a_reg.s89ab, b_lm[lm_offset + 67]); \
127
+ acc8.s4 = dot(a_reg.s89ab, b_lm[lm_offset + 68]); \
128
+ acc8.s5 = dot(a_reg.s89ab, b_lm[lm_offset + 69]); \
129
+ acc8.s6 = dot(a_reg.s89ab, b_lm[lm_offset + 70]); \
130
+ acc8.s7 = dot(a_reg.s89ab, b_lm[lm_offset + 71]); \
131
+ acc8.s0 += dot(a_reg.scdef, b_lm[lm_offset + 96]); \
132
+ acc8.s1 += dot(a_reg.scdef, b_lm[lm_offset + 97]); \
133
+ acc8.s2 += dot(a_reg.scdef, b_lm[lm_offset + 98]); \
134
+ acc8.s3 += dot(a_reg.scdef, b_lm[lm_offset + 99]); \
135
+ acc8.s4 += dot(a_reg.scdef, b_lm[lm_offset + 100]); \
136
+ acc8.s5 += dot(a_reg.scdef, b_lm[lm_offset + 101]); \
137
+ acc8.s6 += dot(a_reg.scdef, b_lm[lm_offset + 102]); \
138
+ acc8.s7 += dot(a_reg.scdef, b_lm[lm_offset + 103]); \
139
+ c_reg += convert_float8(acc8); \
140
+
101
141
 
102
142
  __attribute__((qcom_wave_pair_mode(1))) // 1=force single 2=force pair
103
143
  kernel void kernel_gemm_moe_q5_1_f32_ns(
@@ -111,7 +151,9 @@ kernel void kernel_gemm_moe_q5_1_f32_ns(
111
151
  __write_only image1d_buffer_t dst,
112
152
  __global int * total_tiles,
113
153
  uint ne00,
114
- uint ne01
154
+ uint ne01,
155
+ uint is_ragged,
156
+ uint skip_gran
115
157
  ) {
116
158
  uint block_id_m = get_global_id(1); // m_tile
117
159
  uint block_id_n = get_global_id(2); // n_tile
@@ -121,6 +163,28 @@ kernel void kernel_gemm_moe_q5_1_f32_ns(
121
163
  return;
122
164
  }
123
165
 
166
+ // Ragged tile-skip: when is_ragged and the upper 16 token-slots of this tile are all
167
+ // padding (router 0xFFFFFFFF), skip the second (reg_c.hi) dotx16_reduce8 half -> ~half
168
+ // the GEMM dot for sparse tiles. Numerically identical (the skipped lanes are padding).
169
+ // Ragged tile-skip: tokens are packed contiguously per expert (moe_scatter fills
170
+ // lanes 0..V-1, moe_fill pre-pads the rest), so router padding (0xFFFFFFFF) is always
171
+ // trailing. Find the valid-token count V and round it UP to the skip granularity
172
+ // skip_gran (columns per skip-group: 8 = quarter, 16 = half/legacy, 32 = disabled).
173
+ // A 8-column group g is all-padding iff its first column (8*g) >= n_active, so its
174
+ // dotx8_reduce4 is skipped. Numerically identical (skipped lanes are padding).
175
+ uint n_active = TILESIZE_N;
176
+ if (is_ragged && skip_gran < TILESIZE_N) {
177
+ uint n_valid = TILESIZE_N;
178
+ for (uint _t = 0; _t < TILESIZE_N; ++_t) {
179
+ if (src2[block_id_n * TILESIZE_N + _t] == 0xFFFFFFFFu) { n_valid = _t; break; }
180
+ }
181
+ n_active = min((uint)TILESIZE_N, ((n_valid + skip_gran - 1) / skip_gran) * skip_gran);
182
+ }
183
+ // Group 0 (cols 0-7) always runs; groups 1-3 skip when fully padding.
184
+ bool skip_g1 = (8u >= n_active);
185
+ bool skip_g2 = (16u >= n_active);
186
+ bool skip_g3 = (24u >= n_active);
187
+
124
188
  __private half16 reg_a;
125
189
  __private float32 reg_c = (float32)(0);
126
190
  __local half4 shared_b[128];
@@ -173,9 +237,11 @@ kernel void kernel_gemm_moe_q5_1_f32_ns(
173
237
  sub_group_barrier(CLK_LOCAL_MEM_FENCE);
174
238
 
175
239
  // 32 16x16 fp16 dot product with 8 elements reduction for better precision
176
- half16 acc;
177
- dotx16_reduce8(reg_a, shared_b, reg_c.lo, 0);
178
- dotx16_reduce8(reg_a, shared_b, reg_c.hi, 16);
240
+ half8 acc8;
241
+ dotx8_reduce4(reg_a, shared_b, reg_c.lo.lo, 0);
242
+ if (!skip_g1) { dotx8_reduce4(reg_a, shared_b, reg_c.lo.hi, 8); }
243
+ if (!skip_g2) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.lo, 16); }
244
+ if (!skip_g3) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.hi, 24); }
179
245
 
180
246
  // Repeat for second sub-block
181
247
  uint half_step = step + TILESIZE_K;
@@ -200,8 +266,10 @@ kernel void kernel_gemm_moe_q5_1_f32_ns(
200
266
  sub_group_barrier(CLK_LOCAL_MEM_FENCE);
201
267
 
202
268
  // 32 16x16 fp16 dot product with 3-levels reduction for better precision
203
- dotx16_reduce8(reg_a, shared_b, reg_c.lo, 0);
204
- dotx16_reduce8(reg_a, shared_b, reg_c.hi, 16);
269
+ dotx8_reduce4(reg_a, shared_b, reg_c.lo.lo, 0);
270
+ if (!skip_g1) { dotx8_reduce4(reg_a, shared_b, reg_c.lo.hi, 8); }
271
+ if (!skip_g2) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.lo, 16); }
272
+ if (!skip_g3) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.hi, 24); }
205
273
  }
206
274
 
207
275
  if ((get_global_id(0) + block_id_m * TILESIZE_M) >= ne01) {
@@ -114,6 +114,46 @@ inline void get_scale_min_k4(
114
114
  c_reg.lo += convert_float8(acc.lo); \
115
115
  c_reg.hi += convert_float8(acc.hi); \
116
116
 
117
+ // Quarter-tile variant: computes 8 output columns (one skip-group) into a float8
118
+ // accumulator. Same reduction order / flush cadence as dotx16_reduce8, so the
119
+ // non-skipped path is byte-identical; it just lets the caller skip empty
120
+ // 8-column groups at finer granularity. Uses a private half8 `acc8`.
121
+ #define dotx8_reduce4(a_reg, b_lm, c_reg, lm_offset) \
122
+ acc8.s0 = dot(a_reg.s0123, b_lm[lm_offset + 0]); \
123
+ acc8.s1 = dot(a_reg.s0123, b_lm[lm_offset + 1]); \
124
+ acc8.s2 = dot(a_reg.s0123, b_lm[lm_offset + 2]); \
125
+ acc8.s3 = dot(a_reg.s0123, b_lm[lm_offset + 3]); \
126
+ acc8.s4 = dot(a_reg.s0123, b_lm[lm_offset + 4]); \
127
+ acc8.s5 = dot(a_reg.s0123, b_lm[lm_offset + 5]); \
128
+ acc8.s6 = dot(a_reg.s0123, b_lm[lm_offset + 6]); \
129
+ acc8.s7 = dot(a_reg.s0123, b_lm[lm_offset + 7]); \
130
+ acc8.s0 += dot(a_reg.s4567, b_lm[lm_offset + 32]); \
131
+ acc8.s1 += dot(a_reg.s4567, b_lm[lm_offset + 33]); \
132
+ acc8.s2 += dot(a_reg.s4567, b_lm[lm_offset + 34]); \
133
+ acc8.s3 += dot(a_reg.s4567, b_lm[lm_offset + 35]); \
134
+ acc8.s4 += dot(a_reg.s4567, b_lm[lm_offset + 36]); \
135
+ acc8.s5 += dot(a_reg.s4567, b_lm[lm_offset + 37]); \
136
+ acc8.s6 += dot(a_reg.s4567, b_lm[lm_offset + 38]); \
137
+ acc8.s7 += dot(a_reg.s4567, b_lm[lm_offset + 39]); \
138
+ c_reg += convert_float8(acc8); \
139
+ acc8.s0 = dot(a_reg.s89ab, b_lm[lm_offset + 64]); \
140
+ acc8.s1 = dot(a_reg.s89ab, b_lm[lm_offset + 65]); \
141
+ acc8.s2 = dot(a_reg.s89ab, b_lm[lm_offset + 66]); \
142
+ acc8.s3 = dot(a_reg.s89ab, b_lm[lm_offset + 67]); \
143
+ acc8.s4 = dot(a_reg.s89ab, b_lm[lm_offset + 68]); \
144
+ acc8.s5 = dot(a_reg.s89ab, b_lm[lm_offset + 69]); \
145
+ acc8.s6 = dot(a_reg.s89ab, b_lm[lm_offset + 70]); \
146
+ acc8.s7 = dot(a_reg.s89ab, b_lm[lm_offset + 71]); \
147
+ acc8.s0 += dot(a_reg.scdef, b_lm[lm_offset + 96]); \
148
+ acc8.s1 += dot(a_reg.scdef, b_lm[lm_offset + 97]); \
149
+ acc8.s2 += dot(a_reg.scdef, b_lm[lm_offset + 98]); \
150
+ acc8.s3 += dot(a_reg.scdef, b_lm[lm_offset + 99]); \
151
+ acc8.s4 += dot(a_reg.scdef, b_lm[lm_offset + 100]); \
152
+ acc8.s5 += dot(a_reg.scdef, b_lm[lm_offset + 101]); \
153
+ acc8.s6 += dot(a_reg.scdef, b_lm[lm_offset + 102]); \
154
+ acc8.s7 += dot(a_reg.scdef, b_lm[lm_offset + 103]); \
155
+ c_reg += convert_float8(acc8); \
156
+
117
157
 
118
158
  __attribute__((qcom_wave_pair_mode(1)))
119
159
  kernel void kernel_gemm_moe_q5_k_f32_ns(
@@ -128,7 +168,9 @@ kernel void kernel_gemm_moe_q5_k_f32_ns(
128
168
  __write_only image1d_buffer_t dst,
129
169
  __global int * total_tiles,
130
170
  uint ne00,
131
- uint ne01
171
+ uint ne01,
172
+ uint is_ragged,
173
+ uint skip_gran
132
174
  ) {
133
175
  uint block_id_m = get_global_id(1); // m_tile
134
176
  uint block_id_n = get_global_id(2); // n_tile
@@ -138,6 +180,28 @@ kernel void kernel_gemm_moe_q5_k_f32_ns(
138
180
  return;
139
181
  }
140
182
 
183
+ // Ragged tile-skip: when is_ragged and the upper 16 token-slots of this tile are all
184
+ // padding (router 0xFFFFFFFF), skip the second (reg_c.hi) dotx16_reduce8 half -> ~half
185
+ // the GEMM dot for sparse tiles. Numerically identical (the skipped lanes are padding).
186
+ // Ragged tile-skip: tokens are packed contiguously per expert (moe_scatter fills
187
+ // lanes 0..V-1, moe_fill pre-pads the rest), so router padding (0xFFFFFFFF) is always
188
+ // trailing. Find the valid-token count V and round it UP to the skip granularity
189
+ // skip_gran (columns per skip-group: 8 = quarter, 16 = half/legacy, 32 = disabled).
190
+ // A 8-column group g is all-padding iff its first column (8*g) >= n_active, so its
191
+ // dotx8_reduce4 is skipped. Numerically identical (skipped lanes are padding).
192
+ uint n_active = TILESIZE_N;
193
+ if (is_ragged && skip_gran < TILESIZE_N) {
194
+ uint n_valid = TILESIZE_N;
195
+ for (uint _t = 0; _t < TILESIZE_N; ++_t) {
196
+ if (src2[block_id_n * TILESIZE_N + _t] == 0xFFFFFFFFu) { n_valid = _t; break; }
197
+ }
198
+ n_active = min((uint)TILESIZE_N, ((n_valid + skip_gran - 1) / skip_gran) * skip_gran);
199
+ }
200
+ // Group 0 (cols 0-7) always runs; groups 1-3 skip when fully padding.
201
+ bool skip_g1 = (8u >= n_active);
202
+ bool skip_g2 = (16u >= n_active);
203
+ bool skip_g3 = (24u >= n_active);
204
+
141
205
  __private half16 reg_a;
142
206
  __private float32 reg_c = (float32)(0);
143
207
  __local half4 shared_b[128];
@@ -204,9 +268,11 @@ kernel void kernel_gemm_moe_q5_k_f32_ns(
204
268
 
205
269
  sub_group_barrier(CLK_LOCAL_MEM_FENCE);
206
270
 
207
- half16 acc;
208
- dotx16_reduce8(reg_a, shared_b, reg_c.lo, 0);
209
- dotx16_reduce8(reg_a, shared_b, reg_c.hi, 16);
271
+ half8 acc8;
272
+ dotx8_reduce4(reg_a, shared_b, reg_c.lo.lo, 0);
273
+ if (!skip_g1) { dotx8_reduce4(reg_a, shared_b, reg_c.lo.hi, 8); }
274
+ if (!skip_g2) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.lo, 16); }
275
+ if (!skip_g3) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.hi, 24); }
210
276
 
211
277
  // Second half
212
278
  uint half_step = step + TILESIZE_K;
@@ -226,8 +292,10 @@ kernel void kernel_gemm_moe_q5_k_f32_ns(
226
292
 
227
293
  sub_group_barrier(CLK_LOCAL_MEM_FENCE);
228
294
 
229
- dotx16_reduce8(reg_a, shared_b, reg_c.lo, 0);
230
- dotx16_reduce8(reg_a, shared_b, reg_c.hi, 16);
295
+ dotx8_reduce4(reg_a, shared_b, reg_c.lo.lo, 0);
296
+ if (!skip_g1) { dotx8_reduce4(reg_a, shared_b, reg_c.lo.hi, 8); }
297
+ if (!skip_g2) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.lo, 16); }
298
+ if (!skip_g3) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.hi, 24); }
231
299
  }
232
300
 
233
301
  if ((get_global_id(0) + block_id_m * TILESIZE_M) >= ne01) {
@@ -98,6 +98,46 @@
98
98
  c_reg.lo += convert_float8(acc.lo); \
99
99
  c_reg.hi += convert_float8(acc.hi); \
100
100
 
101
+ // Quarter-tile variant: computes 8 output columns (one skip-group) into a float8
102
+ // accumulator. Same reduction order / flush cadence as dotx16_reduce8, so the
103
+ // non-skipped path is byte-identical; it just lets the caller skip empty
104
+ // 8-column groups at finer granularity. Uses a private half8 `acc8`.
105
+ #define dotx8_reduce4(a_reg, b_lm, c_reg, lm_offset) \
106
+ acc8.s0 = dot(a_reg.s0123, b_lm[lm_offset + 0]); \
107
+ acc8.s1 = dot(a_reg.s0123, b_lm[lm_offset + 1]); \
108
+ acc8.s2 = dot(a_reg.s0123, b_lm[lm_offset + 2]); \
109
+ acc8.s3 = dot(a_reg.s0123, b_lm[lm_offset + 3]); \
110
+ acc8.s4 = dot(a_reg.s0123, b_lm[lm_offset + 4]); \
111
+ acc8.s5 = dot(a_reg.s0123, b_lm[lm_offset + 5]); \
112
+ acc8.s6 = dot(a_reg.s0123, b_lm[lm_offset + 6]); \
113
+ acc8.s7 = dot(a_reg.s0123, b_lm[lm_offset + 7]); \
114
+ acc8.s0 += dot(a_reg.s4567, b_lm[lm_offset + 32]); \
115
+ acc8.s1 += dot(a_reg.s4567, b_lm[lm_offset + 33]); \
116
+ acc8.s2 += dot(a_reg.s4567, b_lm[lm_offset + 34]); \
117
+ acc8.s3 += dot(a_reg.s4567, b_lm[lm_offset + 35]); \
118
+ acc8.s4 += dot(a_reg.s4567, b_lm[lm_offset + 36]); \
119
+ acc8.s5 += dot(a_reg.s4567, b_lm[lm_offset + 37]); \
120
+ acc8.s6 += dot(a_reg.s4567, b_lm[lm_offset + 38]); \
121
+ acc8.s7 += dot(a_reg.s4567, b_lm[lm_offset + 39]); \
122
+ c_reg += convert_float8(acc8); \
123
+ acc8.s0 = dot(a_reg.s89ab, b_lm[lm_offset + 64]); \
124
+ acc8.s1 = dot(a_reg.s89ab, b_lm[lm_offset + 65]); \
125
+ acc8.s2 = dot(a_reg.s89ab, b_lm[lm_offset + 66]); \
126
+ acc8.s3 = dot(a_reg.s89ab, b_lm[lm_offset + 67]); \
127
+ acc8.s4 = dot(a_reg.s89ab, b_lm[lm_offset + 68]); \
128
+ acc8.s5 = dot(a_reg.s89ab, b_lm[lm_offset + 69]); \
129
+ acc8.s6 = dot(a_reg.s89ab, b_lm[lm_offset + 70]); \
130
+ acc8.s7 = dot(a_reg.s89ab, b_lm[lm_offset + 71]); \
131
+ acc8.s0 += dot(a_reg.scdef, b_lm[lm_offset + 96]); \
132
+ acc8.s1 += dot(a_reg.scdef, b_lm[lm_offset + 97]); \
133
+ acc8.s2 += dot(a_reg.scdef, b_lm[lm_offset + 98]); \
134
+ acc8.s3 += dot(a_reg.scdef, b_lm[lm_offset + 99]); \
135
+ acc8.s4 += dot(a_reg.scdef, b_lm[lm_offset + 100]); \
136
+ acc8.s5 += dot(a_reg.scdef, b_lm[lm_offset + 101]); \
137
+ acc8.s6 += dot(a_reg.scdef, b_lm[lm_offset + 102]); \
138
+ acc8.s7 += dot(a_reg.scdef, b_lm[lm_offset + 103]); \
139
+ c_reg += convert_float8(acc8); \
140
+
101
141
 
102
142
  __attribute__((qcom_wave_pair_mode(1)))
103
143
  kernel void kernel_gemm_moe_q6_k_f32_ns(
@@ -111,7 +151,9 @@ kernel void kernel_gemm_moe_q6_k_f32_ns(
111
151
  __write_only image1d_buffer_t dst,
112
152
  __global int * total_tiles,
113
153
  uint ne00,
114
- uint ne01
154
+ uint ne01,
155
+ uint is_ragged,
156
+ uint skip_gran
115
157
  ) {
116
158
  uint block_id_m = get_global_id(1); // m_tile
117
159
  uint block_id_n = get_global_id(2); // n_tile
@@ -121,6 +163,28 @@ kernel void kernel_gemm_moe_q6_k_f32_ns(
121
163
  return;
122
164
  }
123
165
 
166
+ // Ragged tile-skip: when is_ragged and the upper 16 token-slots of this tile are all
167
+ // padding (router 0xFFFFFFFF), skip the second (reg_c.hi) dotx16_reduce8 half -> ~half
168
+ // the GEMM dot for sparse tiles. Numerically identical (the skipped lanes are padding).
169
+ // Ragged tile-skip: tokens are packed contiguously per expert (moe_scatter fills
170
+ // lanes 0..V-1, moe_fill pre-pads the rest), so router padding (0xFFFFFFFF) is always
171
+ // trailing. Find the valid-token count V and round it UP to the skip granularity
172
+ // skip_gran (columns per skip-group: 8 = quarter, 16 = half/legacy, 32 = disabled).
173
+ // A 8-column group g is all-padding iff its first column (8*g) >= n_active, so its
174
+ // dotx8_reduce4 is skipped. Numerically identical (skipped lanes are padding).
175
+ uint n_active = TILESIZE_N;
176
+ if (is_ragged && skip_gran < TILESIZE_N) {
177
+ uint n_valid = TILESIZE_N;
178
+ for (uint _t = 0; _t < TILESIZE_N; ++_t) {
179
+ if (src2[block_id_n * TILESIZE_N + _t] == 0xFFFFFFFFu) { n_valid = _t; break; }
180
+ }
181
+ n_active = min((uint)TILESIZE_N, ((n_valid + skip_gran - 1) / skip_gran) * skip_gran);
182
+ }
183
+ // Group 0 (cols 0-7) always runs; groups 1-3 skip when fully padding.
184
+ bool skip_g1 = (8u >= n_active);
185
+ bool skip_g2 = (16u >= n_active);
186
+ bool skip_g3 = (24u >= n_active);
187
+
124
188
  __private half16 reg_a;
125
189
  __private float32 reg_c = (float32)(0);
126
190
  __local half4 shared_b[128];
@@ -183,9 +247,11 @@ kernel void kernel_gemm_moe_q6_k_f32_ns(
183
247
 
184
248
  sub_group_barrier(CLK_LOCAL_MEM_FENCE);
185
249
 
186
- half16 acc;
187
- dotx16_reduce8(reg_a, shared_b, reg_c.lo, 0);
188
- dotx16_reduce8(reg_a, shared_b, reg_c.hi, 16);
250
+ half8 acc8;
251
+ dotx8_reduce4(reg_a, shared_b, reg_c.lo.lo, 0);
252
+ if (!skip_g1) { dotx8_reduce4(reg_a, shared_b, reg_c.lo.hi, 8); }
253
+ if (!skip_g2) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.lo, 16); }
254
+ if (!skip_g3) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.hi, 24); }
189
255
 
190
256
  // Second half
191
257
  uint half_step = step + TILESIZE_K;
@@ -205,8 +271,10 @@ kernel void kernel_gemm_moe_q6_k_f32_ns(
205
271
 
206
272
  sub_group_barrier(CLK_LOCAL_MEM_FENCE);
207
273
 
208
- dotx16_reduce8(reg_a, shared_b, reg_c.lo, 0);
209
- dotx16_reduce8(reg_a, shared_b, reg_c.hi, 16);
274
+ dotx8_reduce4(reg_a, shared_b, reg_c.lo.lo, 0);
275
+ if (!skip_g1) { dotx8_reduce4(reg_a, shared_b, reg_c.lo.hi, 8); }
276
+ if (!skip_g2) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.lo, 16); }
277
+ if (!skip_g3) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.hi, 24); }
210
278
  }
211
279
 
212
280
  if ((get_global_id(0) + block_id_m * TILESIZE_M) >= ne01) {
@@ -0,0 +1,94 @@
1
+ #pragma OPENCL EXTENSION cl_khr_fp16 : enable
2
+ #pragma OPENCL EXTENSION cl_qcom_reqd_sub_group_size : enable
3
+
4
+ #ifdef cl_qcom_reqd_sub_group_size
5
+ #pragma OPENCL EXTENSION cl_qcom_reqd_sub_group_size : enable
6
+ #define ADRENO_GPU 1
7
+ #define REQD_SUBGROUP_SIZE_128 __attribute__((qcom_reqd_sub_group_size("full")))
8
+ #endif
9
+
10
+ // each work-item computes a 4 (rows of A / m) x 8 (cols of B / n) output tile.
11
+ #ifdef ADRENO_GPU
12
+ REQD_SUBGROUP_SIZE_128
13
+ #endif
14
+ kernel void kernel_gemm_noshuffle_q1_0_f32(
15
+ global const uint * src0_q,
16
+ global const half * src0_d,
17
+ read_only image1d_buffer_t src1,
18
+ global float * dst,
19
+ int k,
20
+ int m,
21
+ int n,
22
+ int n_no_padding,
23
+ ulong offsetd
24
+ ) {
25
+ int n_4 = n >> 2;
26
+
27
+ int gy = get_global_id(0);
28
+ int gx = get_global_id(1);
29
+ int gx_2 = gx << 2;
30
+ dst = (global float *)((global char*)dst + offsetd);
31
+
32
+ half8 c0 = 0, c1 = 0, c2 = 0, c3 = 0;
33
+ half8 B;
34
+
35
+ global const uint* wptr = src0_q + gx_2;
36
+ global const half* sptr = src0_d + gx_2;
37
+
38
+ // 32 weights per uint32, 128 weights (one block / one scale) per 4 uint32.
39
+ for (int i = 0; i < k; i += 32) {
40
+ uint4 pack4 = vload4(0, wptr + (i / 32) * m); // 4 rows, 32 K-values each
41
+ half4 scale = vload4(0, sptr + (i / 128) * m); // 4 rows, one scale per 128
42
+
43
+ for (int j = 0; j < 32; ++j) {
44
+ B.s0123 = read_imageh(src1, gy * 2 + (i + j) * n_4);
45
+ B.s4567 = read_imageh(src1, gy * 2 + (i + j) * n_4 + 1);
46
+
47
+ // sign bit -> +-1 (half arithmetic avoids unsigned underflow)
48
+ half4 wj = (half4)(
49
+ 2.0h * (half)((pack4.s0 >> j) & 1u) - 1.0h,
50
+ 2.0h * (half)((pack4.s1 >> j) & 1u) - 1.0h,
51
+ 2.0h * (half)((pack4.s2 >> j) & 1u) - 1.0h,
52
+ 2.0h * (half)((pack4.s3 >> j) & 1u) - 1.0h) * scale;
53
+
54
+ c0 += B * wj.s0;
55
+ c1 += B * wj.s1;
56
+ c2 += B * wj.s2;
57
+ c3 += B * wj.s3;
58
+ }
59
+ }
60
+
61
+ int idx = (gy << 3) * m + (gx << 2);
62
+
63
+ if(idx+3 < m*n_no_padding){
64
+ vstore4((float4)(c0.s0, c1.s0, c2.s0, c3.s0), 0, dst + idx);
65
+ idx += m;
66
+ }
67
+ if(idx+3 < m*n_no_padding){
68
+ vstore4((float4)(c0.s1, c1.s1, c2.s1, c3.s1), 0, dst + idx);
69
+ idx += m;
70
+ }
71
+ if(idx+3 < m*n_no_padding){
72
+ vstore4((float4)(c0.s2, c1.s2, c2.s2, c3.s2), 0, dst + idx);
73
+ idx += m;
74
+ }
75
+ if(idx+3 < m*n_no_padding){
76
+ vstore4((float4)(c0.s3, c1.s3, c2.s3, c3.s3), 0, dst + idx);
77
+ idx += m;
78
+ }
79
+ if(idx+3 < m*n_no_padding){
80
+ vstore4((float4)(c0.s4, c1.s4, c2.s4, c3.s4), 0, dst + idx);
81
+ idx += m;
82
+ }
83
+ if(idx+3 < m*n_no_padding){
84
+ vstore4((float4)(c0.s5, c1.s5, c2.s5, c3.s5), 0, dst + idx);
85
+ idx += m;
86
+ }
87
+ if(idx+3 < m*n_no_padding){
88
+ vstore4((float4)(c0.s6, c1.s6, c2.s6, c3.s6), 0, dst + idx);
89
+ idx += m;
90
+ }
91
+ if(idx+3 < m*n_no_padding){
92
+ vstore4((float4)(c0.s7, c1.s7, c2.s7, c3.s7), 0, dst + idx);
93
+ }
94
+ }
@@ -0,0 +1,121 @@
1
+ #pragma OPENCL EXTENSION cl_khr_fp16 : enable
2
+ #pragma OPENCL EXTENSION cl_khr_subgroups : enable
3
+
4
+ #ifdef cl_qcom_reqd_sub_group_size
5
+ #pragma OPENCL EXTENSION cl_qcom_reqd_sub_group_size : enable
6
+ #define ADRENO_GPU 1
7
+ #define REQD_SUBGROUP_SIZE_64 __attribute__((qcom_reqd_sub_group_size("half")))
8
+ #endif
9
+
10
+ #define QK1_0 128
11
+ #define N_SIMDGROUP 4
12
+
13
+ #define dequantizeBlockAccum_q1(total, bits, scale, regB, lb) \
14
+ total += (2.0f*(float)((bits >> 0) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s0, lb+0); \
15
+ total += (2.0f*(float)((bits >> 1) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s1, lb+0); \
16
+ total += (2.0f*(float)((bits >> 2) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s2, lb+0); \
17
+ total += (2.0f*(float)((bits >> 3) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s3, lb+0); \
18
+ total += (2.0f*(float)((bits >> 4) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s4, lb+0); \
19
+ total += (2.0f*(float)((bits >> 5) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s5, lb+0); \
20
+ total += (2.0f*(float)((bits >> 6) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s6, lb+0); \
21
+ total += (2.0f*(float)((bits >> 7) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s7, lb+0); \
22
+ total += (2.0f*(float)((bits >> 8) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s0, lb+1); \
23
+ total += (2.0f*(float)((bits >> 9) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s1, lb+1); \
24
+ total += (2.0f*(float)((bits >> 10) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s2, lb+1); \
25
+ total += (2.0f*(float)((bits >> 11) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s3, lb+1); \
26
+ total += (2.0f*(float)((bits >> 12) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s4, lb+1); \
27
+ total += (2.0f*(float)((bits >> 13) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s5, lb+1); \
28
+ total += (2.0f*(float)((bits >> 14) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s6, lb+1); \
29
+ total += (2.0f*(float)((bits >> 15) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s7, lb+1); \
30
+ total += (2.0f*(float)((bits >> 16) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s0, lb+2); \
31
+ total += (2.0f*(float)((bits >> 17) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s1, lb+2); \
32
+ total += (2.0f*(float)((bits >> 18) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s2, lb+2); \
33
+ total += (2.0f*(float)((bits >> 19) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s3, lb+2); \
34
+ total += (2.0f*(float)((bits >> 20) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s4, lb+2); \
35
+ total += (2.0f*(float)((bits >> 21) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s5, lb+2); \
36
+ total += (2.0f*(float)((bits >> 22) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s6, lb+2); \
37
+ total += (2.0f*(float)((bits >> 23) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s7, lb+2); \
38
+ total += (2.0f*(float)((bits >> 24) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s0, lb+3); \
39
+ total += (2.0f*(float)((bits >> 25) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s1, lb+3); \
40
+ total += (2.0f*(float)((bits >> 26) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s2, lb+3); \
41
+ total += (2.0f*(float)((bits >> 27) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s3, lb+3); \
42
+ total += (2.0f*(float)((bits >> 28) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s4, lb+3); \
43
+ total += (2.0f*(float)((bits >> 29) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s5, lb+3); \
44
+ total += (2.0f*(float)((bits >> 30) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s6, lb+3); \
45
+ total += (2.0f*(float)((bits >> 31) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s7, lb+3);
46
+
47
+
48
+ #ifdef ADRENO_GPU
49
+ REQD_SUBGROUP_SIZE_64
50
+ #endif
51
+ __kernel void kernel_gemv_noshuffle_q1_0_f32(
52
+ read_only image1d_buffer_t src0_q,
53
+ global half * src0_d,
54
+ read_only image1d_buffer_t src1,
55
+ ulong offset1,
56
+ global float * dst,
57
+ ulong offsetd,
58
+ int ne00,
59
+ int ne01,
60
+ int ne02,
61
+ int ne10,
62
+ int ne12,
63
+ int ne0,
64
+ int ne1,
65
+ int r2,
66
+ int r3)
67
+ {
68
+ uint groupId = get_local_id(1);
69
+ uint gid = get_global_id(0);
70
+ ushort slid = get_sub_group_local_id();
71
+
72
+ uint K = ne00;
73
+ uint M = ne01;
74
+
75
+ uint LINE_STRIDE_A = M;
76
+ uint BLOCK_STRIDE_A = 4 * M;
77
+
78
+ uint4 regA;
79
+ half regS;
80
+ float8 regB;
81
+
82
+ float totalSum = 0.0f;
83
+
84
+ #pragma unroll 1
85
+ for (uint kb = groupId; kb < (K / QK1_0); kb += N_SIMDGROUP) {
86
+ regS = src0_d[gid + kb * LINE_STRIDE_A]; // each fiber loads its row's scale
87
+
88
+ // first 16 fibers load 8 B values each -> 128 activations for this block
89
+ if (slid < 16) {
90
+ regB.s0123 = read_imagef(src1, (slid * 2 + kb * 32));
91
+ regB.s4567 = read_imagef(src1, (1 + slid * 2 + kb * 32));
92
+ }
93
+
94
+ // load this row's 4 uint32 (128 sign bits)
95
+ regA.s0 = read_imageui(src0_q, (gid + kb * BLOCK_STRIDE_A + LINE_STRIDE_A * 0)).x;
96
+ regA.s1 = read_imageui(src0_q, (gid + kb * BLOCK_STRIDE_A + LINE_STRIDE_A * 1)).x;
97
+ regA.s2 = read_imageui(src0_q, (gid + kb * BLOCK_STRIDE_A + LINE_STRIDE_A * 2)).x;
98
+ regA.s3 = read_imageui(src0_q, (gid + kb * BLOCK_STRIDE_A + LINE_STRIDE_A * 3)).x;
99
+
100
+ float scale = (float)regS;
101
+ dequantizeBlockAccum_q1(totalSum, regA.s0, scale, regB, 0);
102
+ dequantizeBlockAccum_q1(totalSum, regA.s1, scale, regB, 4);
103
+ dequantizeBlockAccum_q1(totalSum, regA.s2, scale, regB, 8);
104
+ dequantizeBlockAccum_q1(totalSum, regA.s3, scale, regB, 12);
105
+ }
106
+
107
+ // reduction in local memory, assumes #wave = N_SIMDGROUP = 4
108
+ local float reduceLM[SIMDGROUP_WIDTH * 3];
109
+ if (groupId == 1) reduceLM[SIMDGROUP_WIDTH * 0 + slid] = totalSum;
110
+ if (groupId == 2) reduceLM[SIMDGROUP_WIDTH * 1 + slid] = totalSum;
111
+ if (groupId == 3) reduceLM[SIMDGROUP_WIDTH * 2 + slid] = totalSum;
112
+ barrier(CLK_LOCAL_MEM_FENCE);
113
+ if (groupId == 0) totalSum += reduceLM[SIMDGROUP_WIDTH * 0 + slid];
114
+ if (groupId == 0) totalSum += reduceLM[SIMDGROUP_WIDTH * 1 + slid];
115
+ if (groupId == 0) totalSum += reduceLM[SIMDGROUP_WIDTH * 2 + slid];
116
+
117
+ if (groupId == 0) {
118
+ dst = (global float*)((global char*)dst + offsetd);
119
+ dst[gid] = totalSum;
120
+ }
121
+ }