whispercpp 1.3.6 → 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 (965) hide show
  1. checksums.yaml +4 -4
  2. data/.document +3 -0
  3. data/.rdoc_options +2 -0
  4. data/README.md +43 -9
  5. data/Rakefile +18 -3
  6. data/ext/dependencies.rb +10 -4
  7. data/ext/dependencies_for_windows.rb +17 -0
  8. data/ext/extconf.rb +20 -8
  9. data/ext/options.rb +54 -14
  10. data/ext/options_for_windows.rb +51 -0
  11. data/ext/ruby_whisper.c +35 -42
  12. data/ext/ruby_whisper.h +141 -0
  13. data/ext/ruby_whisper_context.c +157 -29
  14. data/ext/ruby_whisper_log_queue.c +180 -0
  15. data/ext/ruby_whisper_log_settable.h +46 -0
  16. data/ext/ruby_whisper_parakeet.c +49 -0
  17. data/ext/ruby_whisper_parakeet_context.c +304 -0
  18. data/ext/ruby_whisper_parakeet_context_params.c +117 -0
  19. data/ext/ruby_whisper_parakeet_model.c +84 -0
  20. data/ext/ruby_whisper_parakeet_params.c +548 -0
  21. data/ext/ruby_whisper_parakeet_segment.c +157 -0
  22. data/ext/ruby_whisper_parakeet_token.c +188 -0
  23. data/ext/ruby_whisper_parakeet_transcribe.cpp +58 -0
  24. data/ext/ruby_whisper_params.c +265 -73
  25. data/ext/ruby_whisper_segment.c +6 -6
  26. data/ext/ruby_whisper_transcribe.cpp +23 -15
  27. data/ext/ruby_whisper_vad_context.c +30 -10
  28. data/ext/ruby_whisper_vad_context_detect.cpp +8 -9
  29. data/ext/ruby_whisper_vad_params.c +4 -4
  30. data/ext/ruby_whisper_vad_segment.c +2 -2
  31. data/ext/sources/CMakeLists.txt +42 -3
  32. data/ext/sources/CMakePresets.json +95 -0
  33. data/ext/sources/cmake/parakeet-config.cmake.in +30 -0
  34. data/ext/sources/cmake/parakeet.pc.in +10 -0
  35. data/ext/sources/cmake/whisper.pc.in +2 -2
  36. data/ext/sources/examples/CMakeLists.txt +4 -2
  37. data/ext/sources/examples/bench/bench.cpp +1 -1
  38. data/ext/sources/examples/cli/cli.cpp +52 -10
  39. data/ext/sources/examples/common-ggml.cpp +4 -0
  40. data/ext/sources/examples/common-whisper.cpp +139 -67
  41. data/ext/sources/examples/common-whisper.h +11 -0
  42. data/ext/sources/examples/ffmpeg-transcode.cpp +211 -341
  43. data/ext/sources/examples/parakeet-cli/CMakeLists.txt +8 -0
  44. data/ext/sources/examples/parakeet-cli/parakeet-cli.cpp +243 -0
  45. data/ext/sources/examples/parakeet-quantize/CMakeLists.txt +7 -0
  46. data/ext/sources/examples/parakeet-quantize/parakeet-quantize.cpp +230 -0
  47. data/ext/sources/examples/server/server.cpp +199 -163
  48. data/ext/sources/examples/vad-speech-segments/speech.cpp +3 -2
  49. data/ext/sources/ggml/CMakeLists.txt +21 -14
  50. data/ext/sources/ggml/cmake/FindNCCL.cmake +36 -0
  51. data/ext/sources/ggml/cmake/ggml-config.cmake.in +12 -2
  52. data/ext/sources/ggml/include/ggml-alloc.h +1 -0
  53. data/ext/sources/ggml/include/ggml-backend.h +72 -10
  54. data/ext/sources/ggml/include/ggml-cuda.h +2 -2
  55. data/ext/sources/ggml/include/ggml-rpc.h +3 -3
  56. data/ext/sources/ggml/include/ggml-sycl.h +8 -0
  57. data/ext/sources/ggml/include/ggml.h +103 -9
  58. data/ext/sources/ggml/include/gguf.h +10 -2
  59. data/ext/sources/ggml/src/CMakeLists.txt +30 -6
  60. data/ext/sources/ggml/src/ggml-alloc.c +5 -1
  61. data/ext/sources/ggml/src/ggml-backend-impl.h +22 -2
  62. data/ext/sources/ggml/src/ggml-backend-meta.cpp +2266 -0
  63. data/ext/sources/ggml/src/ggml-backend-reg.cpp +12 -0
  64. data/ext/sources/ggml/src/ggml-backend.cpp +110 -9
  65. data/ext/sources/ggml/src/ggml-blas/ggml-blas.cpp +4 -0
  66. data/ext/sources/ggml/src/ggml-cann/aclnn_ops.cpp +672 -257
  67. data/ext/sources/ggml/src/ggml-cann/aclnn_ops.h +71 -0
  68. data/ext/sources/ggml/src/ggml-cann/common.h +20 -10
  69. data/ext/sources/ggml/src/ggml-cann/ggml-cann.cpp +211 -30
  70. data/ext/sources/ggml/src/ggml-common.h +24 -2
  71. data/ext/sources/ggml/src/ggml-cpu/CMakeLists.txt +59 -30
  72. data/ext/sources/ggml/src/ggml-cpu/amx/amx.cpp +2 -0
  73. data/ext/sources/ggml/src/ggml-cpu/amx/mmq.cpp +21 -22
  74. data/ext/sources/ggml/src/ggml-cpu/arch/arm/quants.c +194 -11
  75. data/ext/sources/ggml/src/ggml-cpu/arch/arm/repack.cpp +65 -0
  76. data/ext/sources/ggml/src/ggml-cpu/arch/loongarch/quants.c +151 -1
  77. data/ext/sources/ggml/src/ggml-cpu/arch/powerpc/quants.c +0 -1
  78. data/ext/sources/ggml/src/ggml-cpu/arch/riscv/quants.c +4279 -1292
  79. data/ext/sources/ggml/src/ggml-cpu/arch/riscv/repack.cpp +5 -35
  80. data/ext/sources/ggml/src/ggml-cpu/arch/s390/quants.c +0 -1
  81. data/ext/sources/ggml/src/ggml-cpu/arch/wasm/quants.c +72 -1
  82. data/ext/sources/ggml/src/ggml-cpu/arch/x86/quants.c +319 -31
  83. data/ext/sources/ggml/src/ggml-cpu/arch/x86/repack.cpp +1 -1
  84. data/ext/sources/ggml/src/ggml-cpu/arch-fallback.h +12 -2
  85. data/ext/sources/ggml/src/ggml-cpu/cmake/FindSMTIME.cmake +32 -0
  86. data/ext/sources/ggml/src/ggml-cpu/ggml-cpu-impl.h +10 -0
  87. data/ext/sources/ggml/src/ggml-cpu/ggml-cpu.c +109 -5
  88. data/ext/sources/ggml/src/ggml-cpu/ggml-cpu.cpp +2 -0
  89. data/ext/sources/ggml/src/ggml-cpu/kleidiai/kleidiai.cpp +146 -134
  90. data/ext/sources/ggml/src/ggml-cpu/llamafile/sgemm.cpp +107 -82
  91. data/ext/sources/ggml/src/ggml-cpu/ops.cpp +501 -119
  92. data/ext/sources/ggml/src/ggml-cpu/ops.h +3 -0
  93. data/ext/sources/ggml/src/ggml-cpu/quants.c +106 -0
  94. data/ext/sources/ggml/src/ggml-cpu/quants.h +6 -0
  95. data/ext/sources/ggml/src/ggml-cpu/repack.cpp +3 -0
  96. data/ext/sources/ggml/src/ggml-cpu/simd-gemm.h +91 -1
  97. data/ext/sources/ggml/src/ggml-cpu/simd-mappings.h +14 -16
  98. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime.cpp +1402 -687
  99. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime.h +8 -0
  100. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime1_kernels.cpp +597 -2766
  101. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime2_kernels.cpp +5768 -0
  102. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime_env.cpp +320 -0
  103. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime_env.h +55 -0
  104. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime_kernels.h +182 -19
  105. data/ext/sources/ggml/src/ggml-cpu/spacemit/repack.cpp +1795 -0
  106. data/ext/sources/ggml/src/ggml-cpu/spacemit/repack.h +14 -0
  107. data/ext/sources/ggml/src/ggml-cpu/spacemit/rvv_kernels.cpp +3178 -0
  108. data/ext/sources/ggml/src/ggml-cpu/spacemit/rvv_kernels.h +95 -0
  109. data/ext/sources/ggml/src/ggml-cpu/spacemit/spine_barrier.h +34 -0
  110. data/ext/sources/ggml/src/ggml-cpu/spacemit/spine_mem_pool.cpp +760 -0
  111. data/ext/sources/ggml/src/ggml-cpu/spacemit/spine_mem_pool.h +32 -0
  112. data/ext/sources/ggml/src/ggml-cpu/spacemit/spine_tcm.h +409 -0
  113. data/ext/sources/ggml/src/ggml-cpu/vec.cpp +39 -55
  114. data/ext/sources/ggml/src/ggml-cpu/vec.h +225 -240
  115. data/ext/sources/ggml/src/ggml-cuda/CMakeLists.txt +17 -7
  116. data/ext/sources/ggml/src/ggml-cuda/allreduce.cu +971 -0
  117. data/ext/sources/ggml/src/ggml-cuda/allreduce.cuh +29 -0
  118. data/ext/sources/ggml/src/ggml-cuda/argsort.cu +62 -26
  119. data/ext/sources/ggml/src/ggml-cuda/binbcast.cu +134 -64
  120. data/ext/sources/ggml/src/ggml-cuda/binbcast.cuh +1 -0
  121. data/ext/sources/ggml/src/ggml-cuda/col2im-1d.cu +81 -0
  122. data/ext/sources/ggml/src/ggml-cuda/col2im-1d.cuh +3 -0
  123. data/ext/sources/ggml/src/ggml-cuda/common.cuh +246 -28
  124. data/ext/sources/ggml/src/ggml-cuda/concat.cu +134 -116
  125. data/ext/sources/ggml/src/ggml-cuda/conv-transpose-1d.cu +14 -12
  126. data/ext/sources/ggml/src/ggml-cuda/conv2d-transpose.cu +45 -21
  127. data/ext/sources/ggml/src/ggml-cuda/conv2d-transpose.cuh +1 -0
  128. data/ext/sources/ggml/src/ggml-cuda/convert.cu +139 -34
  129. data/ext/sources/ggml/src/ggml-cuda/convert.cuh +10 -0
  130. data/ext/sources/ggml/src/ggml-cuda/cpy.cu +88 -29
  131. data/ext/sources/ggml/src/ggml-cuda/dequantize.cuh +22 -0
  132. data/ext/sources/ggml/src/ggml-cuda/fattn-common.cuh +287 -49
  133. data/ext/sources/ggml/src/ggml-cuda/fattn-mma-f16.cuh +335 -130
  134. data/ext/sources/ggml/src/ggml-cuda/fattn-tile.cu +12 -0
  135. data/ext/sources/ggml/src/ggml-cuda/fattn-tile.cuh +127 -24
  136. data/ext/sources/ggml/src/ggml-cuda/fattn-vec.cuh +40 -15
  137. data/ext/sources/ggml/src/ggml-cuda/fattn-wmma-f16.cu +18 -9
  138. data/ext/sources/ggml/src/ggml-cuda/fattn.cu +169 -60
  139. data/ext/sources/ggml/src/ggml-cuda/fattn.cuh +2 -0
  140. data/ext/sources/ggml/src/ggml-cuda/fwht.cu +101 -0
  141. data/ext/sources/ggml/src/ggml-cuda/fwht.cuh +4 -0
  142. data/ext/sources/ggml/src/ggml-cuda/gated_delta_net.cu +109 -45
  143. data/ext/sources/ggml/src/ggml-cuda/gated_delta_net.cuh +10 -0
  144. data/ext/sources/ggml/src/ggml-cuda/getrows.cu +48 -23
  145. data/ext/sources/ggml/src/ggml-cuda/ggml-cuda.cu +2034 -2104
  146. data/ext/sources/ggml/src/ggml-cuda/im2col.cu +32 -29
  147. data/ext/sources/ggml/src/ggml-cuda/mean.cu +4 -2
  148. data/ext/sources/ggml/src/ggml-cuda/mma.cuh +242 -195
  149. data/ext/sources/ggml/src/ggml-cuda/mmf.cuh +3 -3
  150. data/ext/sources/ggml/src/ggml-cuda/mmq.cu +25 -12
  151. data/ext/sources/ggml/src/ggml-cuda/mmq.cuh +502 -423
  152. data/ext/sources/ggml/src/ggml-cuda/mmvf.cu +19 -12
  153. data/ext/sources/ggml/src/ggml-cuda/mmvq.cu +562 -97
  154. data/ext/sources/ggml/src/ggml-cuda/mmvq.cuh +6 -1
  155. data/ext/sources/ggml/src/ggml-cuda/norm.cu +36 -10
  156. data/ext/sources/ggml/src/ggml-cuda/out-prod.cu +66 -7
  157. data/ext/sources/ggml/src/ggml-cuda/quantize.cu +133 -26
  158. data/ext/sources/ggml/src/ggml-cuda/quantize.cuh +1 -1
  159. data/ext/sources/ggml/src/ggml-cuda/reduce_rows.cuh +5 -1
  160. data/ext/sources/ggml/src/ggml-cuda/rope.cu +11 -4
  161. data/ext/sources/ggml/src/ggml-cuda/scale.cu +4 -1
  162. data/ext/sources/ggml/src/ggml-cuda/set-rows.cu +78 -10
  163. data/ext/sources/ggml/src/ggml-cuda/snake.cu +72 -0
  164. data/ext/sources/ggml/src/ggml-cuda/snake.cuh +8 -0
  165. data/ext/sources/ggml/src/ggml-cuda/softcap.cu +4 -1
  166. data/ext/sources/ggml/src/ggml-cuda/ssm-conv.cu +45 -13
  167. data/ext/sources/ggml/src/ggml-cuda/ssm-conv.cuh +1 -1
  168. data/ext/sources/ggml/src/ggml-cuda/ssm-scan.cu +40 -18
  169. data/ext/sources/ggml/src/ggml-cuda/sumrows.cu +8 -4
  170. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_1-ncols2_16.cu +1 -0
  171. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_1-ncols2_32.cu +1 -0
  172. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_1-ncols2_8.cu +2 -0
  173. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_16-ncols2_2.cu +1 -0
  174. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_16-ncols2_4.cu +1 -0
  175. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_2-ncols2_16.cu +1 -0
  176. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_2-ncols2_32.cu +1 -0
  177. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_2-ncols2_4.cu +1 -0
  178. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_2-ncols2_8.cu +2 -0
  179. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_32-ncols2_2.cu +1 -0
  180. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_16.cu +1 -0
  181. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_2.cu +1 -0
  182. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_4.cu +1 -0
  183. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_8.cu +2 -0
  184. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_8-ncols2_2.cu +1 -0
  185. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_8-ncols2_4.cu +1 -0
  186. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_8-ncols2_8.cu +2 -0
  187. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-tile-instance-dkq192-dv128.cu +5 -0
  188. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-tile-instance-dkq320-dv256.cu +5 -0
  189. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-tile-instance-dkq512-dv512.cu +5 -0
  190. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-bf16.cu +7 -0
  191. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-f16.cu +7 -0
  192. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q4_0.cu +7 -0
  193. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q4_1.cu +7 -0
  194. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q5_0.cu +7 -0
  195. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q5_1.cu +7 -0
  196. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q8_0.cu +7 -0
  197. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-f16-bf16.cu +7 -0
  198. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q4_0-bf16.cu +7 -0
  199. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q4_1-bf16.cu +7 -0
  200. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q5_0-bf16.cu +7 -0
  201. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q5_1-bf16.cu +7 -0
  202. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q8_0-bf16.cu +7 -0
  203. data/ext/sources/ggml/src/ggml-cuda/template-instances/mmq-instance-nvfp4.cu +5 -0
  204. data/ext/sources/ggml/src/ggml-cuda/template-instances/mmq-instance-q1_0.cu +5 -0
  205. data/ext/sources/ggml/src/ggml-cuda/top-k.cu +5 -4
  206. data/ext/sources/ggml/src/ggml-cuda/topk-moe.cu +33 -24
  207. data/ext/sources/ggml/src/ggml-cuda/unary.cu +31 -2
  208. data/ext/sources/ggml/src/ggml-cuda/unary.cuh +2 -0
  209. data/ext/sources/ggml/src/ggml-cuda/vecdotq.cuh +80 -0
  210. data/ext/sources/ggml/src/ggml-cuda/vendors/cuda.h +7 -2
  211. data/ext/sources/ggml/src/ggml-cuda/vendors/hip.h +23 -4
  212. data/ext/sources/ggml/src/ggml-cuda/vendors/musa.h +4 -0
  213. data/ext/sources/ggml/src/ggml-hexagon/CMakeLists.txt +1 -5
  214. data/ext/sources/ggml/src/ggml-hexagon/ggml-hexagon.cpp +2788 -1762
  215. data/ext/sources/ggml/src/ggml-hexagon/htp/CMakeLists.txt +13 -4
  216. data/ext/sources/ggml/src/ggml-hexagon/htp/act-ops.c +53 -84
  217. data/ext/sources/ggml/src/ggml-hexagon/htp/argsort-ops.c +25 -12
  218. data/ext/sources/ggml/src/ggml-hexagon/htp/binary-ops.c +165 -184
  219. data/ext/sources/ggml/src/ggml-hexagon/htp/cmake-toolchain.cmake +17 -19
  220. data/ext/sources/ggml/src/ggml-hexagon/htp/concat-ops.c +277 -0
  221. data/ext/sources/ggml/src/ggml-hexagon/htp/cpy-ops.c +170 -127
  222. data/ext/sources/ggml/src/ggml-hexagon/htp/cumsum-ops.c +270 -0
  223. data/ext/sources/ggml/src/ggml-hexagon/htp/diag-ops.c +216 -0
  224. data/ext/sources/ggml/src/ggml-hexagon/htp/fill-ops.c +123 -0
  225. data/ext/sources/ggml/src/ggml-hexagon/htp/flash-attn-ops.c +1774 -396
  226. data/ext/sources/ggml/src/ggml-hexagon/htp/flash-attn-ops.h +303 -0
  227. data/ext/sources/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.c +1148 -0
  228. data/ext/sources/ggml/src/ggml-hexagon/htp/get-rows-ops.c +148 -42
  229. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-common.h +80 -0
  230. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-dma.c +2 -2
  231. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-dma.h +255 -62
  232. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-dump.h +9 -0
  233. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-profile.h +64 -0
  234. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-utils.h +25 -21
  235. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-fa-kernels.h +555 -0
  236. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-mm-kernels-tiled.h +1303 -0
  237. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-queue.c +167 -0
  238. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-queue.h +157 -0
  239. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-utils.h +222 -0
  240. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-ctx.h +104 -13
  241. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-ops.h +222 -57
  242. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-vtcm.h +19 -0
  243. data/ext/sources/ggml/src/ggml-hexagon/htp/htp_iface.idl +10 -3
  244. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-base.h +78 -26
  245. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-copy.h +27 -10
  246. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-div.h +63 -23
  247. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-exp.h +48 -8
  248. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-fa-kernels.h +232 -0
  249. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-flash-attn.h +47 -0
  250. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-log.h +65 -0
  251. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-flat.h +1511 -0
  252. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-tiled.h +1200 -0
  253. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-pow.h +42 -0
  254. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-repl.h +74 -0
  255. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-sigmoid.h +40 -0
  256. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-sin-cos.h +90 -0
  257. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-utils.h +5 -8
  258. data/ext/sources/ggml/src/ggml-hexagon/htp/main.c +625 -816
  259. data/ext/sources/ggml/src/ggml-hexagon/htp/matmul-ops.c +3052 -2166
  260. data/ext/sources/ggml/src/ggml-hexagon/htp/matmul-ops.h +650 -0
  261. data/ext/sources/ggml/src/ggml-hexagon/htp/pad-ops.c +547 -0
  262. data/ext/sources/ggml/src/ggml-hexagon/htp/repeat-ops.c +148 -0
  263. data/ext/sources/ggml/src/ggml-hexagon/htp/rope-ops.c +337 -106
  264. data/ext/sources/ggml/src/ggml-hexagon/htp/set-rows-ops.c +59 -37
  265. data/ext/sources/ggml/src/ggml-hexagon/htp/softmax-ops.c +121 -133
  266. data/ext/sources/ggml/src/ggml-hexagon/htp/solve-tri-ops.c +267 -0
  267. data/ext/sources/ggml/src/ggml-hexagon/htp/ssm-conv.c +245 -151
  268. data/ext/sources/ggml/src/ggml-hexagon/htp/sum-rows-ops.c +6 -6
  269. data/ext/sources/ggml/src/ggml-hexagon/htp/unary-ops.c +719 -45
  270. data/ext/sources/ggml/src/ggml-hexagon/htp/worker-pool.c +15 -3
  271. data/ext/sources/ggml/src/ggml-hexagon/htp/worker-pool.h +8 -0
  272. data/ext/sources/ggml/src/ggml-hexagon/htp-opnode.h +390 -0
  273. data/ext/sources/ggml/src/ggml-hexagon/libggml-htp.inf +3 -5
  274. data/ext/sources/ggml/src/ggml-hip/CMakeLists.txt +27 -9
  275. data/ext/sources/ggml/src/ggml-impl.h +6 -1
  276. data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.cpp +207 -18
  277. data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.h +36 -2
  278. data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.m +186 -29
  279. data/ext/sources/ggml/src/ggml-metal/ggml-metal-impl.h +118 -0
  280. data/ext/sources/ggml/src/ggml-metal/ggml-metal-ops.cpp +322 -21
  281. data/ext/sources/ggml/src/ggml-metal/ggml-metal-ops.h +4 -0
  282. data/ext/sources/ggml/src/ggml-metal/ggml-metal.cpp +39 -26
  283. data/ext/sources/ggml/src/ggml-metal/ggml-metal.metal +1226 -467
  284. data/ext/sources/ggml/src/ggml-musa/CMakeLists.txt +5 -6
  285. data/ext/sources/ggml/src/ggml-opencl/CMakeLists.txt +67 -5
  286. data/ext/sources/ggml/src/ggml-opencl/fa_tune.h +92 -0
  287. data/ext/sources/ggml/src/ggml-opencl/ggml-opencl.cpp +16290 -6246
  288. data/ext/sources/ggml/src/ggml-opencl/kernels/concat.cl +67 -0
  289. data/ext/sources/ggml/src/ggml-opencl/kernels/cpy.cl +59 -0
  290. data/ext/sources/ggml/src/ggml-opencl/kernels/cvt.cl +1997 -92
  291. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f16.cl +81 -41
  292. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32.cl +88 -39
  293. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_f16.cl +1995 -96
  294. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_q4_0.cl +1615 -0
  295. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_q8_0.cl +1486 -0
  296. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_pre_f16.cl +156 -0
  297. data/ext/sources/ggml/src/ggml-opencl/kernels/gated_delta_net.cl +249 -0
  298. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_mxfp4_f32_ns.cl +374 -0
  299. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_0_f32_ns.cl +324 -0
  300. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_1_f32_ns.cl +326 -0
  301. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_k_f32_ns.cl +348 -0
  302. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_0_f32_ns.cl +328 -0
  303. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_1_f32_ns.cl +330 -0
  304. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_k_f32_ns.cl +356 -0
  305. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q6_k_f32_ns.cl +335 -0
  306. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_iq4_nl_f32.cl +150 -0
  307. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q1_0_f32.cl +94 -0
  308. data/ext/sources/ggml/src/ggml-opencl/kernels/{mul_mat_Ab_Bi_8x4.cl → gemm_noshuffle_q4_0_f32.cl} +1 -1
  309. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q4_k_f32.cl +172 -0
  310. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q5_0_f32.cl +131 -0
  311. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q5_1_f32.cl +134 -0
  312. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q5_k_f32.cl +176 -0
  313. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q6_k_f32.cl +140 -0
  314. data/ext/sources/ggml/src/ggml-opencl/kernels/{mul_mm_q8_0_f32_8x4.cl → gemm_noshuffle_q8_0_f32.cl} +1 -1
  315. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_xmem_f16_f32_os8.cl +233 -0
  316. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_mxfp4_f32_ns.cl +165 -0
  317. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q4_0_f32_ns.cl +120 -0
  318. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q4_1_f32_ns.cl +123 -0
  319. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q4_k_f32_ns.cl +155 -0
  320. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q5_0_f32_ns.cl +123 -0
  321. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q5_1_f32_ns.cl +125 -0
  322. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q5_k_f32_ns.cl +160 -0
  323. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q6_k_f32_ns.cl +141 -0
  324. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_iq4_nl_f32.cl +302 -0
  325. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q1_0_f32.cl +121 -0
  326. data/ext/sources/ggml/src/ggml-opencl/kernels/{gemv_noshuffle_general.cl → gemv_noshuffle_q4_0_f32.cl} +5 -5
  327. data/ext/sources/ggml/src/ggml-opencl/kernels/{gemv_noshuffle.cl → gemv_noshuffle_q4_0_f32_spec.cl} +5 -5
  328. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_k_f32.cl +318 -0
  329. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q5_0_f32.cl +291 -0
  330. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q5_1_f32.cl +294 -0
  331. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q5_k_f32.cl +326 -0
  332. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q6_k_f32.cl +293 -0
  333. data/ext/sources/ggml/src/ggml-opencl/kernels/{gemv_noshuffle_general_q8_0_f32.cl → gemv_noshuffle_q8_0_f32.cl} +1 -1
  334. data/ext/sources/ggml/src/ggml-opencl/kernels/get_rows.cl +15 -9
  335. data/ext/sources/ggml/src/ggml-opencl/kernels/moe_reorder_b.cl +30 -0
  336. data/ext/sources/ggml/src/ggml-opencl/kernels/moe_sort_by_expert.cl +82 -0
  337. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_iq4_nl_f32_l4_lm.cl +171 -0
  338. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q1_0_f32_l4_lm.cl +156 -0
  339. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q4_k_f32_l4_lm.cl +179 -0
  340. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q5_0_f32_l4_lm.cl +173 -0
  341. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q5_1_f32_l4_lm.cl +175 -0
  342. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q5_k_f32_l4_lm.cl +192 -0
  343. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_f16_f32_l4.cl +1149 -0
  344. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_iq4_nl_f32.cl +164 -0
  345. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_iq4_nl_f32_flat.cl +202 -0
  346. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q1_0_f32.cl +141 -0
  347. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q1_0_f32_flat.cl +190 -0
  348. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q4_k_f32_flat.cl +196 -0
  349. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_0_f32.cl +241 -0
  350. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_0_f32_flat.cl +243 -0
  351. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_1_f32.cl +243 -0
  352. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_1_f32_flat.cl +247 -0
  353. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_k_f32.cl +187 -0
  354. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_k_f32_flat.cl +203 -0
  355. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q6_k_f32_flat.cl +48 -64
  356. data/ext/sources/ggml/src/ggml-opencl/kernels/norm.cl +5 -2
  357. data/ext/sources/ggml/src/ggml-opencl/kernels/set_rows.cl +500 -0
  358. data/ext/sources/ggml/src/ggml-opencl/libdl.h +79 -0
  359. data/ext/sources/ggml/src/ggml-openvino/.clang-format +0 -5
  360. data/ext/sources/ggml/src/ggml-openvino/CMakeLists.txt +2 -4
  361. data/ext/sources/ggml/src/ggml-openvino/ggml-decoder.cpp +740 -127
  362. data/ext/sources/ggml/src/ggml-openvino/ggml-decoder.h +76 -23
  363. data/ext/sources/ggml/src/ggml-openvino/ggml-openvino-extra.cpp +75 -14
  364. data/ext/sources/ggml/src/ggml-openvino/ggml-openvino-extra.h +29 -8
  365. data/ext/sources/ggml/src/ggml-openvino/ggml-openvino.cpp +339 -69
  366. data/ext/sources/ggml/src/ggml-openvino/ggml-quants.cpp +330 -192
  367. data/ext/sources/ggml/src/ggml-openvino/ggml-quants.h +10 -4
  368. data/ext/sources/ggml/src/ggml-openvino/openvino/decoder.h +56 -16
  369. data/ext/sources/ggml/src/ggml-openvino/openvino/frontend.h +1 -1
  370. data/ext/sources/ggml/src/ggml-openvino/openvino/input_model.h +4 -4
  371. data/ext/sources/ggml/src/ggml-openvino/openvino/node_context.h +94 -37
  372. data/ext/sources/ggml/src/ggml-openvino/openvino/op/add_id.cpp +76 -0
  373. data/ext/sources/ggml/src/ggml-openvino/openvino/op/argsort.cpp +47 -0
  374. data/ext/sources/ggml/src/ggml-openvino/openvino/op/clamp.cpp +33 -0
  375. data/ext/sources/ggml/src/ggml-openvino/openvino/op/concat.cpp +48 -0
  376. data/ext/sources/ggml/src/ggml-openvino/openvino/op/cont.cpp +8 -16
  377. data/ext/sources/ggml/src/ggml-openvino/openvino/op/cpy.cpp +14 -1
  378. data/ext/sources/ggml/src/ggml-openvino/openvino/op/div.cpp +146 -0
  379. data/ext/sources/ggml/src/ggml-openvino/openvino/op/flash_attn_ext.cpp +108 -21
  380. data/ext/sources/ggml/src/ggml-openvino/openvino/op/gated_delta_net.cpp +282 -0
  381. data/ext/sources/ggml/src/ggml-openvino/openvino/op/gated_delta_net.hpp +65 -0
  382. data/ext/sources/ggml/src/ggml-openvino/openvino/op/get_rows.cpp +2 -9
  383. data/ext/sources/ggml/src/ggml-openvino/openvino/op/glu_geglu.cpp +21 -7
  384. data/ext/sources/ggml/src/ggml-openvino/openvino/op/glu_swiglu.cpp +41 -8
  385. data/ext/sources/ggml/src/ggml-openvino/openvino/op/im2col.cpp +120 -0
  386. data/ext/sources/ggml/src/ggml-openvino/openvino/op/l2_norm.cpp +44 -0
  387. data/ext/sources/ggml/src/ggml-openvino/openvino/op/mul_mat_id.cpp +226 -0
  388. data/ext/sources/ggml/src/ggml-openvino/openvino/op/mulmat.cpp +19 -9
  389. data/ext/sources/ggml/src/ggml-openvino/openvino/op/norm.cpp +58 -0
  390. data/ext/sources/ggml/src/ggml-openvino/openvino/op/pad.cpp +95 -0
  391. data/ext/sources/ggml/src/ggml-openvino/openvino/op/permute.cpp +58 -13
  392. data/ext/sources/ggml/src/ggml-openvino/openvino/op/repeat.cpp +74 -0
  393. data/ext/sources/ggml/src/ggml-openvino/openvino/op/reshape.cpp +13 -6
  394. data/ext/sources/ggml/src/ggml-openvino/openvino/op/rms_norm.cpp +1 -1
  395. data/ext/sources/ggml/src/ggml-openvino/openvino/op/rope.cpp +161 -39
  396. data/ext/sources/ggml/src/ggml-openvino/openvino/op/set_rows.cpp +3 -3
  397. data/ext/sources/ggml/src/ggml-openvino/openvino/op/softmax.cpp +126 -49
  398. data/ext/sources/ggml/src/ggml-openvino/openvino/op/ssm_conv.cpp +59 -0
  399. data/ext/sources/ggml/src/ggml-openvino/openvino/op/sum_rows.cpp +27 -0
  400. data/ext/sources/ggml/src/ggml-openvino/openvino/op/transpose.cpp +32 -1
  401. data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_silu.cpp +1 -1
  402. data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_softplus.cpp +38 -0
  403. data/ext/sources/ggml/src/ggml-openvino/openvino/op/view.cpp +90 -25
  404. data/ext/sources/ggml/src/ggml-openvino/openvino/op_table.cpp +41 -22
  405. data/ext/sources/ggml/src/ggml-openvino/openvino/op_table.h +18 -4
  406. data/ext/sources/ggml/src/ggml-openvino/openvino/pass/mark_decompression_convert_constant_folding.h +1 -1
  407. data/ext/sources/ggml/src/ggml-openvino/openvino/rt_info/weightless_caching_attributes.hpp +41 -0
  408. data/ext/sources/ggml/src/ggml-openvino/openvino/translate_session.cpp +70 -43
  409. data/ext/sources/ggml/src/ggml-openvino/openvino/translate_session.h +5 -4
  410. data/ext/sources/ggml/src/ggml-openvino/openvino/utils.cpp +612 -36
  411. data/ext/sources/ggml/src/ggml-openvino/openvino/utils.h +29 -26
  412. data/ext/sources/ggml/src/ggml-openvino/utils.cpp +460 -114
  413. data/ext/sources/ggml/src/ggml-openvino/utils.h +32 -9
  414. data/ext/sources/ggml/src/ggml-opt.cpp +1 -0
  415. data/ext/sources/ggml/src/ggml-quants.c +365 -114
  416. data/ext/sources/ggml/src/ggml-quants.h +6 -0
  417. data/ext/sources/ggml/src/ggml-rpc/CMakeLists.txt +24 -0
  418. data/ext/sources/ggml/src/ggml-rpc/ggml-rpc.cpp +167 -311
  419. data/ext/sources/ggml/src/ggml-rpc/transport.cpp +683 -0
  420. data/ext/sources/ggml/src/ggml-rpc/transport.h +34 -0
  421. data/ext/sources/ggml/src/ggml-sycl/CMakeLists.txt +50 -4
  422. data/ext/sources/ggml/src/ggml-sycl/add-id.cpp +1 -1
  423. data/ext/sources/ggml/src/ggml-sycl/backend.hpp +5 -1
  424. data/ext/sources/ggml/src/ggml-sycl/binbcast.cpp +12 -0
  425. data/ext/sources/ggml/src/ggml-sycl/col2im-1d.cpp +102 -0
  426. data/ext/sources/ggml/src/ggml-sycl/col2im-1d.hpp +8 -0
  427. data/ext/sources/ggml/src/ggml-sycl/common.cpp +72 -2
  428. data/ext/sources/ggml/src/ggml-sycl/common.hpp +59 -2
  429. data/ext/sources/ggml/src/ggml-sycl/concat.cpp +21 -1
  430. data/ext/sources/ggml/src/ggml-sycl/conv2d-dw.cpp +158 -0
  431. data/ext/sources/ggml/src/ggml-sycl/conv2d-dw.hpp +10 -0
  432. data/ext/sources/ggml/src/ggml-sycl/conv2d-transpose.cpp +125 -0
  433. data/ext/sources/ggml/src/ggml-sycl/conv2d-transpose.hpp +10 -0
  434. data/ext/sources/ggml/src/ggml-sycl/conv2d.cpp +150 -0
  435. data/ext/sources/ggml/src/ggml-sycl/conv2d.hpp +10 -0
  436. data/ext/sources/ggml/src/ggml-sycl/conv3d.cpp +224 -0
  437. data/ext/sources/ggml/src/ggml-sycl/conv3d.hpp +8 -0
  438. data/ext/sources/ggml/src/ggml-sycl/convert.cpp +121 -13
  439. data/ext/sources/ggml/src/ggml-sycl/convert.hpp +9 -0
  440. data/ext/sources/ggml/src/ggml-sycl/cpy.cpp +706 -0
  441. data/ext/sources/ggml/src/ggml-sycl/cpy.hpp +281 -0
  442. data/ext/sources/ggml/src/ggml-sycl/cross_entropy_loss.cpp +255 -0
  443. data/ext/sources/ggml/src/ggml-sycl/cross_entropy_loss.hpp +7 -0
  444. data/ext/sources/ggml/src/ggml-sycl/cumsum.cpp +148 -0
  445. data/ext/sources/ggml/src/ggml-sycl/cumsum.hpp +5 -0
  446. data/ext/sources/ggml/src/ggml-sycl/dequantize.hpp +678 -0
  447. data/ext/sources/ggml/src/ggml-sycl/diag.cpp +67 -0
  448. data/ext/sources/ggml/src/ggml-sycl/diag.hpp +5 -0
  449. data/ext/sources/ggml/src/ggml-sycl/dmmv.cpp +997 -244
  450. data/ext/sources/ggml/src/ggml-sycl/dpct/helper.hpp +15 -7
  451. data/ext/sources/ggml/src/ggml-sycl/element_wise.cpp +215 -204
  452. data/ext/sources/ggml/src/ggml-sycl/element_wise.hpp +2 -2
  453. data/ext/sources/ggml/src/ggml-sycl/fattn-buffers.cpp +56 -0
  454. data/ext/sources/ggml/src/ggml-sycl/fattn-buffers.hpp +63 -0
  455. data/ext/sources/ggml/src/ggml-sycl/fattn-common.hpp +7 -5
  456. data/ext/sources/ggml/src/ggml-sycl/fattn-tile.cpp +4 -0
  457. data/ext/sources/ggml/src/ggml-sycl/fattn-tile.hpp +76 -168
  458. data/ext/sources/ggml/src/ggml-sycl/fattn-vec.hpp +7 -0
  459. data/ext/sources/ggml/src/ggml-sycl/fattn.cpp +3 -1
  460. data/ext/sources/ggml/src/ggml-sycl/fill.cpp +55 -0
  461. data/ext/sources/ggml/src/ggml-sycl/fill.hpp +5 -0
  462. data/ext/sources/ggml/src/ggml-sycl/gated_delta_net.cpp +69 -31
  463. data/ext/sources/ggml/src/ggml-sycl/gated_delta_net.hpp +1 -0
  464. data/ext/sources/ggml/src/ggml-sycl/gemm.hpp +3 -0
  465. data/ext/sources/ggml/src/ggml-sycl/getrows.cpp +79 -3
  466. data/ext/sources/ggml/src/ggml-sycl/ggml-sycl.cpp +1758 -455
  467. data/ext/sources/ggml/src/ggml-sycl/im2col.cpp +353 -89
  468. data/ext/sources/ggml/src/ggml-sycl/im2col.hpp +5 -3
  469. data/ext/sources/ggml/src/ggml-sycl/mmvq.cpp +1542 -39
  470. data/ext/sources/ggml/src/ggml-sycl/mmvq.hpp +33 -0
  471. data/ext/sources/ggml/src/ggml-sycl/norm.cpp +103 -49
  472. data/ext/sources/ggml/src/ggml-sycl/outprod.cpp +45 -9
  473. data/ext/sources/ggml/src/ggml-sycl/pad.cpp +27 -27
  474. data/ext/sources/ggml/src/ggml-sycl/pool.cpp +185 -0
  475. data/ext/sources/ggml/src/ggml-sycl/pool.hpp +22 -0
  476. data/ext/sources/ggml/src/ggml-sycl/presets.hpp +3 -1
  477. data/ext/sources/ggml/src/ggml-sycl/quants.hpp +71 -0
  478. data/ext/sources/ggml/src/ggml-sycl/set_rows.cpp +17 -3
  479. data/ext/sources/ggml/src/ggml-sycl/softmax.cpp +9 -10
  480. data/ext/sources/ggml/src/ggml-sycl/solve_tri.cpp +172 -0
  481. data/ext/sources/ggml/src/ggml-sycl/solve_tri.hpp +8 -0
  482. data/ext/sources/ggml/src/ggml-sycl/ssm_conv.cpp +6 -1
  483. data/ext/sources/ggml/src/ggml-sycl/ssm_scan.cpp +156 -0
  484. data/ext/sources/ggml/src/ggml-sycl/ssm_scan.hpp +5 -0
  485. data/ext/sources/ggml/src/ggml-sycl/sycl_hw.cpp +62 -10
  486. data/ext/sources/ggml/src/ggml-sycl/sycl_hw.hpp +18 -6
  487. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-tile-instance-dkq512-dv512.cpp +6 -0
  488. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-f16.cpp +1 -0
  489. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q4_0.cpp +1 -0
  490. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q4_1.cpp +1 -0
  491. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q5_0.cpp +1 -0
  492. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q5_1.cpp +1 -0
  493. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q8_0.cpp +1 -0
  494. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-f16.cpp +1 -0
  495. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q4_0.cpp +1 -0
  496. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q4_1.cpp +1 -0
  497. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q5_0.cpp +1 -0
  498. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q5_1.cpp +1 -0
  499. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q8_0.cpp +1 -0
  500. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-f16.cpp +1 -0
  501. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q4_0.cpp +1 -0
  502. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q4_1.cpp +1 -0
  503. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q5_0.cpp +1 -0
  504. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q5_1.cpp +1 -0
  505. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q8_0.cpp +1 -0
  506. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-f16.cpp +1 -0
  507. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q4_0.cpp +1 -0
  508. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q4_1.cpp +1 -0
  509. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q5_0.cpp +1 -0
  510. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q5_1.cpp +1 -0
  511. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q8_0.cpp +1 -0
  512. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-f16.cpp +1 -0
  513. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q4_0.cpp +1 -0
  514. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q4_1.cpp +1 -0
  515. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q5_0.cpp +1 -0
  516. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q5_1.cpp +1 -0
  517. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q8_0.cpp +1 -0
  518. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-f16.cpp +1 -0
  519. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q4_0.cpp +1 -0
  520. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q4_1.cpp +1 -0
  521. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q5_0.cpp +1 -0
  522. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q5_1.cpp +1 -0
  523. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q8_0.cpp +1 -0
  524. data/ext/sources/ggml/src/ggml-sycl/type.hpp +112 -0
  525. data/ext/sources/ggml/src/ggml-sycl/upscale.cpp +410 -0
  526. data/ext/sources/ggml/src/ggml-sycl/upscale.hpp +9 -0
  527. data/ext/sources/ggml/src/ggml-sycl/vecdotq.hpp +242 -45
  528. data/ext/sources/ggml/src/ggml-virtgpu/ggml-backend-buffer.cpp +4 -0
  529. data/ext/sources/ggml/src/ggml-virtgpu/ggml-backend-device.cpp +2 -0
  530. data/ext/sources/ggml/src/ggml-virtgpu/ggml-backend.cpp +2 -0
  531. data/ext/sources/ggml/src/ggml-virtgpu/virtgpu-shm.cpp +1 -0
  532. data/ext/sources/ggml/src/ggml-virtgpu/virtgpu.cpp +1 -0
  533. data/ext/sources/ggml/src/ggml-virtgpu/virtgpu.h +0 -2
  534. data/ext/sources/ggml/src/ggml-vulkan/CMakeLists.txt +16 -0
  535. data/ext/sources/ggml/src/ggml-vulkan/ggml-vulkan.cpp +2843 -700
  536. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/CMakeLists.txt +4 -0
  537. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/col2im_1d.comp +61 -0
  538. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/contig_copy.comp +6 -2
  539. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/conv2d_mm.comp +146 -13
  540. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/conv3d_mm.comp +431 -0
  541. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/copy.comp +3 -1
  542. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/copy_from_quant.comp +1 -1
  543. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/copy_to_quant.comp +25 -1
  544. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs.glsl +88 -0
  545. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs_cm2.glsl +643 -1
  546. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dequant_nvfp4.comp +32 -0
  547. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q1_0.comp +29 -0
  548. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/diag.comp +3 -4
  549. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dot_product_funcs.glsl +27 -0
  550. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/feature-tests/coopmat2_decode_vector.comp +7 -0
  551. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp +198 -48
  552. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl +60 -59
  553. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp +116 -113
  554. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm2.comp +122 -31
  555. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_dequant.glsl +131 -0
  556. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_mmq_funcs.glsl +203 -0
  557. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/fwht.comp +115 -0
  558. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gated_delta_net.comp +125 -64
  559. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/generic_binary_head.glsl +0 -1
  560. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/generic_unary_head.glsl +21 -19
  561. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/get_rows_back.comp +25 -0
  562. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/glu_head.glsl +29 -1
  563. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/glu_main.glsl +17 -11
  564. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/im2col.comp +76 -54
  565. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/im2col_3d.comp +0 -1
  566. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/l2_norm.comp +4 -7
  567. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/log.comp +0 -1
  568. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec.comp +122 -27
  569. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_iface.glsl +6 -6
  570. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_q2_k.comp +1 -1
  571. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_q4_k.comp +1 -1
  572. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_q5_k.comp +1 -1
  573. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq.comp +22 -24
  574. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq_funcs.glsl +88 -55
  575. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp +42 -40
  576. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp +49 -15
  577. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl +222 -171
  578. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_funcs.glsl +8 -8
  579. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_shmem_types.glsl +24 -9
  580. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/multi_add.comp +0 -1
  581. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/norm.comp +10 -10
  582. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/repeat_back.comp +3 -3
  583. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/roll.comp +3 -3
  584. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/rope_funcs.glsl +5 -2
  585. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/rope_head.glsl +0 -1
  586. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/rope_params.glsl +3 -2
  587. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/snake.comp +49 -0
  588. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/ssm_conv.comp +11 -1
  589. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/tri.comp +3 -4
  590. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/types.glsl +79 -2
  591. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/unary.comp +168 -0
  592. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +282 -211
  593. data/ext/sources/ggml/src/ggml-webgpu/CMakeLists.txt +5 -2
  594. data/ext/sources/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp +2209 -283
  595. data/ext/sources/ggml/src/ggml-webgpu/ggml-webgpu.cpp +2618 -1416
  596. data/ext/sources/ggml/src/ggml-webgpu/pre_wgsl.hpp +37 -7
  597. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/add_id.wgsl +64 -0
  598. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/binary.wgsl +8 -7
  599. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/common_decls.tmpl +90 -95
  600. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/concat.wgsl +19 -1
  601. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/conv2d.wgsl +165 -0
  602. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/{cpy.tmpl.wgsl → cpy.wgsl} +25 -50
  603. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn.wgsl +107 -184
  604. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_quant_staging.tmpl +124 -0
  605. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_tile.wgsl +397 -0
  606. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_blk.wgsl +101 -0
  607. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_reduce.wgsl +84 -0
  608. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_split.wgsl +619 -0
  609. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/gated_delta_net.wgsl +149 -0
  610. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/get_rows.wgsl +204 -78
  611. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/glu.wgsl +155 -0
  612. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/im2col.wgsl +101 -0
  613. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_decls.tmpl +805 -526
  614. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_id.wgsl +195 -0
  615. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_id_gather.wgsl +52 -0
  616. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_id_vec.wgsl +154 -0
  617. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_reg_tile.wgsl +8 -6
  618. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_subgroup_matrix.wgsl +5 -1
  619. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec.wgsl +90 -413
  620. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_acc.tmpl +1553 -0
  621. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_q_acc.tmpl +297 -0
  622. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/quant_inner_loops.tmpl +21 -0
  623. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/quantize_q8.wgsl +178 -0
  624. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/rms_norm_mul.wgsl +152 -0
  625. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/{rope.tmpl.wgsl → rope.wgsl} +71 -142
  626. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/row_norm.wgsl +153 -0
  627. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/scale.wgsl +6 -4
  628. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/set.wgsl +109 -0
  629. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/set_rows.wgsl +2 -3
  630. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/set_rows_quant.wgsl +224 -0
  631. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/{soft_max.tmpl.wgsl → soft_max.wgsl} +106 -206
  632. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/solve_tri.wgsl +121 -0
  633. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/ssm_conv.wgsl +65 -0
  634. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/ssm_scan.wgsl +193 -0
  635. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/unary.wgsl +68 -48
  636. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/upscale.wgsl +240 -0
  637. data/ext/sources/ggml/src/ggml-zdnn/ggml-zdnn.cpp +18 -14
  638. data/ext/sources/ggml/src/ggml-zendnn/CMakeLists.txt +1 -1
  639. data/ext/sources/ggml/src/ggml-zendnn/ggml-zendnn.cpp +244 -10
  640. data/ext/sources/ggml/src/ggml.c +146 -42
  641. data/ext/sources/ggml/src/gguf.cpp +173 -28
  642. data/ext/sources/include/parakeet.h +342 -0
  643. data/ext/sources/include/whisper.h +31 -0
  644. data/ext/sources/media/matmul.png +0 -0
  645. data/ext/sources/src/CMakeLists.txt +23 -0
  646. data/ext/sources/src/parakeet-arch.h +188 -0
  647. data/ext/sources/src/parakeet.cpp +3838 -0
  648. data/ext/sources/src/whisper.cpp +220 -26
  649. data/extsources.rb +26 -10
  650. data/lib/whisper/log_settable.rb +33 -0
  651. data/lib/whisper/model/uri.rb +13 -8
  652. data/lib/whisper/output.rb +74 -0
  653. data/sig/whisper.rbs +417 -62
  654. data/test/helper.rb +2 -0
  655. data/test/jfk_reader/jfk_reader.c +50 -7
  656. data/test/test_callback.rb +1 -0
  657. data/test/test_package.rb +6 -5
  658. data/test/test_parakeet.rb +28 -0
  659. data/test/test_parakeet_callback.rb +107 -0
  660. data/test/test_parakeet_context.rb +116 -0
  661. data/test/test_parakeet_context_params.rb +24 -0
  662. data/test/test_parakeet_model.rb +21 -0
  663. data/test/test_parakeet_params.rb +78 -0
  664. data/test/test_parakeet_segment.rb +42 -0
  665. data/test/test_parakeet_token.rb +73 -0
  666. data/test/test_params.rb +2 -0
  667. data/test/test_vad.rb +9 -0
  668. data/test/test_vad_context.rb +2 -2
  669. data/test/test_vad_segment.rb +1 -1
  670. data/test/test_whisper.rb +24 -6
  671. data/whispercpp.gemspec +2 -2
  672. metadata +263 -304
  673. data/ext/sources/bindings/javascript/CMakeLists.txt +0 -41
  674. data/ext/sources/bindings/javascript/emscripten.cpp +0 -93
  675. data/ext/sources/bindings/javascript/libwhisper.worker.js +0 -1
  676. data/ext/sources/bindings/javascript/package.json +0 -26
  677. data/ext/sources/bindings/javascript/whisper.js +0 -19
  678. data/ext/sources/examples/addon.node/CMakeLists.txt +0 -31
  679. data/ext/sources/examples/addon.node/__test__/whisper.spec.js +0 -133
  680. data/ext/sources/examples/addon.node/addon.cpp +0 -557
  681. data/ext/sources/examples/addon.node/index.js +0 -59
  682. data/ext/sources/examples/addon.node/package.json +0 -16
  683. data/ext/sources/examples/addon.node/vad-example.js +0 -132
  684. data/ext/sources/examples/bench.wasm/CMakeLists.txt +0 -49
  685. data/ext/sources/examples/bench.wasm/emscripten.cpp +0 -87
  686. data/ext/sources/examples/bench.wasm/index-tmpl.html +0 -285
  687. data/ext/sources/examples/coi-serviceworker.js +0 -146
  688. data/ext/sources/examples/command/CMakeLists.txt +0 -10
  689. data/ext/sources/examples/command/command.cpp +0 -802
  690. data/ext/sources/examples/command/commands.txt +0 -9
  691. data/ext/sources/examples/command.wasm/CMakeLists.txt +0 -50
  692. data/ext/sources/examples/command.wasm/emscripten.cpp +0 -327
  693. data/ext/sources/examples/command.wasm/index-tmpl.html +0 -415
  694. data/ext/sources/examples/generate-karaoke.sh +0 -57
  695. data/ext/sources/examples/helpers.js +0 -191
  696. data/ext/sources/examples/livestream.sh +0 -112
  697. data/ext/sources/examples/lsp/CMakeLists.txt +0 -10
  698. data/ext/sources/examples/lsp/lsp.cpp +0 -471
  699. data/ext/sources/examples/lsp/whisper.vim +0 -362
  700. data/ext/sources/examples/python/test_whisper_processor.py +0 -7
  701. data/ext/sources/examples/python/whisper_processor.py +0 -54
  702. data/ext/sources/examples/server/bench.js +0 -29
  703. data/ext/sources/examples/server.py +0 -120
  704. data/ext/sources/examples/stream/CMakeLists.txt +0 -10
  705. data/ext/sources/examples/stream/stream.cpp +0 -437
  706. data/ext/sources/examples/stream.wasm/CMakeLists.txt +0 -49
  707. data/ext/sources/examples/stream.wasm/emscripten.cpp +0 -216
  708. data/ext/sources/examples/stream.wasm/index-tmpl.html +0 -491
  709. data/ext/sources/examples/sycl/CMakeLists.txt +0 -9
  710. data/ext/sources/examples/sycl/build.sh +0 -22
  711. data/ext/sources/examples/sycl/ls-sycl-device.cpp +0 -11
  712. data/ext/sources/examples/sycl/run-whisper.sh +0 -17
  713. data/ext/sources/examples/talk-llama/CMakeLists.txt +0 -48
  714. data/ext/sources/examples/talk-llama/eleven-labs.py +0 -80
  715. data/ext/sources/examples/talk-llama/llama-adapter.cpp +0 -488
  716. data/ext/sources/examples/talk-llama/llama-adapter.h +0 -89
  717. data/ext/sources/examples/talk-llama/llama-arch.cpp +0 -2877
  718. data/ext/sources/examples/talk-llama/llama-arch.h +0 -628
  719. data/ext/sources/examples/talk-llama/llama-batch.cpp +0 -919
  720. data/ext/sources/examples/talk-llama/llama-batch.h +0 -173
  721. data/ext/sources/examples/talk-llama/llama-chat.cpp +0 -896
  722. data/ext/sources/examples/talk-llama/llama-chat.h +0 -71
  723. data/ext/sources/examples/talk-llama/llama-context.cpp +0 -3633
  724. data/ext/sources/examples/talk-llama/llama-context.h +0 -359
  725. data/ext/sources/examples/talk-llama/llama-cparams.cpp +0 -5
  726. data/ext/sources/examples/talk-llama/llama-cparams.h +0 -47
  727. data/ext/sources/examples/talk-llama/llama-ext.h +0 -12
  728. data/ext/sources/examples/talk-llama/llama-grammar.cpp +0 -1464
  729. data/ext/sources/examples/talk-llama/llama-grammar.h +0 -194
  730. data/ext/sources/examples/talk-llama/llama-graph.cpp +0 -2735
  731. data/ext/sources/examples/talk-llama/llama-graph.h +0 -1031
  732. data/ext/sources/examples/talk-llama/llama-hparams.cpp +0 -258
  733. data/ext/sources/examples/talk-llama/llama-hparams.h +0 -353
  734. data/ext/sources/examples/talk-llama/llama-impl.cpp +0 -171
  735. data/ext/sources/examples/talk-llama/llama-impl.h +0 -75
  736. data/ext/sources/examples/talk-llama/llama-io.cpp +0 -15
  737. data/ext/sources/examples/talk-llama/llama-io.h +0 -35
  738. data/ext/sources/examples/talk-llama/llama-kv-cache-iswa.cpp +0 -330
  739. data/ext/sources/examples/talk-llama/llama-kv-cache-iswa.h +0 -137
  740. data/ext/sources/examples/talk-llama/llama-kv-cache.cpp +0 -2285
  741. data/ext/sources/examples/talk-llama/llama-kv-cache.h +0 -389
  742. data/ext/sources/examples/talk-llama/llama-kv-cells.h +0 -533
  743. data/ext/sources/examples/talk-llama/llama-memory-hybrid-iswa.cpp +0 -275
  744. data/ext/sources/examples/talk-llama/llama-memory-hybrid-iswa.h +0 -140
  745. data/ext/sources/examples/talk-llama/llama-memory-hybrid.cpp +0 -268
  746. data/ext/sources/examples/talk-llama/llama-memory-hybrid.h +0 -139
  747. data/ext/sources/examples/talk-llama/llama-memory-recurrent.cpp +0 -1165
  748. data/ext/sources/examples/talk-llama/llama-memory-recurrent.h +0 -182
  749. data/ext/sources/examples/talk-llama/llama-memory.cpp +0 -59
  750. data/ext/sources/examples/talk-llama/llama-memory.h +0 -122
  751. data/ext/sources/examples/talk-llama/llama-mmap.cpp +0 -752
  752. data/ext/sources/examples/talk-llama/llama-mmap.h +0 -73
  753. data/ext/sources/examples/talk-llama/llama-model-loader.cpp +0 -1655
  754. data/ext/sources/examples/talk-llama/llama-model-loader.h +0 -206
  755. data/ext/sources/examples/talk-llama/llama-model-saver.cpp +0 -299
  756. data/ext/sources/examples/talk-llama/llama-model-saver.h +0 -40
  757. data/ext/sources/examples/talk-llama/llama-model.cpp +0 -9056
  758. data/ext/sources/examples/talk-llama/llama-model.h +0 -597
  759. data/ext/sources/examples/talk-llama/llama-quant.cpp +0 -1304
  760. data/ext/sources/examples/talk-llama/llama-quant.h +0 -1
  761. data/ext/sources/examples/talk-llama/llama-sampler.cpp +0 -3885
  762. data/ext/sources/examples/talk-llama/llama-sampler.h +0 -42
  763. data/ext/sources/examples/talk-llama/llama-vocab.cpp +0 -3970
  764. data/ext/sources/examples/talk-llama/llama-vocab.h +0 -187
  765. data/ext/sources/examples/talk-llama/llama.cpp +0 -1194
  766. data/ext/sources/examples/talk-llama/llama.h +0 -1573
  767. data/ext/sources/examples/talk-llama/models/afmoe.cpp +0 -190
  768. data/ext/sources/examples/talk-llama/models/apertus.cpp +0 -125
  769. data/ext/sources/examples/talk-llama/models/arcee.cpp +0 -135
  770. data/ext/sources/examples/talk-llama/models/arctic.cpp +0 -137
  771. data/ext/sources/examples/talk-llama/models/arwkv7.cpp +0 -86
  772. data/ext/sources/examples/talk-llama/models/baichuan.cpp +0 -123
  773. data/ext/sources/examples/talk-llama/models/bailingmoe.cpp +0 -143
  774. data/ext/sources/examples/talk-llama/models/bailingmoe2.cpp +0 -133
  775. data/ext/sources/examples/talk-llama/models/bert.cpp +0 -184
  776. data/ext/sources/examples/talk-llama/models/bitnet.cpp +0 -145
  777. data/ext/sources/examples/talk-llama/models/bloom.cpp +0 -101
  778. data/ext/sources/examples/talk-llama/models/chameleon.cpp +0 -178
  779. data/ext/sources/examples/talk-llama/models/chatglm.cpp +0 -132
  780. data/ext/sources/examples/talk-llama/models/codeshell.cpp +0 -111
  781. data/ext/sources/examples/talk-llama/models/cogvlm.cpp +0 -102
  782. data/ext/sources/examples/talk-llama/models/cohere2-iswa.cpp +0 -134
  783. data/ext/sources/examples/talk-llama/models/command-r.cpp +0 -122
  784. data/ext/sources/examples/talk-llama/models/dbrx.cpp +0 -122
  785. data/ext/sources/examples/talk-llama/models/deci.cpp +0 -135
  786. data/ext/sources/examples/talk-llama/models/deepseek.cpp +0 -142
  787. data/ext/sources/examples/talk-llama/models/deepseek2.cpp +0 -262
  788. data/ext/sources/examples/talk-llama/models/delta-net-base.cpp +0 -445
  789. data/ext/sources/examples/talk-llama/models/dots1.cpp +0 -132
  790. data/ext/sources/examples/talk-llama/models/dream.cpp +0 -105
  791. data/ext/sources/examples/talk-llama/models/ernie4-5-moe.cpp +0 -148
  792. data/ext/sources/examples/talk-llama/models/ernie4-5.cpp +0 -110
  793. data/ext/sources/examples/talk-llama/models/eurobert.cpp +0 -97
  794. data/ext/sources/examples/talk-llama/models/exaone-moe.cpp +0 -145
  795. data/ext/sources/examples/talk-llama/models/exaone.cpp +0 -114
  796. data/ext/sources/examples/talk-llama/models/exaone4.cpp +0 -123
  797. data/ext/sources/examples/talk-llama/models/falcon-h1.cpp +0 -111
  798. data/ext/sources/examples/talk-llama/models/falcon.cpp +0 -120
  799. data/ext/sources/examples/talk-llama/models/gemma-embedding.cpp +0 -116
  800. data/ext/sources/examples/talk-llama/models/gemma.cpp +0 -112
  801. data/ext/sources/examples/talk-llama/models/gemma2-iswa.cpp +0 -128
  802. data/ext/sources/examples/talk-llama/models/gemma3.cpp +0 -155
  803. data/ext/sources/examples/talk-llama/models/gemma3n-iswa.cpp +0 -384
  804. data/ext/sources/examples/talk-llama/models/glm4-moe.cpp +0 -170
  805. data/ext/sources/examples/talk-llama/models/glm4.cpp +0 -157
  806. data/ext/sources/examples/talk-llama/models/gpt2.cpp +0 -105
  807. data/ext/sources/examples/talk-llama/models/gptneox.cpp +0 -144
  808. data/ext/sources/examples/talk-llama/models/granite-hybrid.cpp +0 -195
  809. data/ext/sources/examples/talk-llama/models/granite.cpp +0 -210
  810. data/ext/sources/examples/talk-llama/models/grok.cpp +0 -159
  811. data/ext/sources/examples/talk-llama/models/grovemoe.cpp +0 -139
  812. data/ext/sources/examples/talk-llama/models/hunyuan-dense.cpp +0 -132
  813. data/ext/sources/examples/talk-llama/models/hunyuan-moe.cpp +0 -153
  814. data/ext/sources/examples/talk-llama/models/internlm2.cpp +0 -120
  815. data/ext/sources/examples/talk-llama/models/jais.cpp +0 -86
  816. data/ext/sources/examples/talk-llama/models/jais2.cpp +0 -123
  817. data/ext/sources/examples/talk-llama/models/jamba.cpp +0 -106
  818. data/ext/sources/examples/talk-llama/models/kimi-linear.cpp +0 -381
  819. data/ext/sources/examples/talk-llama/models/lfm2.cpp +0 -196
  820. data/ext/sources/examples/talk-llama/models/llada-moe.cpp +0 -122
  821. data/ext/sources/examples/talk-llama/models/llada.cpp +0 -99
  822. data/ext/sources/examples/talk-llama/models/llama-iswa.cpp +0 -178
  823. data/ext/sources/examples/talk-llama/models/llama.cpp +0 -175
  824. data/ext/sources/examples/talk-llama/models/maincoder.cpp +0 -117
  825. data/ext/sources/examples/talk-llama/models/mamba-base.cpp +0 -289
  826. data/ext/sources/examples/talk-llama/models/mamba.cpp +0 -54
  827. data/ext/sources/examples/talk-llama/models/mimo2-iswa.cpp +0 -129
  828. data/ext/sources/examples/talk-llama/models/minicpm3.cpp +0 -200
  829. data/ext/sources/examples/talk-llama/models/minimax-m2.cpp +0 -123
  830. data/ext/sources/examples/talk-llama/models/mistral3.cpp +0 -160
  831. data/ext/sources/examples/talk-llama/models/models.h +0 -704
  832. data/ext/sources/examples/talk-llama/models/modern-bert.cpp +0 -109
  833. data/ext/sources/examples/talk-llama/models/mpt.cpp +0 -126
  834. data/ext/sources/examples/talk-llama/models/nemotron-h.cpp +0 -162
  835. data/ext/sources/examples/talk-llama/models/nemotron.cpp +0 -122
  836. data/ext/sources/examples/talk-llama/models/neo-bert.cpp +0 -104
  837. data/ext/sources/examples/talk-llama/models/olmo.cpp +0 -121
  838. data/ext/sources/examples/talk-llama/models/olmo2.cpp +0 -150
  839. data/ext/sources/examples/talk-llama/models/olmoe.cpp +0 -124
  840. data/ext/sources/examples/talk-llama/models/openai-moe-iswa.cpp +0 -127
  841. data/ext/sources/examples/talk-llama/models/openelm.cpp +0 -124
  842. data/ext/sources/examples/talk-llama/models/orion.cpp +0 -123
  843. data/ext/sources/examples/talk-llama/models/paddleocr.cpp +0 -122
  844. data/ext/sources/examples/talk-llama/models/pangu-embedded.cpp +0 -121
  845. data/ext/sources/examples/talk-llama/models/phi2.cpp +0 -121
  846. data/ext/sources/examples/talk-llama/models/phi3.cpp +0 -152
  847. data/ext/sources/examples/talk-llama/models/plamo.cpp +0 -110
  848. data/ext/sources/examples/talk-llama/models/plamo2.cpp +0 -320
  849. data/ext/sources/examples/talk-llama/models/plamo3.cpp +0 -128
  850. data/ext/sources/examples/talk-llama/models/plm.cpp +0 -169
  851. data/ext/sources/examples/talk-llama/models/qwen.cpp +0 -108
  852. data/ext/sources/examples/talk-llama/models/qwen2.cpp +0 -126
  853. data/ext/sources/examples/talk-llama/models/qwen2moe.cpp +0 -151
  854. data/ext/sources/examples/talk-llama/models/qwen2vl.cpp +0 -117
  855. data/ext/sources/examples/talk-llama/models/qwen3.cpp +0 -120
  856. data/ext/sources/examples/talk-llama/models/qwen35.cpp +0 -381
  857. data/ext/sources/examples/talk-llama/models/qwen35moe.cpp +0 -422
  858. data/ext/sources/examples/talk-llama/models/qwen3moe.cpp +0 -131
  859. data/ext/sources/examples/talk-llama/models/qwen3next.cpp +0 -525
  860. data/ext/sources/examples/talk-llama/models/qwen3vl-moe.cpp +0 -140
  861. data/ext/sources/examples/talk-llama/models/qwen3vl.cpp +0 -132
  862. data/ext/sources/examples/talk-llama/models/refact.cpp +0 -94
  863. data/ext/sources/examples/talk-llama/models/rnd1.cpp +0 -126
  864. data/ext/sources/examples/talk-llama/models/rwkv6-base.cpp +0 -164
  865. data/ext/sources/examples/talk-llama/models/rwkv6.cpp +0 -94
  866. data/ext/sources/examples/talk-llama/models/rwkv6qwen2.cpp +0 -86
  867. data/ext/sources/examples/talk-llama/models/rwkv7-base.cpp +0 -137
  868. data/ext/sources/examples/talk-llama/models/rwkv7.cpp +0 -90
  869. data/ext/sources/examples/talk-llama/models/seed-oss.cpp +0 -124
  870. data/ext/sources/examples/talk-llama/models/smallthinker.cpp +0 -126
  871. data/ext/sources/examples/talk-llama/models/smollm3.cpp +0 -128
  872. data/ext/sources/examples/talk-llama/models/stablelm.cpp +0 -146
  873. data/ext/sources/examples/talk-llama/models/starcoder.cpp +0 -100
  874. data/ext/sources/examples/talk-llama/models/starcoder2.cpp +0 -121
  875. data/ext/sources/examples/talk-llama/models/step35-iswa.cpp +0 -165
  876. data/ext/sources/examples/talk-llama/models/t5-dec.cpp +0 -166
  877. data/ext/sources/examples/talk-llama/models/t5-enc.cpp +0 -96
  878. data/ext/sources/examples/talk-llama/models/wavtokenizer-dec.cpp +0 -149
  879. data/ext/sources/examples/talk-llama/models/xverse.cpp +0 -108
  880. data/ext/sources/examples/talk-llama/prompts/talk-alpaca.txt +0 -23
  881. data/ext/sources/examples/talk-llama/speak +0 -40
  882. data/ext/sources/examples/talk-llama/speak.bat +0 -1
  883. data/ext/sources/examples/talk-llama/speak.ps1 +0 -14
  884. data/ext/sources/examples/talk-llama/talk-llama.cpp +0 -813
  885. data/ext/sources/examples/talk-llama/unicode-data.cpp +0 -7034
  886. data/ext/sources/examples/talk-llama/unicode-data.h +0 -20
  887. data/ext/sources/examples/talk-llama/unicode.cpp +0 -1103
  888. data/ext/sources/examples/talk-llama/unicode.h +0 -111
  889. data/ext/sources/examples/wchess/CMakeLists.txt +0 -10
  890. data/ext/sources/examples/wchess/libwchess/CMakeLists.txt +0 -19
  891. data/ext/sources/examples/wchess/libwchess/Chessboard.cpp +0 -803
  892. data/ext/sources/examples/wchess/libwchess/Chessboard.h +0 -33
  893. data/ext/sources/examples/wchess/libwchess/WChess.cpp +0 -193
  894. data/ext/sources/examples/wchess/libwchess/WChess.h +0 -63
  895. data/ext/sources/examples/wchess/libwchess/test-chessboard.cpp +0 -117
  896. data/ext/sources/examples/wchess/wchess.cmd/CMakeLists.txt +0 -8
  897. data/ext/sources/examples/wchess/wchess.cmd/wchess.cmd.cpp +0 -253
  898. data/ext/sources/examples/whisper.wasm/CMakeLists.txt +0 -50
  899. data/ext/sources/examples/whisper.wasm/emscripten.cpp +0 -118
  900. data/ext/sources/examples/whisper.wasm/index-tmpl.html +0 -659
  901. data/ext/sources/ggml/src/ggml-cuda/template-instances/generate_cu_files.py +0 -99
  902. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-msg.h +0 -155
  903. data/ext/sources/ggml/src/ggml-hexagon/op-desc.h +0 -153
  904. data/ext/sources/ggml/src/ggml-opencl/kernels/embed_kernel.py +0 -26
  905. data/ext/sources/ggml/src/ggml-openvino/openvino/pass/eliminate_zp.cpp +0 -123
  906. data/ext/sources/ggml/src/ggml-openvino/openvino/pass/eliminate_zp.h +0 -17
  907. data/ext/sources/ggml/src/ggml-virtgpu/regenerate_remoting.py +0 -333
  908. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/abs.comp +0 -21
  909. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/ceil.comp +0 -22
  910. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/clamp.comp +0 -17
  911. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/cos.comp +0 -17
  912. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/elu.comp +0 -27
  913. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/exp.comp +0 -21
  914. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/floor.comp +0 -22
  915. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu.comp +0 -25
  916. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu_erf.comp +0 -39
  917. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu_quick.comp +0 -23
  918. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/hardsigmoid.comp +0 -22
  919. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/hardswish.comp +0 -22
  920. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/leaky_relu.comp +0 -22
  921. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/neg.comp +0 -20
  922. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/relu.comp +0 -21
  923. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/round.comp +0 -29
  924. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/rte.glsl +0 -5
  925. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sgn.comp +0 -21
  926. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sigmoid.comp +0 -20
  927. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/silu.comp +0 -22
  928. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sin.comp +0 -17
  929. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/softplus.comp +0 -23
  930. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sqrt.comp +0 -17
  931. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/square.comp +0 -17
  932. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/step.comp +0 -22
  933. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/tanh.comp +0 -20
  934. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/trunc.comp +0 -22
  935. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/xielu.comp +0 -35
  936. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/embed_wgsl.py +0 -182
  937. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/glu.tmpl.wgsl +0 -323
  938. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat.wgsl +0 -718
  939. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/rms_norm.wgsl +0 -123
  940. data/ext/sources/tests/CMakeLists.txt +0 -112
  941. data/ext/sources/tests/earnings21/eval.mk +0 -58
  942. data/ext/sources/tests/earnings21/eval.py +0 -68
  943. data/ext/sources/tests/earnings21/normalizers/__init__.py +0 -2
  944. data/ext/sources/tests/earnings21/normalizers/basic.py +0 -80
  945. data/ext/sources/tests/earnings21/normalizers/english.json +0 -1741
  946. data/ext/sources/tests/earnings21/normalizers/english.py +0 -550
  947. data/ext/sources/tests/earnings21/requirements.txt +0 -6
  948. data/ext/sources/tests/en-0-ref.txt +0 -1
  949. data/ext/sources/tests/en-1-ref.txt +0 -1
  950. data/ext/sources/tests/en-2-ref.txt +0 -1
  951. data/ext/sources/tests/es-0-ref.txt +0 -1
  952. data/ext/sources/tests/librispeech/eval.mk +0 -39
  953. data/ext/sources/tests/librispeech/eval.py +0 -47
  954. data/ext/sources/tests/librispeech/normalizers/__init__.py +0 -2
  955. data/ext/sources/tests/librispeech/normalizers/basic.py +0 -80
  956. data/ext/sources/tests/librispeech/normalizers/english.json +0 -1741
  957. data/ext/sources/tests/librispeech/normalizers/english.py +0 -550
  958. data/ext/sources/tests/librispeech/requirements.txt +0 -6
  959. data/ext/sources/tests/run-tests.sh +0 -130
  960. data/ext/sources/tests/test-c.c +0 -3
  961. data/ext/sources/tests/test-vad-full.cpp +0 -56
  962. data/ext/sources/tests/test-vad.cpp +0 -83
  963. data/ext/sources/tests/test-whisper.js +0 -58
  964. data/lib/whisper/context.rb +0 -15
  965. data/lib/whisper/segment.rb +0 -58
@@ -56,6 +56,65 @@ static void mul_mat_vec_q_reorder(const void * __restrict__ vx, const void * __r
56
56
  }
57
57
  }
58
58
 
59
+ template <typename reorder_vec_dot_q_sycl, int ncols_dst>
60
+ static void mul_mat_vec_q_reorder_ncols(const void * __restrict__ vx, const void * __restrict__ vy,
61
+ float * __restrict__ dst, const int ncols, const int nrows,
62
+ const int stride_col_y_bytes, const int stride_col_dst,
63
+ const sycl::nd_item<3> & nd_item) {
64
+ using block_type = ggml_sycl_reordered::block_q_t<reorder_vec_dot_q_sycl::gtype>;
65
+ using block_traits = typename block_type::traits;
66
+
67
+ const auto sg = nd_item.get_sub_group();
68
+ const int sg_range = sg.get_group_linear_range();
69
+ const int workgroup_id = nd_item.get_group_linear_id();
70
+ const int sg_id = sg.get_group_linear_id();
71
+ const int row = workgroup_id * sg_range + sg_id;
72
+
73
+ if (row >= nrows) {
74
+ return;
75
+ }
76
+
77
+ const int blocks_per_row = ncols / block_traits::qk;
78
+ constexpr int blocks_per_subgroup = ceil_div(block_traits::vdr_mmvq * WARP_SIZE, block_traits::qi);
79
+ constexpr int block_elements_per_subgroup = block_traits::qi / block_traits::vdr_mmvq;
80
+ const int nblocks = nrows * (ncols / block_traits::qk);
81
+
82
+ static_assert(blocks_per_subgroup > 0);
83
+ static_assert(block_elements_per_subgroup > 0);
84
+
85
+ float partial_sum[ncols_dst] = {0.0f};
86
+ for (int i = sg.get_local_linear_id() / block_elements_per_subgroup; i < blocks_per_row; i += blocks_per_subgroup) {
87
+ const int ibx = row * blocks_per_row + i;
88
+
89
+ const auto bx_offset = block_type::get_block_offset(ibx, nblocks);
90
+ const auto d_offset = block_type::get_d_offset(nrows, ncols, ibx);
91
+ const int iby = i * block_type::block_to_q8_1_ratio();
92
+
93
+ #pragma unroll
94
+ for (int elem = 0; elem < block_elements_per_subgroup; elem += WARP_SIZE) {
95
+ const int iqs = elem + block_traits::vdr_mmvq * (sg.get_local_linear_id() % block_elements_per_subgroup);
96
+
97
+ #pragma unroll
98
+ for (int j = 0; j < ncols_dst; ++j) {
99
+ const char * vy_j = (const char *)vy + j * stride_col_y_bytes;
100
+ const int8_t * q8_1_quant_ptr = (const int8_t *)vy_j + iby * QK8_1;
101
+ const sycl::half2* q8_1_ds_ptr = (const sycl::half2 *)(vy_j + ncols + iby * sizeof(sycl::half2));
102
+
103
+ partial_sum[j] += reorder_vec_dot_q_sycl()(vx, bx_offset, d_offset, q8_1_quant_ptr, q8_1_ds_ptr, iqs);
104
+ }
105
+ }
106
+ }
107
+
108
+ #pragma unroll
109
+ for (int j = 0; j < ncols_dst; ++j) {
110
+ float sum = sycl::reduce_over_group(nd_item.get_sub_group(), partial_sum[j], std::plus<>());
111
+
112
+ if (sg.leader()) {
113
+ dst[j * stride_col_dst + row] = sum;
114
+ }
115
+ }
116
+ }
117
+
59
118
  template <int qk, int qi, typename block_q_t, int vdr, vec_dot_q_sycl_t vec_dot_q_sycl>
60
119
  static void mul_mat_vec_q(const void * __restrict__ vx, const void * __restrict__ vy, float * __restrict__ dst,
61
120
  const int ncols, const int nrows, const sycl::nd_item<3> & item_ct1) {
@@ -100,6 +159,70 @@ static void mul_mat_vec_q(const void * __restrict__ vx, const void * __restrict_
100
159
  }
101
160
  }
102
161
 
162
+ template <int qk, int qi, typename block_q_t, int vdr,
163
+ vec_dot_q_sycl_t vec_dot_q_sycl, int ncols_dst>
164
+ static void mul_mat_vec_q_ncols(
165
+ const void * __restrict__ vx,
166
+ const void * __restrict__ vy,
167
+ float * __restrict__ dst,
168
+ const int ncols,
169
+ const int nrows,
170
+ const int stride_col_y,
171
+ const int stride_col_dst,
172
+ const sycl::nd_item<3> & item_ct1) {
173
+
174
+ const int row = item_ct1.get_group(2) * item_ct1.get_local_range(1)
175
+ + item_ct1.get_local_id(1);
176
+
177
+ if (row >= nrows) {
178
+ return;
179
+ }
180
+
181
+ const int blocks_per_row = ncols / qk;
182
+ constexpr int blocks_per_warp = (vdr * WARP_SIZE + qi - 1) / qi;
183
+
184
+ // partial sums: one per output column
185
+ float tmp[ncols_dst] = {0.0f};
186
+
187
+ const block_q_t * x = (const block_q_t *) vx;
188
+ const block_q8_1 * y = (const block_q8_1 *) vy;
189
+
190
+ for (int i = item_ct1.get_local_id(2) / (qi / vdr);
191
+ i < blocks_per_row;
192
+ i += blocks_per_warp) {
193
+
194
+ const int ibx = row * blocks_per_row + i;
195
+ const int iby = i * (qk / QK8_1);
196
+
197
+ // read weight block once, dot against all columns
198
+ for (size_t elem = 0; elem < qi / vdr; elem += WARP_SIZE) {
199
+ const int iqs = elem + vdr * (item_ct1.get_local_id(2) % (qi / vdr));
200
+
201
+ #pragma unroll
202
+ for (int j = 0; j < ncols_dst; ++j) {
203
+ tmp[j] += vec_dot_q_sycl(&x[ibx], &y[j * stride_col_y + iby], iqs);
204
+ }
205
+ }
206
+ }
207
+
208
+ // reduce within subgroup
209
+ #pragma unroll
210
+ for (int j = 0; j < ncols_dst; ++j) {
211
+ #pragma unroll
212
+ for (int mask = WARP_SIZE / 2; mask > 0; mask >>= 1) {
213
+ tmp[j] += dpct::permute_sub_group_by_xor(
214
+ item_ct1.get_sub_group(), tmp[j], mask);
215
+ }
216
+ }
217
+
218
+ if (item_ct1.get_local_id(2) == 0) {
219
+ #pragma unroll
220
+ for (int j = 0; j < ncols_dst; ++j) {
221
+ dst[j * stride_col_dst + row] = tmp[j];
222
+ }
223
+ }
224
+ }
225
+
103
226
  template <int qk, int qi, typename block_q_t, int vdr>
104
227
  static void mul_mat_vec_q_iq2_xxs_q8_1(const void *__restrict__ vx,
105
228
  const void *__restrict__ vy,
@@ -537,15 +660,14 @@ static void mul_mat_vec_q_iq4_xs_q8_1(const void *__restrict__ vx,
537
660
  static void reorder_mul_mat_vec_q4_0_q8_1_sycl(const void * vx, const void * vy, float * dst, const int ncols,
538
661
  const int nrows, dpct::queue_ptr stream) {
539
662
  GGML_ASSERT(ncols % QK4_0 == 0);
540
- const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y);
541
- constexpr size_t num_subgroups = 16;
542
- GGML_ASSERT(block_num_y % num_subgroups == 0);
543
-
544
- const sycl::range<3> global_size(1, GGML_SYCL_MMV_Y, (block_num_y * WARP_SIZE));
545
- const sycl::range<3> workgroup_size(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
663
+ // Round up to a whole number of subgroup-sized workgroups; out-of-range rows are skipped inside the kernel.
664
+ constexpr size_t num_subgroups = WARP_SIZE;
665
+ const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
666
+ const sycl::range<3> block_nums(1, 1, block_num_y);
667
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
546
668
 
547
669
  stream->submit([&](sycl::handler & cgh) {
548
- cgh.parallel_for(sycl::nd_range<3>(global_size, workgroup_size),
670
+ cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
549
671
  [=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
550
672
  mul_mat_vec_q_reorder<reorder_vec_dot_q_sycl<GGML_TYPE_Q4_0>>(vx, vy, dst, ncols, nrows,
551
673
  nd_item);
@@ -553,6 +675,45 @@ static void reorder_mul_mat_vec_q4_0_q8_1_sycl(const void * vx, const void * vy,
553
675
  });
554
676
  }
555
677
 
678
+ template <int ncols_dst>
679
+ static void reorder_mul_mat_vec_q4_0_q8_1_sycl_ncols(
680
+ const void * vx, const void * vy, float * dst,
681
+ const int ncols, const int nrows,
682
+ const int stride_col_y_bytes, const int stride_col_dst,
683
+ dpct::queue_ptr stream) {
684
+ GGML_ASSERT(ncols % QK4_0 == 0);
685
+ constexpr size_t num_subgroups = WARP_SIZE;
686
+ const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
687
+ const sycl::range<3> block_nums(1, 1, block_num_y);
688
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
689
+
690
+ stream->submit([&](sycl::handler & cgh) {
691
+ cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
692
+ [=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
693
+ mul_mat_vec_q_reorder_ncols<reorder_vec_dot_q_sycl<GGML_TYPE_Q4_0>, ncols_dst>(
694
+ vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, nd_item);
695
+ });
696
+ });
697
+ }
698
+
699
+ static void reorder_mul_mat_vec_q4_0_q8_1_sycl_switch_ncols(
700
+ const void * vx, const void * vy, float * dst,
701
+ const int ncols, const int nrows, const int ncols_dst,
702
+ const int stride_col_y_bytes, const int stride_col_dst,
703
+ dpct::queue_ptr stream) {
704
+ switch (ncols_dst) {
705
+ case 1: reorder_mul_mat_vec_q4_0_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
706
+ case 2: reorder_mul_mat_vec_q4_0_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
707
+ case 3: reorder_mul_mat_vec_q4_0_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
708
+ case 4: reorder_mul_mat_vec_q4_0_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
709
+ case 5: reorder_mul_mat_vec_q4_0_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
710
+ case 6: reorder_mul_mat_vec_q4_0_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
711
+ case 7: reorder_mul_mat_vec_q4_0_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
712
+ case 8: reorder_mul_mat_vec_q4_0_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
713
+ default: GGML_ABORT("unsupported ncols_dst=%d for Q4_0 reorder multi-col MMVQ", ncols_dst);
714
+ }
715
+ }
716
+
556
717
  static void mul_mat_vec_q4_0_q8_1_sycl(const void * vx, const void * vy, float * dst, const int ncols, const int nrows,
557
718
  dpct::queue_ptr stream) {
558
719
  GGML_ASSERT(ncols % QK4_0 == 0);
@@ -571,6 +732,45 @@ static void mul_mat_vec_q4_0_q8_1_sycl(const void * vx, const void * vy, float *
571
732
  }
572
733
  }
573
734
 
735
+ template <int ncols_dst>
736
+ static void mul_mat_vec_q4_0_q8_1_sycl_ncols(
737
+ const void * vx, const void * vy, float * dst,
738
+ const int ncols, const int nrows,
739
+ const int stride_col_y, const int stride_col_dst,
740
+ dpct::queue_ptr stream) {
741
+ GGML_ASSERT(ncols % QK4_0 == 0);
742
+ const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
743
+ const sycl::range<3> block_nums(1, 1, block_num_y);
744
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
745
+ stream->submit([&](sycl::handler & cgh) {
746
+ cgh.parallel_for(
747
+ sycl::nd_range<3>(block_nums * block_dims, block_dims),
748
+ [=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
749
+ mul_mat_vec_q_ncols<QK4_0, QI4_0, block_q4_0,
750
+ VDR_Q4_0_Q8_1_MMVQ, vec_dot_q4_0_q8_1, ncols_dst>(
751
+ vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, item_ct1);
752
+ });
753
+ });
754
+ }
755
+
756
+ static void mul_mat_vec_q4_0_q8_1_sycl_switch_ncols(
757
+ const void * vx, const void * vy, float * dst,
758
+ const int ncols, const int nrows, const int ncols_dst,
759
+ const int stride_col_y, const int stride_col_dst,
760
+ dpct::queue_ptr stream) {
761
+ switch (ncols_dst) {
762
+ case 1: mul_mat_vec_q4_0_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
763
+ case 2: mul_mat_vec_q4_0_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
764
+ case 3: mul_mat_vec_q4_0_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
765
+ case 4: mul_mat_vec_q4_0_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
766
+ case 5: mul_mat_vec_q4_0_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
767
+ case 6: mul_mat_vec_q4_0_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
768
+ case 7: mul_mat_vec_q4_0_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
769
+ case 8: mul_mat_vec_q4_0_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
770
+ default: GGML_ABORT("unsupported ncols_dst=%d for Q4_0 multi-col MMVQ", ncols_dst);
771
+ }
772
+ }
773
+
574
774
  static void mul_mat_vec_q4_1_q8_1_sycl(const void *vx, const void *vy,
575
775
  float *dst, const int ncols,
576
776
  const int nrows,
@@ -595,6 +795,45 @@ static void mul_mat_vec_q4_1_q8_1_sycl(const void *vx, const void *vy,
595
795
  }
596
796
  }
597
797
 
798
+ template <int ncols_dst>
799
+ static void mul_mat_vec_q4_1_q8_1_sycl_ncols(
800
+ const void * vx, const void * vy, float * dst,
801
+ const int ncols, const int nrows,
802
+ const int stride_col_y, const int stride_col_dst,
803
+ dpct::queue_ptr stream) {
804
+ GGML_ASSERT(ncols % QK4_1 == 0);
805
+ const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
806
+ const sycl::range<3> block_nums(1, 1, block_num_y);
807
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
808
+ stream->submit([&](sycl::handler & cgh) {
809
+ cgh.parallel_for(
810
+ sycl::nd_range<3>(block_nums * block_dims, block_dims),
811
+ [=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
812
+ mul_mat_vec_q_ncols<QK4_0, QI4_1, block_q4_1,
813
+ VDR_Q4_1_Q8_1_MMVQ, vec_dot_q4_1_q8_1, ncols_dst>(
814
+ vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, item_ct1);
815
+ });
816
+ });
817
+ }
818
+
819
+ static void mul_mat_vec_q4_1_q8_1_sycl_switch_ncols(
820
+ const void * vx, const void * vy, float * dst,
821
+ const int ncols, const int nrows, const int ncols_dst,
822
+ const int stride_col_y, const int stride_col_dst,
823
+ dpct::queue_ptr stream) {
824
+ switch (ncols_dst) {
825
+ case 1: mul_mat_vec_q4_1_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
826
+ case 2: mul_mat_vec_q4_1_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
827
+ case 3: mul_mat_vec_q4_1_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
828
+ case 4: mul_mat_vec_q4_1_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
829
+ case 5: mul_mat_vec_q4_1_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
830
+ case 6: mul_mat_vec_q4_1_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
831
+ case 7: mul_mat_vec_q4_1_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
832
+ case 8: mul_mat_vec_q4_1_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
833
+ default: GGML_ABORT("unsupported ncols_dst=%d for Q4_1 multi-col MMVQ", ncols_dst);
834
+ }
835
+ }
836
+
598
837
  static void mul_mat_vec_mxfp4_q8_1_sycl(const void * vx, const void * vy, float * dst, const int ncols, const int nrows,
599
838
  dpct::queue_ptr stream) {
600
839
  GGML_ASSERT(ncols % QK_MXFP4 == 0);
@@ -613,6 +852,101 @@ static void mul_mat_vec_mxfp4_q8_1_sycl(const void * vx, const void * vy, float
613
852
  }
614
853
  }
615
854
 
855
+ template <int ncols_dst>
856
+ static void mul_mat_vec_mxfp4_q8_1_sycl_ncols(
857
+ const void * vx, const void * vy, float * dst,
858
+ const int ncols, const int nrows,
859
+ const int stride_col_y, const int stride_col_dst,
860
+ dpct::queue_ptr stream) {
861
+ GGML_ASSERT(ncols % QK_MXFP4 == 0);
862
+ const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
863
+ const sycl::range<3> block_nums(1, 1, block_num_y);
864
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
865
+ stream->submit([&](sycl::handler & cgh) {
866
+ cgh.parallel_for(
867
+ sycl::nd_range<3>(block_nums * block_dims, block_dims),
868
+ [=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
869
+ mul_mat_vec_q_ncols<QK_MXFP4, QI_MXFP4, block_mxfp4,
870
+ VDR_MXFP4_Q8_1_MMVQ, vec_dot_mxfp4_q8_1, ncols_dst>(
871
+ vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, item_ct1);
872
+ });
873
+ });
874
+ }
875
+
876
+ static void mul_mat_vec_mxfp4_q8_1_sycl_switch_ncols(
877
+ const void * vx, const void * vy, float * dst,
878
+ const int ncols, const int nrows, const int ncols_dst,
879
+ const int stride_col_y, const int stride_col_dst,
880
+ dpct::queue_ptr stream) {
881
+ switch (ncols_dst) {
882
+ case 1: mul_mat_vec_mxfp4_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
883
+ case 2: mul_mat_vec_mxfp4_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
884
+ case 3: mul_mat_vec_mxfp4_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
885
+ case 4: mul_mat_vec_mxfp4_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
886
+ case 5: mul_mat_vec_mxfp4_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
887
+ case 6: mul_mat_vec_mxfp4_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
888
+ case 7: mul_mat_vec_mxfp4_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
889
+ case 8: mul_mat_vec_mxfp4_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
890
+ default: GGML_ABORT("unsupported ncols_dst=%d for MXFP4 multi-col MMVQ", ncols_dst);
891
+ }
892
+ }
893
+
894
+ static void mul_mat_vec_nvfp4_q8_1_sycl(const void * vx, const void * vy, float * dst, const int ncols, const int nrows,
895
+ dpct::queue_ptr stream) {
896
+ GGML_ASSERT(ncols % QK_NVFP4 == 0);
897
+ const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
898
+ const sycl::range<3> block_nums(1, 1, block_num_y);
899
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
900
+
901
+ {
902
+ stream->submit([&](sycl::handler & cgh) {
903
+ cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
904
+ [=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
905
+ mul_mat_vec_q<QK_NVFP4, QI_NVFP4, block_nvfp4, VDR_NVFP4_Q8_1_MMVQ, vec_dot_nvfp4_q8_1>(
906
+ vx, vy, dst, ncols, nrows, item_ct1);
907
+ });
908
+ });
909
+ }
910
+ }
911
+
912
+ template <int ncols_dst>
913
+ static void mul_mat_vec_nvfp4_q8_1_sycl_ncols(
914
+ const void * vx, const void * vy, float * dst,
915
+ const int ncols, const int nrows,
916
+ const int stride_col_y, const int stride_col_dst,
917
+ dpct::queue_ptr stream) {
918
+ GGML_ASSERT(ncols % QK_NVFP4 == 0);
919
+ const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
920
+ const sycl::range<3> block_nums(1, 1, block_num_y);
921
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
922
+ stream->submit([&](sycl::handler & cgh) {
923
+ cgh.parallel_for(
924
+ sycl::nd_range<3>(block_nums * block_dims, block_dims),
925
+ [=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
926
+ mul_mat_vec_q_ncols<QK_NVFP4, QI_NVFP4, block_nvfp4,
927
+ VDR_NVFP4_Q8_1_MMVQ, vec_dot_nvfp4_q8_1, ncols_dst>(
928
+ vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, item_ct1);
929
+ });
930
+ });
931
+ }
932
+
933
+ static void mul_mat_vec_nvfp4_q8_1_sycl_switch_ncols(
934
+ const void * vx, const void * vy, float * dst,
935
+ const int ncols, const int nrows, const int ncols_dst,
936
+ const int stride_col_y, const int stride_col_dst,
937
+ dpct::queue_ptr stream) {
938
+ switch (ncols_dst) {
939
+ case 1: mul_mat_vec_nvfp4_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
940
+ case 2: mul_mat_vec_nvfp4_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
941
+ case 3: mul_mat_vec_nvfp4_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
942
+ case 4: mul_mat_vec_nvfp4_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
943
+ case 5: mul_mat_vec_nvfp4_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
944
+ case 6: mul_mat_vec_nvfp4_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
945
+ case 7: mul_mat_vec_nvfp4_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
946
+ case 8: mul_mat_vec_nvfp4_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
947
+ default: GGML_ABORT("unsupported ncols_dst=%d for NVFP4 multi-col MMVQ", ncols_dst);
948
+ }
949
+ }
616
950
 
617
951
  static void mul_mat_vec_q5_0_q8_1_sycl(const void *vx, const void *vy,
618
952
  float *dst, const int ncols,
@@ -638,6 +972,45 @@ static void mul_mat_vec_q5_0_q8_1_sycl(const void *vx, const void *vy,
638
972
  }
639
973
  }
640
974
 
975
+ template <int ncols_dst>
976
+ static void mul_mat_vec_q5_0_q8_1_sycl_ncols(
977
+ const void * vx, const void * vy, float * dst,
978
+ const int ncols, const int nrows,
979
+ const int stride_col_y, const int stride_col_dst,
980
+ dpct::queue_ptr stream) {
981
+ GGML_ASSERT(ncols % QK5_0 == 0);
982
+ const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
983
+ const sycl::range<3> block_nums(1, 1, block_num_y);
984
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
985
+ stream->submit([&](sycl::handler & cgh) {
986
+ cgh.parallel_for(
987
+ sycl::nd_range<3>(block_nums * block_dims, block_dims),
988
+ [=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
989
+ mul_mat_vec_q_ncols<QK5_0, QI5_0, block_q5_0,
990
+ VDR_Q5_0_Q8_1_MMVQ, vec_dot_q5_0_q8_1, ncols_dst>(
991
+ vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, item_ct1);
992
+ });
993
+ });
994
+ }
995
+
996
+ static void mul_mat_vec_q5_0_q8_1_sycl_switch_ncols(
997
+ const void * vx, const void * vy, float * dst,
998
+ const int ncols, const int nrows, const int ncols_dst,
999
+ const int stride_col_y, const int stride_col_dst,
1000
+ dpct::queue_ptr stream) {
1001
+ switch (ncols_dst) {
1002
+ case 1: mul_mat_vec_q5_0_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
1003
+ case 2: mul_mat_vec_q5_0_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1004
+ case 3: mul_mat_vec_q5_0_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1005
+ case 4: mul_mat_vec_q5_0_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1006
+ case 5: mul_mat_vec_q5_0_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1007
+ case 6: mul_mat_vec_q5_0_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1008
+ case 7: mul_mat_vec_q5_0_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1009
+ case 8: mul_mat_vec_q5_0_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1010
+ default: GGML_ABORT("unsupported ncols_dst=%d for Q5_0 multi-col MMVQ", ncols_dst);
1011
+ }
1012
+ }
1013
+
641
1014
  static void mul_mat_vec_q5_1_q8_1_sycl(const void *vx, const void *vy,
642
1015
  float *dst, const int ncols,
643
1016
  const int nrows,
@@ -662,6 +1035,102 @@ static void mul_mat_vec_q5_1_q8_1_sycl(const void *vx, const void *vy,
662
1035
  }
663
1036
  }
664
1037
 
1038
+ template <int ncols_dst>
1039
+ static void mul_mat_vec_q5_1_q8_1_sycl_ncols(
1040
+ const void * vx, const void * vy, float * dst,
1041
+ const int ncols, const int nrows,
1042
+ const int stride_col_y, const int stride_col_dst,
1043
+ dpct::queue_ptr stream) {
1044
+ GGML_ASSERT(ncols % QK5_1 == 0);
1045
+ const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
1046
+ const sycl::range<3> block_nums(1, 1, block_num_y);
1047
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
1048
+ stream->submit([&](sycl::handler & cgh) {
1049
+ cgh.parallel_for(
1050
+ sycl::nd_range<3>(block_nums * block_dims, block_dims),
1051
+ [=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
1052
+ mul_mat_vec_q_ncols<QK5_1, QI5_1, block_q5_1,
1053
+ VDR_Q5_1_Q8_1_MMVQ, vec_dot_q5_1_q8_1, ncols_dst>(
1054
+ vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, item_ct1);
1055
+ });
1056
+ });
1057
+ }
1058
+
1059
+ static void mul_mat_vec_q5_1_q8_1_sycl_switch_ncols(
1060
+ const void * vx, const void * vy, float * dst,
1061
+ const int ncols, const int nrows, const int ncols_dst,
1062
+ const int stride_col_y, const int stride_col_dst,
1063
+ dpct::queue_ptr stream) {
1064
+ switch (ncols_dst) {
1065
+ case 1: mul_mat_vec_q5_1_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
1066
+ case 2: mul_mat_vec_q5_1_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1067
+ case 3: mul_mat_vec_q5_1_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1068
+ case 4: mul_mat_vec_q5_1_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1069
+ case 5: mul_mat_vec_q5_1_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1070
+ case 6: mul_mat_vec_q5_1_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1071
+ case 7: mul_mat_vec_q5_1_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1072
+ case 8: mul_mat_vec_q5_1_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1073
+ default: GGML_ABORT("unsupported ncols_dst=%d for Q5_1 multi-col MMVQ", ncols_dst);
1074
+ }
1075
+ }
1076
+
1077
+ static void reorder_mul_mat_vec_q8_0_q8_1_sycl(const void * vx, const void * vy, float * dst, const int ncols,
1078
+ const int nrows, dpct::queue_ptr stream) {
1079
+ GGML_ASSERT(ncols % QK8_0 == 0);
1080
+ // Round up to a whole number of subgroup-sized workgroups; out-of-range rows are skipped inside the kernel.
1081
+ constexpr size_t num_subgroups = WARP_SIZE;
1082
+ const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
1083
+ const sycl::range<3> block_nums(1, 1, block_num_y);
1084
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
1085
+
1086
+ stream->submit([&](sycl::handler & cgh) {
1087
+ cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
1088
+ [=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
1089
+ mul_mat_vec_q_reorder<reorder_vec_dot_q_sycl<GGML_TYPE_Q8_0>>(vx, vy, dst, ncols, nrows,
1090
+ nd_item);
1091
+ });
1092
+ });
1093
+ }
1094
+
1095
+ template <int ncols_dst>
1096
+ static void reorder_mul_mat_vec_q8_0_q8_1_sycl_ncols(
1097
+ const void * vx, const void * vy, float * dst,
1098
+ const int ncols, const int nrows,
1099
+ const int stride_col_y_bytes, const int stride_col_dst,
1100
+ dpct::queue_ptr stream) {
1101
+ GGML_ASSERT(ncols % QK8_0 == 0);
1102
+ constexpr size_t num_subgroups = WARP_SIZE;
1103
+ const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
1104
+ const sycl::range<3> block_nums(1, 1, block_num_y);
1105
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
1106
+
1107
+ stream->submit([&](sycl::handler & cgh) {
1108
+ cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
1109
+ [=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
1110
+ mul_mat_vec_q_reorder_ncols<reorder_vec_dot_q_sycl<GGML_TYPE_Q8_0>, ncols_dst>(
1111
+ vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, nd_item);
1112
+ });
1113
+ });
1114
+ }
1115
+
1116
+ static void reorder_mul_mat_vec_q8_0_q8_1_sycl_switch_ncols(
1117
+ const void * vx, const void * vy, float * dst,
1118
+ const int ncols, const int nrows, const int ncols_dst,
1119
+ const int stride_col_y_bytes, const int stride_col_dst,
1120
+ dpct::queue_ptr stream) {
1121
+ switch (ncols_dst) {
1122
+ case 1: reorder_mul_mat_vec_q8_0_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
1123
+ case 2: reorder_mul_mat_vec_q8_0_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1124
+ case 3: reorder_mul_mat_vec_q8_0_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1125
+ case 4: reorder_mul_mat_vec_q8_0_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1126
+ case 5: reorder_mul_mat_vec_q8_0_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1127
+ case 6: reorder_mul_mat_vec_q8_0_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1128
+ case 7: reorder_mul_mat_vec_q8_0_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1129
+ case 8: reorder_mul_mat_vec_q8_0_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1130
+ default: GGML_ABORT("unsupported ncols_dst=%d for Q8_0 reorder multi-col MMVQ", ncols_dst);
1131
+ }
1132
+ }
1133
+
665
1134
  static void mul_mat_vec_q8_0_q8_1_sycl(const void *vx, const void *vy,
666
1135
  float *dst, const int ncols,
667
1136
  const int nrows,
@@ -686,6 +1155,105 @@ static void mul_mat_vec_q8_0_q8_1_sycl(const void *vx, const void *vy,
686
1155
  }
687
1156
  }
688
1157
 
1158
+ template <int ncols_dst>
1159
+ static void mul_mat_vec_q8_0_q8_1_sycl_ncols(
1160
+ const void * vx, const void * vy, float * dst,
1161
+ const int ncols, const int nrows,
1162
+ const int stride_col_y, const int stride_col_dst,
1163
+ dpct::queue_ptr stream) {
1164
+ GGML_ASSERT(ncols % QK8_0 == 0);
1165
+ const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
1166
+ const sycl::range<3> block_nums(1, 1, block_num_y);
1167
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
1168
+ stream->submit([&](sycl::handler & cgh) {
1169
+ cgh.parallel_for(
1170
+ sycl::nd_range<3>(block_nums * block_dims, block_dims),
1171
+ [=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
1172
+ mul_mat_vec_q_ncols<QK8_0, QI8_0, block_q8_0,
1173
+ VDR_Q8_0_Q8_1_MMVQ, vec_dot_q8_0_q8_1, ncols_dst>(
1174
+ vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, item_ct1);
1175
+ });
1176
+ });
1177
+ }
1178
+
1179
+ static void mul_mat_vec_q8_0_q8_1_sycl_switch_ncols(
1180
+ const void * vx, const void * vy, float * dst,
1181
+ const int ncols, const int nrows, const int ncols_dst,
1182
+ const int stride_col_y, const int stride_col_dst,
1183
+ dpct::queue_ptr stream) {
1184
+ switch (ncols_dst) {
1185
+ case 1: mul_mat_vec_q8_0_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
1186
+ case 2: mul_mat_vec_q8_0_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1187
+ case 3: mul_mat_vec_q8_0_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1188
+ case 4: mul_mat_vec_q8_0_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1189
+ case 5: mul_mat_vec_q8_0_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1190
+ case 6: mul_mat_vec_q8_0_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1191
+ case 7: mul_mat_vec_q8_0_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1192
+ case 8: mul_mat_vec_q8_0_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1193
+ default: GGML_ABORT("unsupported ncols_dst=%d for Q8_0 multi-col MMVQ", ncols_dst);
1194
+ }
1195
+ }
1196
+
1197
+ static void mul_mat_vec_q1_0_q8_1_sycl(const void * vx, const void * vy,
1198
+ float * dst, const int ncols,
1199
+ const int nrows,
1200
+ dpct::queue_ptr stream) {
1201
+ GGML_ASSERT(ncols % QK1_0 == 0);
1202
+ const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
1203
+ const sycl::range<3> block_nums(1, 1, block_num_y);
1204
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
1205
+
1206
+ stream->submit([&](sycl::handler & cgh) {
1207
+ cgh.parallel_for(
1208
+ sycl::nd_range<3>(block_nums * block_dims, block_dims),
1209
+ [=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
1210
+ mul_mat_vec_q<QK1_0, QI1_0, block_q1_0,
1211
+ VDR_Q1_0_Q8_1_MMVQ, vec_dot_q1_0_q8_1>(
1212
+ vx, vy, dst, ncols, nrows, item_ct1);
1213
+ });
1214
+ });
1215
+ }
1216
+
1217
+ template <int ncols_dst>
1218
+ static void mul_mat_vec_q1_0_q8_1_sycl_ncols(
1219
+ const void * vx, const void * vy, float * dst,
1220
+ const int ncols, const int nrows,
1221
+ const int stride_col_y, const int stride_col_dst,
1222
+ dpct::queue_ptr stream) {
1223
+ GGML_ASSERT(ncols % QK1_0 == 0);
1224
+ const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
1225
+ const sycl::range<3> block_nums(1, 1, block_num_y);
1226
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
1227
+
1228
+ stream->submit([&](sycl::handler & cgh) {
1229
+ cgh.parallel_for(
1230
+ sycl::nd_range<3>(block_nums * block_dims, block_dims),
1231
+ [=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
1232
+ mul_mat_vec_q_ncols<QK1_0, QI1_0, block_q1_0,
1233
+ VDR_Q1_0_Q8_1_MMVQ, vec_dot_q1_0_q8_1, ncols_dst>(
1234
+ vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, item_ct1);
1235
+ });
1236
+ });
1237
+ }
1238
+
1239
+ static void mul_mat_vec_q1_0_q8_1_sycl_switch_ncols(
1240
+ const void * vx, const void * vy, float * dst,
1241
+ const int ncols, const int nrows, const int ncols_dst,
1242
+ const int stride_col_y, const int stride_col_dst,
1243
+ dpct::queue_ptr stream) {
1244
+ switch (ncols_dst) {
1245
+ case 1: mul_mat_vec_q1_0_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
1246
+ case 2: mul_mat_vec_q1_0_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1247
+ case 3: mul_mat_vec_q1_0_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1248
+ case 4: mul_mat_vec_q1_0_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1249
+ case 5: mul_mat_vec_q1_0_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1250
+ case 6: mul_mat_vec_q1_0_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1251
+ case 7: mul_mat_vec_q1_0_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1252
+ case 8: mul_mat_vec_q1_0_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1253
+ default: GGML_ABORT("unsupported ncols_dst=%d for Q1_0 multi-col MMVQ", ncols_dst);
1254
+ }
1255
+ }
1256
+
689
1257
  static void mul_mat_vec_q2_K_q8_1_sycl(const void *vx, const void *vy,
690
1258
  float *dst, const int ncols,
691
1259
  const int nrows,
@@ -710,6 +1278,45 @@ static void mul_mat_vec_q2_K_q8_1_sycl(const void *vx, const void *vy,
710
1278
  }
711
1279
  }
712
1280
 
1281
+ template <int ncols_dst>
1282
+ static void mul_mat_vec_q2_K_q8_1_sycl_ncols(
1283
+ const void * vx, const void * vy, float * dst,
1284
+ const int ncols, const int nrows,
1285
+ const int stride_col_y, const int stride_col_dst,
1286
+ dpct::queue_ptr stream) {
1287
+ GGML_ASSERT(ncols % QK_K == 0);
1288
+ const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
1289
+ const sycl::range<3> block_nums(1, 1, block_num_y);
1290
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
1291
+ stream->submit([&](sycl::handler & cgh) {
1292
+ cgh.parallel_for(
1293
+ sycl::nd_range<3>(block_nums * block_dims, block_dims),
1294
+ [=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
1295
+ mul_mat_vec_q_ncols<QK_K, QI2_K, block_q2_K,
1296
+ VDR_Q2_K_Q8_1_MMVQ, vec_dot_q2_K_q8_1, ncols_dst>(
1297
+ vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, item_ct1);
1298
+ });
1299
+ });
1300
+ }
1301
+
1302
+ static void mul_mat_vec_q2_K_q8_1_sycl_switch_ncols(
1303
+ const void * vx, const void * vy, float * dst,
1304
+ const int ncols, const int nrows, const int ncols_dst,
1305
+ const int stride_col_y, const int stride_col_dst,
1306
+ dpct::queue_ptr stream) {
1307
+ switch (ncols_dst) {
1308
+ case 1: mul_mat_vec_q2_K_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
1309
+ case 2: mul_mat_vec_q2_K_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1310
+ case 3: mul_mat_vec_q2_K_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1311
+ case 4: mul_mat_vec_q2_K_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1312
+ case 5: mul_mat_vec_q2_K_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1313
+ case 6: mul_mat_vec_q2_K_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1314
+ case 7: mul_mat_vec_q2_K_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1315
+ case 8: mul_mat_vec_q2_K_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1316
+ default: GGML_ABORT("unsupported ncols_dst=%d for Q2_K multi-col MMVQ", ncols_dst);
1317
+ }
1318
+ }
1319
+
713
1320
  static void mul_mat_vec_q3_K_q8_1_sycl(const void *vx, const void *vy,
714
1321
  float *dst, const int ncols,
715
1322
  const int nrows,
@@ -734,6 +1341,104 @@ static void mul_mat_vec_q3_K_q8_1_sycl(const void *vx, const void *vy,
734
1341
  }
735
1342
  }
736
1343
 
1344
+ static void reorder_mul_mat_vec_q3_k_q8_1_sycl(const void * vx, const void * vy, float * dst, const int ncols,
1345
+ const int nrows, dpct::queue_ptr stream) {
1346
+ GGML_ASSERT(ncols % QK_K == 0);
1347
+
1348
+ // Round up to a whole number of subgroup-sized workgroups; out-of-range rows are skipped inside the kernel.
1349
+ constexpr size_t num_subgroups = WARP_SIZE;
1350
+ const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
1351
+ const sycl::range<3> block_nums(1, 1, block_num_y);
1352
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
1353
+
1354
+ stream->submit([&](sycl::handler & cgh) {
1355
+ cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
1356
+ [=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
1357
+ mul_mat_vec_q_reorder<reorder_vec_dot_q_sycl<GGML_TYPE_Q3_K>>(vx, vy, dst, ncols, nrows,
1358
+ nd_item);
1359
+ });
1360
+ });
1361
+ }
1362
+
1363
+ template <int ncols_dst>
1364
+ static void reorder_mul_mat_vec_q3_k_q8_1_sycl_ncols(
1365
+ const void * vx, const void * vy, float * dst,
1366
+ const int ncols, const int nrows,
1367
+ const int stride_col_y_bytes, const int stride_col_dst,
1368
+ dpct::queue_ptr stream) {
1369
+ GGML_ASSERT(ncols % QK_K == 0);
1370
+ constexpr size_t num_subgroups = WARP_SIZE;
1371
+ const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
1372
+ const sycl::range<3> block_nums(1, 1, block_num_y);
1373
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
1374
+
1375
+ stream->submit([&](sycl::handler & cgh) {
1376
+ cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
1377
+ [=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
1378
+ mul_mat_vec_q_reorder_ncols<reorder_vec_dot_q_sycl<GGML_TYPE_Q3_K>, ncols_dst>(
1379
+ vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, nd_item);
1380
+ });
1381
+ });
1382
+ }
1383
+
1384
+ static void reorder_mul_mat_vec_q3_k_q8_1_sycl_switch_ncols(
1385
+ const void * vx, const void * vy, float * dst,
1386
+ const int ncols, const int nrows, const int ncols_dst,
1387
+ const int stride_col_y_bytes, const int stride_col_dst,
1388
+ dpct::queue_ptr stream) {
1389
+ switch (ncols_dst) {
1390
+ case 1: reorder_mul_mat_vec_q3_k_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
1391
+ case 2: reorder_mul_mat_vec_q3_k_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1392
+ case 3: reorder_mul_mat_vec_q3_k_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1393
+ case 4: reorder_mul_mat_vec_q3_k_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1394
+ case 5: reorder_mul_mat_vec_q3_k_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1395
+ case 6: reorder_mul_mat_vec_q3_k_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1396
+ case 7: reorder_mul_mat_vec_q3_k_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1397
+ case 8: reorder_mul_mat_vec_q3_k_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1398
+ default: GGML_ABORT("unsupported ncols_dst=%d for Q3_K reorder multi-col MMVQ", ncols_dst);
1399
+ }
1400
+ }
1401
+
1402
+ template <int ncols_dst>
1403
+ static void mul_mat_vec_q3_K_q8_1_sycl_ncols(
1404
+ const void * vx, const void * vy, float * dst,
1405
+ const int ncols, const int nrows,
1406
+ const int stride_col_y, const int stride_col_dst,
1407
+ dpct::queue_ptr stream) {
1408
+ GGML_ASSERT(ncols % QK_K == 0);
1409
+ const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
1410
+ const sycl::range<3> block_nums(1, 1, block_num_y);
1411
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
1412
+ stream->submit([&](sycl::handler & cgh) {
1413
+ cgh.parallel_for(
1414
+ sycl::nd_range<3>(block_nums * block_dims, block_dims),
1415
+ [=](sycl::nd_item<3> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
1416
+ mul_mat_vec_q_ncols<QK_K, QI3_K, block_q3_K,
1417
+ VDR_Q3_K_Q8_1_MMVQ, vec_dot_q3_K_q8_1, ncols_dst>(
1418
+ vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, item_ct1);
1419
+ });
1420
+ });
1421
+ }
1422
+
1423
+ static void mul_mat_vec_q3_K_q8_1_sycl_switch_ncols(
1424
+ const void * vx, const void * vy, float * dst,
1425
+ const int ncols, const int nrows, const int ncols_dst,
1426
+ const int stride_col_y, const int stride_col_dst,
1427
+ dpct::queue_ptr stream) {
1428
+ switch (ncols_dst) {
1429
+ case 1: mul_mat_vec_q3_K_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
1430
+ case 2: mul_mat_vec_q3_K_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1431
+ case 3: mul_mat_vec_q3_K_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1432
+ case 4: mul_mat_vec_q3_K_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1433
+ case 5: mul_mat_vec_q3_K_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1434
+ case 6: mul_mat_vec_q3_K_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1435
+ case 7: mul_mat_vec_q3_K_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1436
+ case 8: mul_mat_vec_q3_K_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1437
+ default: GGML_ABORT("unsupported ncols_dst=%d for Q3_K multi-col MMVQ", ncols_dst);
1438
+ }
1439
+ }
1440
+
1441
+
737
1442
  static void mul_mat_vec_q4_K_q8_1_sycl(const void *vx, const void *vy,
738
1443
  float *dst, const int ncols,
739
1444
  const int nrows,
@@ -758,19 +1463,63 @@ static void mul_mat_vec_q4_K_q8_1_sycl(const void *vx, const void *vy,
758
1463
  }
759
1464
  }
760
1465
 
1466
+ template <int ncols_dst>
1467
+ static void mul_mat_vec_q4_K_q8_1_sycl_ncols(
1468
+ const void * vx, const void * vy, float * dst,
1469
+ const int ncols, const int nrows,
1470
+ const int stride_col_y, const int stride_col_dst,
1471
+ dpct::queue_ptr stream) {
1472
+ GGML_ASSERT(ncols % QK_K == 0);
1473
+ const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
1474
+ const sycl::range<3> block_nums(1, 1, block_num_y);
1475
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
1476
+
1477
+ stream->submit([&](sycl::handler & cgh) {
1478
+ cgh.parallel_for(
1479
+ sycl::nd_range<3>(block_nums * block_dims, block_dims),
1480
+ [=](sycl::nd_item<3> item_ct1)
1481
+ [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
1482
+ mul_mat_vec_q_ncols<QK_K, QI4_K, block_q4_K,
1483
+ VDR_Q4_K_Q8_1_MMVQ,
1484
+ vec_dot_q4_K_q8_1,
1485
+ ncols_dst>(
1486
+ vx, vy, dst, ncols, nrows,
1487
+ stride_col_y, stride_col_dst, item_ct1);
1488
+ });
1489
+ });
1490
+ }
1491
+
1492
+ static void mul_mat_vec_q4_K_q8_1_sycl_switch_ncols(
1493
+ const void * vx, const void * vy, float * dst,
1494
+ const int ncols, const int nrows,
1495
+ const int ncols_dst,
1496
+ const int stride_col_y, const int stride_col_dst,
1497
+ dpct::queue_ptr stream) {
1498
+ switch (ncols_dst) {
1499
+ case 1: mul_mat_vec_q4_K_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
1500
+ case 2: mul_mat_vec_q4_K_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1501
+ case 3: mul_mat_vec_q4_K_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1502
+ case 4: mul_mat_vec_q4_K_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1503
+ case 5: mul_mat_vec_q4_K_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1504
+ case 6: mul_mat_vec_q4_K_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1505
+ case 7: mul_mat_vec_q4_K_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1506
+ case 8: mul_mat_vec_q4_K_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1507
+ default: GGML_ABORT("unsupported ncols_dst=%d for Q4_K multi-col MMVQ", ncols_dst);
1508
+ }
1509
+ }
1510
+
761
1511
  static void reorder_mul_mat_vec_q4_k_q8_1_sycl(const void * vx, const void * vy, float * dst, const int ncols,
762
1512
  const int nrows, dpct::queue_ptr stream) {
763
1513
  GGML_ASSERT(ncols % QK_K == 0);
764
1514
 
765
- const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y);
766
- constexpr size_t num_subgroups = 16;
767
- GGML_ASSERT(block_num_y % num_subgroups == 0);
768
-
769
- const sycl::range<3> global_size(1, GGML_SYCL_MMV_Y, block_num_y * WARP_SIZE);
770
- const sycl::range<3> workgroup_size(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
1515
+ // Round up to a whole number of subgroup-sized workgroups; out-of-range rows are skipped inside the kernel.
1516
+ constexpr size_t num_subgroups = WARP_SIZE;
1517
+ const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
1518
+ const sycl::range<3> block_nums(1, 1, block_num_y);
1519
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
771
1520
 
772
1521
  stream->submit([&](sycl::handler & cgh) {
773
- cgh.parallel_for(sycl::nd_range<3>(global_size, workgroup_size),
1522
+ cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
774
1523
  [=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
775
1524
  mul_mat_vec_q_reorder<reorder_vec_dot_q_sycl<GGML_TYPE_Q4_K>>(vx, vy, dst, ncols,
776
1525
  nrows, nd_item);
@@ -778,6 +1527,45 @@ static void reorder_mul_mat_vec_q4_k_q8_1_sycl(const void * vx, const void * vy,
778
1527
  });
779
1528
  }
780
1529
 
1530
+ template <int ncols_dst>
1531
+ static void reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols(
1532
+ const void * vx, const void * vy, float * dst,
1533
+ const int ncols, const int nrows,
1534
+ const int stride_col_y_bytes, const int stride_col_dst,
1535
+ dpct::queue_ptr stream) {
1536
+ GGML_ASSERT(ncols % QK_K == 0);
1537
+
1538
+ constexpr size_t num_subgroups = WARP_SIZE;
1539
+ const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
1540
+ const sycl::range<3> block_nums(1, 1, block_num_y);
1541
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
1542
+
1543
+ stream->submit([&](sycl::handler & cgh) {
1544
+ cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
1545
+ [=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
1546
+ mul_mat_vec_q_reorder_ncols<reorder_vec_dot_q_sycl<GGML_TYPE_Q4_K>, ncols_dst>(
1547
+ vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, nd_item);
1548
+ });
1549
+ });
1550
+ }
1551
+
1552
+ static void reorder_mul_mat_vec_q4_k_q8_1_sycl_switch_ncols(
1553
+ const void * vx, const void * vy, float * dst,
1554
+ const int ncols, const int nrows, const int ncols_dst,
1555
+ const int stride_col_y_bytes, const int stride_col_dst,
1556
+ dpct::queue_ptr stream) {
1557
+ switch (ncols_dst) {
1558
+ case 1: reorder_mul_mat_vec_q4_k_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
1559
+ case 2: reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1560
+ case 3: reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1561
+ case 4: reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1562
+ case 5: reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1563
+ case 6: reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1564
+ case 7: reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1565
+ case 8: reorder_mul_mat_vec_q4_k_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1566
+ default: GGML_ABORT("unsupported ncols_dst=%d for Q4_K reorder multi-col MMVQ", ncols_dst);
1567
+ }
1568
+ }
781
1569
 
782
1570
  static void mul_mat_vec_q5_K_q8_1_sycl(const void *vx, const void *vy,
783
1571
  float *dst, const int ncols,
@@ -803,24 +1591,167 @@ static void mul_mat_vec_q5_K_q8_1_sycl(const void *vx, const void *vy,
803
1591
  }
804
1592
  }
805
1593
 
1594
+ template <int ncols_dst>
1595
+ static void mul_mat_vec_q5_K_q8_1_sycl_ncols(
1596
+ const void * vx, const void * vy, float * dst,
1597
+ const int ncols, const int nrows,
1598
+ const int stride_col_y, const int stride_col_dst,
1599
+ dpct::queue_ptr stream) {
1600
+ GGML_ASSERT(ncols % QK_K == 0);
1601
+ const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
1602
+ const sycl::range<3> block_nums(1, 1, block_num_y);
1603
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
1604
+
1605
+ stream->submit([&](sycl::handler & cgh) {
1606
+ cgh.parallel_for(
1607
+ sycl::nd_range<3>(block_nums * block_dims, block_dims),
1608
+ [=](sycl::nd_item<3> item_ct1)
1609
+ [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
1610
+ mul_mat_vec_q_ncols<QK_K, QI5_K, block_q5_K,
1611
+ VDR_Q5_K_Q8_1_MMVQ,
1612
+ vec_dot_q5_K_q8_1,
1613
+ ncols_dst>(
1614
+ vx, vy, dst, ncols, nrows,
1615
+ stride_col_y, stride_col_dst, item_ct1);
1616
+ });
1617
+ });
1618
+ }
1619
+
1620
+ static void mul_mat_vec_q5_K_q8_1_sycl_switch_ncols(
1621
+ const void * vx, const void * vy, float * dst,
1622
+ const int ncols, const int nrows,
1623
+ const int ncols_dst,
1624
+ const int stride_col_y, const int stride_col_dst,
1625
+ dpct::queue_ptr stream) {
1626
+ switch (ncols_dst) {
1627
+ case 1: mul_mat_vec_q5_K_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
1628
+ case 2: mul_mat_vec_q5_K_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1629
+ case 3: mul_mat_vec_q5_K_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1630
+ case 4: mul_mat_vec_q5_K_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1631
+ case 5: mul_mat_vec_q5_K_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1632
+ case 6: mul_mat_vec_q5_K_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1633
+ case 7: mul_mat_vec_q5_K_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1634
+ case 8: mul_mat_vec_q5_K_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1635
+ default: GGML_ABORT("unsupported ncols_dst=%d for Q5_K multi-col MMVQ", ncols_dst);
1636
+ }
1637
+ }
1638
+
1639
+ static void reorder_mul_mat_vec_q5_k_q8_1_sycl(const void * vx, const void * vy, float * dst, const int ncols,
1640
+ const int nrows, dpct::queue_ptr stream) {
1641
+ GGML_ASSERT(ncols % QK_K == 0);
1642
+
1643
+ constexpr size_t num_subgroups = WARP_SIZE;
1644
+ const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
1645
+ const sycl::range<3> block_nums(1, 1, block_num_y);
1646
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
1647
+
1648
+ stream->submit([&](sycl::handler & cgh) {
1649
+ cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
1650
+ [=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
1651
+ mul_mat_vec_q_reorder<reorder_vec_dot_q_sycl<GGML_TYPE_Q5_K>>(vx, vy, dst, ncols,
1652
+ nrows, nd_item);
1653
+ });
1654
+ });
1655
+ }
1656
+
1657
+ template <int ncols_dst>
1658
+ static void reorder_mul_mat_vec_q5_k_q8_1_sycl_ncols(
1659
+ const void * vx, const void * vy, float * dst,
1660
+ const int ncols, const int nrows,
1661
+ const int stride_col_y_bytes, const int stride_col_dst,
1662
+ dpct::queue_ptr stream) {
1663
+ GGML_ASSERT(ncols % QK_K == 0);
1664
+
1665
+ constexpr size_t num_subgroups = WARP_SIZE;
1666
+ const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
1667
+ const sycl::range<3> block_nums(1, 1, block_num_y);
1668
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
1669
+
1670
+ stream->submit([&](sycl::handler & cgh) {
1671
+ cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
1672
+ [=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
1673
+ mul_mat_vec_q_reorder_ncols<reorder_vec_dot_q_sycl<GGML_TYPE_Q5_K>, ncols_dst>(
1674
+ vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, nd_item);
1675
+ });
1676
+ });
1677
+ }
1678
+
1679
+ static void reorder_mul_mat_vec_q5_k_q8_1_sycl_switch_ncols(
1680
+ const void * vx, const void * vy, float * dst,
1681
+ const int ncols, const int nrows, const int ncols_dst,
1682
+ const int stride_col_y_bytes, const int stride_col_dst,
1683
+ dpct::queue_ptr stream) {
1684
+ switch (ncols_dst) {
1685
+ case 1: reorder_mul_mat_vec_q5_k_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
1686
+ case 2: reorder_mul_mat_vec_q5_k_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1687
+ case 3: reorder_mul_mat_vec_q5_k_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1688
+ case 4: reorder_mul_mat_vec_q5_k_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1689
+ case 5: reorder_mul_mat_vec_q5_k_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1690
+ case 6: reorder_mul_mat_vec_q5_k_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1691
+ case 7: reorder_mul_mat_vec_q5_k_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1692
+ case 8: reorder_mul_mat_vec_q5_k_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1693
+ default: GGML_ABORT("unsupported ncols_dst=%d for Q5_K reorder multi-col MMVQ", ncols_dst);
1694
+ }
1695
+ }
1696
+
806
1697
  static void reorder_mul_mat_vec_q6_k_q8_1_sycl(const void * vx, const void * vy, float * dst, const int ncols,
807
1698
  const int nrows, dpct::queue_ptr stream) {
808
1699
  GGML_ASSERT(ncols % QK_K == 0);
809
- const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y);
810
- constexpr size_t num_subgroups = 16;
811
- GGML_ASSERT(block_num_y % num_subgroups == 0);
1700
+ // Round up to a whole number of subgroup-sized workgroups; out-of-range rows are skipped inside the kernel.
1701
+ constexpr size_t num_subgroups = WARP_SIZE;
1702
+ const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
1703
+ const sycl::range<3> block_nums(1, 1, block_num_y);
1704
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
812
1705
 
813
- const sycl::range<3> global_size(1, GGML_SYCL_MMV_Y, block_num_y * WARP_SIZE);
814
- const sycl::range<3> workgroup_size(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
815
1706
 
816
1707
  stream->submit([&](sycl::handler & cgh) {
817
- cgh.parallel_for(sycl::nd_range<3>(global_size, workgroup_size),
1708
+ cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
818
1709
  [=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
819
1710
  mul_mat_vec_q_reorder<reorder_vec_dot_q_sycl<GGML_TYPE_Q6_K>>(vx, vy, dst, ncols, nrows,
820
1711
  nd_item);
821
1712
  });
822
1713
  });
823
1714
  }
1715
+
1716
+ template <int ncols_dst>
1717
+ static void reorder_mul_mat_vec_q6_k_q8_1_sycl_ncols(
1718
+ const void * vx, const void * vy, float * dst,
1719
+ const int ncols, const int nrows,
1720
+ const int stride_col_y_bytes, const int stride_col_dst,
1721
+ dpct::queue_ptr stream) {
1722
+ GGML_ASSERT(ncols % QK_K == 0);
1723
+ constexpr size_t num_subgroups = WARP_SIZE;
1724
+ const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
1725
+ const sycl::range<3> block_nums(1, 1, block_num_y);
1726
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
1727
+
1728
+ stream->submit([&](sycl::handler & cgh) {
1729
+ cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
1730
+ [=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
1731
+ mul_mat_vec_q_reorder_ncols<reorder_vec_dot_q_sycl<GGML_TYPE_Q6_K>, ncols_dst>(
1732
+ vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, nd_item);
1733
+ });
1734
+ });
1735
+ }
1736
+
1737
+ static void reorder_mul_mat_vec_q6_k_q8_1_sycl_switch_ncols(
1738
+ const void * vx, const void * vy, float * dst,
1739
+ const int ncols, const int nrows, const int ncols_dst,
1740
+ const int stride_col_y_bytes, const int stride_col_dst,
1741
+ dpct::queue_ptr stream) {
1742
+ switch (ncols_dst) {
1743
+ case 1: reorder_mul_mat_vec_q6_k_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
1744
+ case 2: reorder_mul_mat_vec_q6_k_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1745
+ case 3: reorder_mul_mat_vec_q6_k_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1746
+ case 4: reorder_mul_mat_vec_q6_k_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1747
+ case 5: reorder_mul_mat_vec_q6_k_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1748
+ case 6: reorder_mul_mat_vec_q6_k_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1749
+ case 7: reorder_mul_mat_vec_q6_k_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1750
+ case 8: reorder_mul_mat_vec_q6_k_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream); break;
1751
+ default: GGML_ABORT("unsupported ncols_dst=%d for Q6_K reorder multi-col MMVQ", ncols_dst);
1752
+ }
1753
+ }
1754
+
824
1755
  static void mul_mat_vec_q6_K_q8_1_sycl(const void *vx, const void *vy,
825
1756
  float *dst, const int ncols,
826
1757
  const int nrows,
@@ -845,6 +1776,51 @@ static void mul_mat_vec_q6_K_q8_1_sycl(const void *vx, const void *vy,
845
1776
  }
846
1777
  }
847
1778
 
1779
+ template <int ncols_dst>
1780
+ static void mul_mat_vec_q6_K_q8_1_sycl_ncols(
1781
+ const void * vx, const void * vy, float * dst,
1782
+ const int ncols, const int nrows,
1783
+ const int stride_col_y, const int stride_col_dst,
1784
+ dpct::queue_ptr stream) {
1785
+ GGML_ASSERT(ncols % QK_K == 0);
1786
+ const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
1787
+ const sycl::range<3> block_nums(1, 1, block_num_y);
1788
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
1789
+
1790
+ stream->submit([&](sycl::handler & cgh) {
1791
+ cgh.parallel_for(
1792
+ sycl::nd_range<3>(block_nums * block_dims, block_dims),
1793
+ [=](sycl::nd_item<3> item_ct1)
1794
+ [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
1795
+ mul_mat_vec_q_ncols<QK_K, QI6_K, block_q6_K,
1796
+ VDR_Q6_K_Q8_1_MMVQ,
1797
+ vec_dot_q6_K_q8_1,
1798
+ ncols_dst>(
1799
+ vx, vy, dst, ncols, nrows,
1800
+ stride_col_y, stride_col_dst, item_ct1);
1801
+ });
1802
+ });
1803
+ }
1804
+
1805
+ static void mul_mat_vec_q6_K_q8_1_sycl_switch_ncols(
1806
+ const void * vx, const void * vy, float * dst,
1807
+ const int ncols, const int nrows,
1808
+ const int ncols_dst,
1809
+ const int stride_col_y, const int stride_col_dst,
1810
+ dpct::queue_ptr stream) {
1811
+ switch (ncols_dst) {
1812
+ case 1: mul_mat_vec_q6_K_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
1813
+ case 2: mul_mat_vec_q6_K_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1814
+ case 3: mul_mat_vec_q6_K_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1815
+ case 4: mul_mat_vec_q6_K_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1816
+ case 5: mul_mat_vec_q6_K_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1817
+ case 6: mul_mat_vec_q6_K_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1818
+ case 7: mul_mat_vec_q6_K_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1819
+ case 8: mul_mat_vec_q6_K_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
1820
+ default: GGML_ABORT("unsupported ncols_dst=%d for Q6_K multi-col MMVQ", ncols_dst);
1821
+ }
1822
+ }
1823
+
848
1824
 
849
1825
  static void mul_mat_vec_iq2_xxs_q8_1_sycl(const void *vx, const void *vy,
850
1826
  float *dst, const int ncols,
@@ -1041,6 +2017,51 @@ static void mul_mat_vec_iq4_xs_q8_1_sycl(const void *vx, const void *vy,
1041
2017
  }
1042
2018
  }
1043
2019
 
2020
+ template <int ncols_dst>
2021
+ static void mul_mat_vec_iq4_xs_q8_1_sycl_ncols(
2022
+ const void * vx, const void * vy, float * dst,
2023
+ const int ncols, const int nrows,
2024
+ const int stride_col_y, const int stride_col_dst,
2025
+ dpct::queue_ptr stream) {
2026
+ GGML_ASSERT(ncols % QK_K == 0);
2027
+ const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
2028
+ const sycl::range<3> block_nums(1, 1, block_num_y);
2029
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
2030
+
2031
+ stream->submit([&](sycl::handler & cgh) {
2032
+ cgh.parallel_for(
2033
+ sycl::nd_range<3>(block_nums * block_dims, block_dims),
2034
+ [=](sycl::nd_item<3> item_ct1)
2035
+ [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
2036
+ mul_mat_vec_q_ncols<QK_K, QI4_XS/4, block_iq4_xs,
2037
+ 1,
2038
+ vec_dot_iq4_xs_q8_1,
2039
+ ncols_dst>(
2040
+ vx, vy, dst, ncols, nrows,
2041
+ stride_col_y, stride_col_dst, item_ct1);
2042
+ });
2043
+ });
2044
+ }
2045
+
2046
+ static void mul_mat_vec_iq4_xs_q8_1_sycl_switch_ncols(
2047
+ const void * vx, const void * vy, float * dst,
2048
+ const int ncols, const int nrows,
2049
+ const int ncols_dst,
2050
+ const int stride_col_y, const int stride_col_dst,
2051
+ dpct::queue_ptr stream) {
2052
+ switch (ncols_dst) {
2053
+ case 1: mul_mat_vec_iq4_xs_q8_1_sycl(vx, vy, dst, ncols, nrows, stream); break;
2054
+ case 2: mul_mat_vec_iq4_xs_q8_1_sycl_ncols<2>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
2055
+ case 3: mul_mat_vec_iq4_xs_q8_1_sycl_ncols<3>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
2056
+ case 4: mul_mat_vec_iq4_xs_q8_1_sycl_ncols<4>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
2057
+ case 5: mul_mat_vec_iq4_xs_q8_1_sycl_ncols<5>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
2058
+ case 6: mul_mat_vec_iq4_xs_q8_1_sycl_ncols<6>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
2059
+ case 7: mul_mat_vec_iq4_xs_q8_1_sycl_ncols<7>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
2060
+ case 8: mul_mat_vec_iq4_xs_q8_1_sycl_ncols<8>(vx, vy, dst, ncols, nrows, stride_col_y, stride_col_dst, stream); break;
2061
+ default: GGML_ABORT("unsupported ncols_dst=%d for IQ4_XS multi-col MMVQ", ncols_dst);
2062
+ }
2063
+ }
2064
+
1044
2065
  void ggml_sycl_op_mul_mat_vec_q(ggml_backend_sycl_context & ctx, const ggml_tensor * src0, const ggml_tensor * src1,
1045
2066
  ggml_tensor * dst, const char * src0_dd_i, const float * src1_ddf_i,
1046
2067
  const char * src1_ddq_i, float * dst_dd_i, const int64_t row_low,
@@ -1067,50 +2088,233 @@ void ggml_sycl_op_mul_mat_vec_q(ggml_backend_sycl_context & ctx, const ggml_tens
1067
2088
  case GGML_TYPE_Q4_0:
1068
2089
  if ((ggml_tensor_extra_gpu *) dst->src[0]->extra &&
1069
2090
  ((ggml_tensor_extra_gpu *) dst->src[0]->extra)->optimized_feature.reorder) {
1070
- GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q4_0_q8_1_sycl\n");
1071
- reorder_mul_mat_vec_q4_0_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
1072
- } else {
2091
+ if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
2092
+ const int stride_col_y_bytes = src1_padded_col_size * q8_1_ts / q8_1_bs;
2093
+ const int stride_col_dst = dst->ne[0];
2094
+ GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q4_0_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
2095
+ reorder_mul_mat_vec_q4_0_q8_1_sycl_switch_ncols(
2096
+ src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
2097
+ src1_ncols, stride_col_y_bytes, stride_col_dst, stream);
2098
+ return;
2099
+ } else {
2100
+ GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q4_0_q8_1_sycl\n");
2101
+ reorder_mul_mat_vec_q4_0_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2102
+ }
2103
+ } else if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
2104
+ const int stride_col_y = src1_padded_col_size / QK8_1;
2105
+ const int stride_col_dst = dst->ne[0];
2106
+ GGML_SYCL_DEBUG("Calling mul_mat_vec_q4_0_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
2107
+ mul_mat_vec_q4_0_q8_1_sycl_switch_ncols(
2108
+ src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
2109
+ src1_ncols, stride_col_y, stride_col_dst, stream);
2110
+ return;
2111
+ } else if (i == 0 || src1_ncols == 1) {
1073
2112
  GGML_SYCL_DEBUG("Calling mul_mat_vec_q4_0_q8_1_sycl\n");
1074
2113
  mul_mat_vec_q4_0_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
1075
2114
  }
1076
2115
  break;
1077
2116
  case GGML_TYPE_Q4_1:
1078
- mul_mat_vec_q4_1_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2117
+ if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
2118
+ const int stride_col_y = src1_padded_col_size / QK8_1;
2119
+ const int stride_col_dst = dst->ne[0];
2120
+ GGML_SYCL_DEBUG("Calling mul_mat_vec_q4_1_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
2121
+ mul_mat_vec_q4_1_q8_1_sycl_switch_ncols(
2122
+ src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
2123
+ src1_ncols, stride_col_y, stride_col_dst, stream);
2124
+ return;
2125
+ } else if (i == 0 || src1_ncols == 1) {
2126
+ mul_mat_vec_q4_1_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2127
+ }
1079
2128
  break;
1080
2129
  case GGML_TYPE_Q5_0:
1081
- mul_mat_vec_q5_0_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2130
+ if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
2131
+ const int stride_col_y = src1_padded_col_size / QK8_1;
2132
+ const int stride_col_dst = dst->ne[0];
2133
+ GGML_SYCL_DEBUG("Calling mul_mat_vec_q5_0_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
2134
+ mul_mat_vec_q5_0_q8_1_sycl_switch_ncols(
2135
+ src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
2136
+ src1_ncols, stride_col_y, stride_col_dst, stream);
2137
+ return;
2138
+ } else if (i == 0 || src1_ncols == 1) {
2139
+ mul_mat_vec_q5_0_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2140
+ }
1082
2141
  break;
1083
2142
  case GGML_TYPE_Q5_1:
1084
- mul_mat_vec_q5_1_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2143
+ if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
2144
+ const int stride_col_y = src1_padded_col_size / QK8_1;
2145
+ const int stride_col_dst = dst->ne[0];
2146
+ GGML_SYCL_DEBUG("Calling mul_mat_vec_q5_1_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
2147
+ mul_mat_vec_q5_1_q8_1_sycl_switch_ncols(
2148
+ src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
2149
+ src1_ncols, stride_col_y, stride_col_dst, stream);
2150
+ return;
2151
+ } else if (i == 0 || src1_ncols == 1) {
2152
+ mul_mat_vec_q5_1_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2153
+ }
1085
2154
  break;
1086
2155
  case GGML_TYPE_Q8_0:
1087
- mul_mat_vec_q8_0_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2156
+ if ((ggml_tensor_extra_gpu *) dst->src[0]->extra &&
2157
+ ((ggml_tensor_extra_gpu *) dst->src[0]->extra)->optimized_feature.reorder) {
2158
+ if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
2159
+ const int stride_col_y_bytes = src1_padded_col_size * q8_1_ts / q8_1_bs;
2160
+ const int stride_col_dst = dst->ne[0];
2161
+ GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q8_0_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
2162
+ reorder_mul_mat_vec_q8_0_q8_1_sycl_switch_ncols(
2163
+ src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
2164
+ src1_ncols, stride_col_y_bytes, stride_col_dst, stream);
2165
+ return;
2166
+ } else {
2167
+ GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q8_0_q8_1_sycl\n");
2168
+ reorder_mul_mat_vec_q8_0_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2169
+ }
2170
+ } else if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
2171
+ const int stride_col_y = src1_padded_col_size / QK8_1;
2172
+ const int stride_col_dst = dst->ne[0];
2173
+ GGML_SYCL_DEBUG("Calling mul_mat_vec_q8_0_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
2174
+ mul_mat_vec_q8_0_q8_1_sycl_switch_ncols(
2175
+ src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
2176
+ src1_ncols, stride_col_y, stride_col_dst, stream);
2177
+ return;
2178
+ } else if (i == 0 || src1_ncols == 1) {
2179
+ GGML_SYCL_DEBUG("Calling mul_mat_vec_q8_0_q8_1_sycl\n");
2180
+ mul_mat_vec_q8_0_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2181
+ }
2182
+ break;
2183
+ case GGML_TYPE_Q1_0:
2184
+ if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
2185
+ const int stride_col_y = src1_padded_col_size / QK8_1;
2186
+ const int stride_col_dst = dst->ne[0];
2187
+ GGML_SYCL_DEBUG("Calling mul_mat_vec_q1_0_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
2188
+ mul_mat_vec_q1_0_q8_1_sycl_switch_ncols(
2189
+ src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
2190
+ src1_ncols, stride_col_y, stride_col_dst, stream);
2191
+ return;
2192
+ } else if (i == 0 || src1_ncols == 1) {
2193
+ GGML_SYCL_DEBUG("Calling mul_mat_vec_q1_0_q8_1_sycl\n");
2194
+ mul_mat_vec_q1_0_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2195
+ }
1088
2196
  break;
1089
2197
  case GGML_TYPE_Q2_K:
1090
- mul_mat_vec_q2_K_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2198
+ if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
2199
+ const int stride_col_y = src1_padded_col_size / QK8_1;
2200
+ const int stride_col_dst = dst->ne[0];
2201
+ GGML_SYCL_DEBUG("Calling mul_mat_vec_q2_K_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
2202
+ mul_mat_vec_q2_K_q8_1_sycl_switch_ncols(
2203
+ src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
2204
+ src1_ncols, stride_col_y, stride_col_dst, stream);
2205
+ return;
2206
+ } else if (i == 0 || src1_ncols == 1) {
2207
+ mul_mat_vec_q2_K_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2208
+ }
1091
2209
  break;
1092
2210
  case GGML_TYPE_Q3_K:
1093
- mul_mat_vec_q3_K_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2211
+ if ((ggml_tensor_extra_gpu *) dst->src[0]->extra &&
2212
+ ((ggml_tensor_extra_gpu *) dst->src[0]->extra)->optimized_feature.reorder) {
2213
+ if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
2214
+ const int stride_col_y_bytes = src1_padded_col_size * q8_1_ts / q8_1_bs;
2215
+ const int stride_col_dst = dst->ne[0];
2216
+ GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q3_k_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
2217
+ reorder_mul_mat_vec_q3_k_q8_1_sycl_switch_ncols(
2218
+ src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
2219
+ src1_ncols, stride_col_y_bytes, stride_col_dst, stream);
2220
+ return;
2221
+ } else {
2222
+ GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q3_k_q8_1_sycl\n");
2223
+ reorder_mul_mat_vec_q3_k_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2224
+ }
2225
+ } else if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
2226
+ const int stride_col_y = src1_padded_col_size / QK8_1;
2227
+ const int stride_col_dst = dst->ne[0];
2228
+ GGML_SYCL_DEBUG("Calling mul_mat_vec_q3_K_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
2229
+ mul_mat_vec_q3_K_q8_1_sycl_switch_ncols(
2230
+ src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
2231
+ src1_ncols, stride_col_y, stride_col_dst, stream);
2232
+ return;
2233
+ } else if (i == 0 || src1_ncols == 1) {
2234
+ GGML_SYCL_DEBUG("Calling mul_mat_vec_q3_K_q8_1_sycl\n");
2235
+ mul_mat_vec_q3_K_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2236
+ }
1094
2237
  break;
1095
2238
  case GGML_TYPE_Q4_K:
1096
2239
  if ((ggml_tensor_extra_gpu *) dst->src[0]->extra &&
1097
2240
  ((ggml_tensor_extra_gpu *) dst->src[0]->extra)->optimized_feature.reorder) {
1098
- GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q4_k_q8_1_sycl\n");
1099
- reorder_mul_mat_vec_q4_k_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
1100
- } else {
2241
+ if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
2242
+ const int stride_col_y_bytes = src1_padded_col_size * q8_1_ts / q8_1_bs;
2243
+ const int stride_col_dst = dst->ne[0];
2244
+ GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q4_k_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
2245
+ reorder_mul_mat_vec_q4_k_q8_1_sycl_switch_ncols(
2246
+ src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
2247
+ src1_ncols, stride_col_y_bytes, stride_col_dst, stream);
2248
+ return;
2249
+ } else {
2250
+ GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q4_k_q8_1_sycl\n");
2251
+ reorder_mul_mat_vec_q4_k_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2252
+ }
2253
+ } else if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
2254
+ const int stride_col_y = src1_padded_col_size / QK8_1;
2255
+ const int stride_col_dst = dst->ne[0];
2256
+ GGML_SYCL_DEBUG("Calling mul_mat_vec_q4_K_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
2257
+ mul_mat_vec_q4_K_q8_1_sycl_switch_ncols(
2258
+ src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
2259
+ src1_ncols, stride_col_y, stride_col_dst, stream);
2260
+ return;
2261
+ } else if (i == 0 || src1_ncols == 1) {
1101
2262
  GGML_SYCL_DEBUG("Calling mul_mat_vec_q4_K_q8_1_sycl\n");
1102
2263
  mul_mat_vec_q4_K_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
1103
2264
  }
1104
2265
  break;
1105
2266
  case GGML_TYPE_Q5_K:
1106
- mul_mat_vec_q5_K_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2267
+ if ((ggml_tensor_extra_gpu *) dst->src[0]->extra &&
2268
+ ((ggml_tensor_extra_gpu *) dst->src[0]->extra)->optimized_feature.reorder) {
2269
+ if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
2270
+ const int stride_col_y_bytes = src1_padded_col_size * q8_1_ts / q8_1_bs;
2271
+ const int stride_col_dst = dst->ne[0];
2272
+ GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q5_k_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
2273
+ reorder_mul_mat_vec_q5_k_q8_1_sycl_switch_ncols(
2274
+ src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
2275
+ src1_ncols, stride_col_y_bytes, stride_col_dst, stream);
2276
+ return;
2277
+ } else {
2278
+ GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q5_k_q8_1_sycl\n");
2279
+ reorder_mul_mat_vec_q5_k_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2280
+ }
2281
+ } else if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
2282
+ const int stride_col_y = src1_padded_col_size / QK8_1;
2283
+ const int stride_col_dst = dst->ne[0];
2284
+ GGML_SYCL_DEBUG("Calling mul_mat_vec_q5_K_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
2285
+ mul_mat_vec_q5_K_q8_1_sycl_switch_ncols(
2286
+ src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
2287
+ src1_ncols, stride_col_y, stride_col_dst, stream);
2288
+ return;
2289
+ } else if (i == 0 || src1_ncols == 1) {
2290
+ GGML_SYCL_DEBUG("Calling mul_mat_vec_q5_K_q8_1_sycl\n");
2291
+ mul_mat_vec_q5_K_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2292
+ }
1107
2293
  break;
1108
2294
  case GGML_TYPE_Q6_K:
1109
2295
  if ((ggml_tensor_extra_gpu *) dst->src[0]->extra &&
1110
2296
  ((ggml_tensor_extra_gpu *) dst->src[0]->extra)->optimized_feature.reorder) {
1111
- GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q6_k_q8_1_sycl\n");
1112
- reorder_mul_mat_vec_q6_k_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
1113
- } else {
2297
+ if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
2298
+ const int stride_col_y_bytes = src1_padded_col_size * q8_1_ts / q8_1_bs;
2299
+ const int stride_col_dst = dst->ne[0];
2300
+ GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q6_k_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
2301
+ reorder_mul_mat_vec_q6_k_q8_1_sycl_switch_ncols(
2302
+ src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
2303
+ src1_ncols, stride_col_y_bytes, stride_col_dst, stream);
2304
+ return;
2305
+ } else {
2306
+ GGML_SYCL_DEBUG("Calling reorder_mul_mat_vec_q6_k_q8_1_sycl\n");
2307
+ reorder_mul_mat_vec_q6_k_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2308
+ }
2309
+ } else if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
2310
+ const int stride_col_y = src1_padded_col_size / QK8_1;
2311
+ const int stride_col_dst = dst->ne[0];
2312
+ GGML_SYCL_DEBUG("Calling mul_mat_vec_q6_K_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
2313
+ mul_mat_vec_q6_K_q8_1_sycl_switch_ncols(
2314
+ src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
2315
+ src1_ncols, stride_col_y, stride_col_dst, stream);
2316
+ return;
2317
+ } else if (i == 0 || src1_ncols == 1) {
1114
2318
  GGML_SYCL_DEBUG("Calling mul_mat_vec_q6_k_q8_1_sycl\n");
1115
2319
  mul_mat_vec_q6_K_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
1116
2320
  }
@@ -1140,13 +2344,46 @@ void ggml_sycl_op_mul_mat_vec_q(ggml_backend_sycl_context & ctx, const ggml_tens
1140
2344
  mul_mat_vec_iq4_nl_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
1141
2345
  break;
1142
2346
  case GGML_TYPE_IQ4_XS:
1143
- mul_mat_vec_iq4_xs_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2347
+ if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
2348
+ const int stride_col_y = src1_padded_col_size / QK8_1;
2349
+ const int stride_col_dst = dst->ne[0];
2350
+ GGML_SYCL_DEBUG("Calling mul_mat_vec_iq4_xs_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
2351
+ mul_mat_vec_iq4_xs_q8_1_sycl_switch_ncols(
2352
+ src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
2353
+ src1_ncols, stride_col_y, stride_col_dst, stream);
2354
+ return;
2355
+ } else if (i == 0 || src1_ncols == 1) {
2356
+ mul_mat_vec_iq4_xs_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2357
+ }
1144
2358
  break;
1145
2359
  case GGML_TYPE_MXFP4:
1146
- mul_mat_vec_mxfp4_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2360
+ if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
2361
+ const int stride_col_y = src1_padded_col_size / QK8_1;
2362
+ const int stride_col_dst = dst->ne[0];
2363
+ GGML_SYCL_DEBUG("Calling mul_mat_vec_mxfp4_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
2364
+ mul_mat_vec_mxfp4_q8_1_sycl_switch_ncols(
2365
+ src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
2366
+ src1_ncols, stride_col_y, stride_col_dst, stream);
2367
+ return;
2368
+ } else if (i == 0 || src1_ncols == 1) {
2369
+ mul_mat_vec_mxfp4_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2370
+ }
2371
+ break;
2372
+ case GGML_TYPE_NVFP4:
2373
+ if (i == 0 && src1_ncols > 1 && src1_ncols <= 8) {
2374
+ const int stride_col_y = src1_padded_col_size / QK8_1;
2375
+ const int stride_col_dst = dst->ne[0];
2376
+ GGML_SYCL_DEBUG("Calling mul_mat_vec_nvfp4_q8_1_sycl_switch_ncols ncols=%d\n", (int)src1_ncols);
2377
+ mul_mat_vec_nvfp4_q8_1_sycl_switch_ncols(
2378
+ src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff,
2379
+ src1_ncols, stride_col_y, stride_col_dst, stream);
2380
+ return;
2381
+ } else if (i == 0 || src1_ncols == 1) {
2382
+ mul_mat_vec_nvfp4_q8_1_sycl(src0_dd_i, src1_ddq_i_bs, dst_dd_i_bs, ne00, row_diff, stream);
2383
+ }
1147
2384
  break;
1148
2385
  default:
1149
- GGML_ABORT("fatal error");
2386
+ GGML_ABORT("fatal error: unsupport data type=%s\n", ggml_type_name(src0->type));
1150
2387
  }
1151
2388
  }
1152
2389
  GGML_UNUSED(src1);
@@ -1154,3 +2391,269 @@ void ggml_sycl_op_mul_mat_vec_q(ggml_backend_sycl_context & ctx, const ggml_tens
1154
2391
  GGML_UNUSED(src1_ddf_i);
1155
2392
  GGML_UNUSED(ctx);
1156
2393
  }
2394
+
2395
+ // src1_row_stride: 0 for shared src1 (gate/up proj), else per-expert stride (down proj).
2396
+ template <int qk, int qi, typename block_q_t, int vdr, vec_dot_q_sycl_t vec_dot_q_sycl>
2397
+ static void mul_mat_vec_q_moe(
2398
+ const void * __restrict__ vx_base, const void * __restrict__ vy_base,
2399
+ float * __restrict__ dst_base, const int32_t * __restrict__ ids_dev,
2400
+ const int ncols, const int nrows,
2401
+ const size_t expert_weight_stride, const size_t dst_row_stride,
2402
+ const size_t src1_row_stride,
2403
+ const sycl::nd_item<3> & item_ct1) {
2404
+
2405
+ const int expert_idx = item_ct1.get_group(1);
2406
+ const int i02 = ids_dev[expert_idx];
2407
+
2408
+ const char * vx = (const char *) vx_base + (size_t) i02 * expert_weight_stride;
2409
+ const char * vy = (const char *) vy_base + (size_t) expert_idx * src1_row_stride;
2410
+ float * dst = (float *) ((char *) dst_base + (size_t) expert_idx * dst_row_stride);
2411
+
2412
+ const int row = item_ct1.get_group(2) * item_ct1.get_local_range(1) + item_ct1.get_local_id(1);
2413
+
2414
+ if (row >= nrows) {
2415
+ return;
2416
+ }
2417
+
2418
+ const int blocks_per_row = ncols / qk;
2419
+ constexpr int blocks_per_warp = (vdr * WARP_SIZE + qi - 1) / qi;
2420
+
2421
+ float tmp = 0.0f;
2422
+
2423
+ const block_q_t * x = (const block_q_t *) vx;
2424
+ const block_q8_1 * y = (const block_q8_1 *) vy;
2425
+
2426
+ for (int i = item_ct1.get_local_id(2) / (qi / vdr); i < blocks_per_row; i += blocks_per_warp) {
2427
+ const int ibx = row * blocks_per_row + i;
2428
+ const int iby = i * (qk / QK8_1);
2429
+
2430
+ for (size_t elem = 0; elem < qi / vdr; elem += WARP_SIZE) {
2431
+ const int iqs = elem + vdr * (item_ct1.get_local_id(2) % (qi / vdr));
2432
+ tmp += vec_dot_q_sycl(&x[ibx], &y[iby], iqs);
2433
+ }
2434
+ }
2435
+
2436
+ #pragma unroll
2437
+ for (int mask = WARP_SIZE / 2; mask > 0; mask >>= 1) {
2438
+ tmp += dpct::permute_sub_group_by_xor(item_ct1.get_sub_group(), tmp, mask);
2439
+ }
2440
+
2441
+ if (item_ct1.get_local_id(2) == 0) {
2442
+ dst[row] = tmp;
2443
+ }
2444
+ }
2445
+
2446
+ template <int qk, int qi, typename block_q_t, int vdr, vec_dot_q_sycl_t vec_dot_q_sycl>
2447
+ static void launch_mul_mat_vec_q_moe(
2448
+ const void * vx_base, const void * vy, const int32_t * ids_dev,
2449
+ float * dst_base, const int ncols, const int nrows, const int n_experts_used,
2450
+ const size_t expert_weight_stride, const size_t dst_row_stride,
2451
+ const size_t src1_row_stride,
2452
+ dpct::queue_ptr stream) {
2453
+ const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
2454
+ const sycl::range<3> block_nums(1, (unsigned) n_experts_used, (unsigned) block_num_y);
2455
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
2456
+ stream->submit([&](sycl::handler & cgh) {
2457
+ cgh.parallel_for(
2458
+ sycl::nd_range<3>(block_nums * block_dims, block_dims),
2459
+ [=](sycl::nd_item<3> item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
2460
+ mul_mat_vec_q_moe<qk, qi, block_q_t, vdr, vec_dot_q_sycl>(
2461
+ vx_base, vy, dst_base, ids_dev, ncols, nrows,
2462
+ expert_weight_stride, dst_row_stride, src1_row_stride, item);
2463
+ });
2464
+ });
2465
+ }
2466
+
2467
+ bool ggml_sycl_mul_mat_vec_q_id(
2468
+ enum ggml_type src0_type,
2469
+ const void * vx_base,
2470
+ const void * vy,
2471
+ const int32_t * ids_dev,
2472
+ float * dst_base,
2473
+ int ncols,
2474
+ int nrows,
2475
+ int n_experts_used,
2476
+ size_t expert_weight_stride,
2477
+ size_t dst_row_stride,
2478
+ size_t src1_row_stride,
2479
+ dpct::queue_ptr stream) {
2480
+ switch (src0_type) {
2481
+ case GGML_TYPE_Q4_0:
2482
+ launch_mul_mat_vec_q_moe<QK4_0, QI4_0, block_q4_0, VDR_Q4_0_Q8_1_MMVQ, vec_dot_q4_0_q8_1>(
2483
+ vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
2484
+ expert_weight_stride, dst_row_stride, src1_row_stride, stream);
2485
+ return true;
2486
+ case GGML_TYPE_Q4_1:
2487
+ launch_mul_mat_vec_q_moe<QK4_1, QI4_1, block_q4_1, VDR_Q4_1_Q8_1_MMVQ, vec_dot_q4_1_q8_1>(
2488
+ vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
2489
+ expert_weight_stride, dst_row_stride, src1_row_stride, stream);
2490
+ return true;
2491
+ case GGML_TYPE_Q5_0:
2492
+ launch_mul_mat_vec_q_moe<QK5_0, QI5_0, block_q5_0, VDR_Q5_0_Q8_1_MMVQ, vec_dot_q5_0_q8_1>(
2493
+ vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
2494
+ expert_weight_stride, dst_row_stride, src1_row_stride, stream);
2495
+ return true;
2496
+ case GGML_TYPE_Q5_1:
2497
+ launch_mul_mat_vec_q_moe<QK5_1, QI5_1, block_q5_1, VDR_Q5_1_Q8_1_MMVQ, vec_dot_q5_1_q8_1>(
2498
+ vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
2499
+ expert_weight_stride, dst_row_stride, src1_row_stride, stream);
2500
+ return true;
2501
+ case GGML_TYPE_Q8_0:
2502
+ launch_mul_mat_vec_q_moe<QK8_0, QI8_0, block_q8_0, VDR_Q8_0_Q8_1_MMVQ, vec_dot_q8_0_q8_1>(
2503
+ vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
2504
+ expert_weight_stride, dst_row_stride, src1_row_stride, stream);
2505
+ return true;
2506
+ case GGML_TYPE_Q2_K:
2507
+ launch_mul_mat_vec_q_moe<QK_K, QI2_K, block_q2_K, VDR_Q2_K_Q8_1_MMVQ, vec_dot_q2_K_q8_1>(
2508
+ vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
2509
+ expert_weight_stride, dst_row_stride, src1_row_stride, stream);
2510
+ return true;
2511
+ case GGML_TYPE_Q3_K:
2512
+ launch_mul_mat_vec_q_moe<QK_K, QI3_K, block_q3_K, VDR_Q3_K_Q8_1_MMVQ, vec_dot_q3_K_q8_1>(
2513
+ vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
2514
+ expert_weight_stride, dst_row_stride, src1_row_stride, stream);
2515
+ return true;
2516
+ case GGML_TYPE_Q4_K:
2517
+ launch_mul_mat_vec_q_moe<QK_K, QI4_K, block_q4_K, VDR_Q4_K_Q8_1_MMVQ, vec_dot_q4_K_q8_1>(
2518
+ vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
2519
+ expert_weight_stride, dst_row_stride, src1_row_stride, stream);
2520
+ return true;
2521
+ case GGML_TYPE_Q5_K:
2522
+ launch_mul_mat_vec_q_moe<QK_K, QI5_K, block_q5_K, VDR_Q5_K_Q8_1_MMVQ, vec_dot_q5_K_q8_1>(
2523
+ vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
2524
+ expert_weight_stride, dst_row_stride, src1_row_stride, stream);
2525
+ return true;
2526
+ case GGML_TYPE_Q6_K:
2527
+ launch_mul_mat_vec_q_moe<QK_K, QI6_K, block_q6_K, VDR_Q6_K_Q8_1_MMVQ, vec_dot_q6_K_q8_1>(
2528
+ vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
2529
+ expert_weight_stride, dst_row_stride, src1_row_stride, stream);
2530
+ return true;
2531
+ case GGML_TYPE_MXFP4:
2532
+ launch_mul_mat_vec_q_moe<QK_MXFP4, QI_MXFP4, block_mxfp4, VDR_MXFP4_Q8_1_MMVQ, vec_dot_mxfp4_q8_1>(
2533
+ vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
2534
+ expert_weight_stride, dst_row_stride, src1_row_stride, stream);
2535
+ return true;
2536
+ case GGML_TYPE_NVFP4:
2537
+ launch_mul_mat_vec_q_moe<QK_NVFP4, QI_NVFP4, block_nvfp4, VDR_NVFP4_Q8_1_MMVQ, vec_dot_nvfp4_q8_1>(
2538
+ vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
2539
+ expert_weight_stride, dst_row_stride, src1_row_stride, stream);
2540
+ return true;
2541
+ default:
2542
+ return false;
2543
+ }
2544
+ }
2545
+
2546
+ // Reorder (SoA) MoE expert GEMV: MoE expert/row/lane indexing (from mul_mat_vec_q_moe) with the
2547
+ // dense-reorder per-block reads (from mul_mat_vec_q_reorder). Each expert slice in vx_base is a
2548
+ // self-contained SoA, so nblocks = nrows*(ncols/qk) per expert and the constant expert stride holds.
2549
+ template <typename reorder_vec_dot_q_sycl>
2550
+ static void mul_mat_vec_q_moe_reorder(
2551
+ const void * __restrict__ vx_base, const void * __restrict__ vy_base,
2552
+ float * __restrict__ dst_base, const int32_t * __restrict__ ids_dev,
2553
+ const int ncols, const int nrows,
2554
+ const size_t expert_weight_stride, const size_t dst_row_stride,
2555
+ const size_t src1_row_stride,
2556
+ const sycl::nd_item<3> & item_ct1) {
2557
+ using block_type = ggml_sycl_reordered::block_q_t<reorder_vec_dot_q_sycl::gtype>;
2558
+ using block_traits = typename block_type::traits;
2559
+
2560
+ const int expert_idx = item_ct1.get_group(1);
2561
+ const int i02 = ids_dev[expert_idx];
2562
+
2563
+ const char * vx = (const char *) vx_base + (size_t) i02 * expert_weight_stride;
2564
+ const char * vy = (const char *) vy_base + (size_t) expert_idx * src1_row_stride;
2565
+ float * dst = (float *) ((char *) dst_base + (size_t) expert_idx * dst_row_stride);
2566
+
2567
+ const int row = item_ct1.get_group(2) * item_ct1.get_local_range(1) + item_ct1.get_local_id(1);
2568
+ if (row >= nrows) {
2569
+ return;
2570
+ }
2571
+
2572
+ const auto sg = item_ct1.get_sub_group();
2573
+
2574
+ const int blocks_per_row = ncols / block_traits::qk;
2575
+ constexpr int blocks_per_subgroup = ceil_div(block_traits::vdr_mmvq * WARP_SIZE, block_traits::qi);
2576
+ constexpr int block_elements_per_subgroup = block_traits::qi / block_traits::vdr_mmvq;
2577
+ const int nblocks = nrows * (ncols / block_traits::qk);
2578
+
2579
+ static_assert(blocks_per_subgroup > 0);
2580
+ static_assert(block_elements_per_subgroup > 0);
2581
+
2582
+ float partial_sum = 0.0f;
2583
+ for (int i = sg.get_local_linear_id() / block_elements_per_subgroup; i < blocks_per_row; i += blocks_per_subgroup) {
2584
+ const int ibx = row * blocks_per_row + i;
2585
+
2586
+ const auto bx_offset = block_type::get_block_offset(ibx, nblocks);
2587
+ const auto d_offset = block_type::get_d_offset(nrows, ncols, ibx);
2588
+
2589
+ const int iby = i * block_type::block_to_q8_1_ratio();
2590
+ const int8_t * q8_1_quant_ptr = (const int8_t *) vy + iby * QK8_1;
2591
+ const sycl::half2 * q8_1_ds_ptr = (const sycl::half2 *) ((const char *) vy + ncols + iby * sizeof(sycl::half2));
2592
+
2593
+ #pragma unroll
2594
+ for (int elem = 0; elem < block_elements_per_subgroup; elem += WARP_SIZE) {
2595
+ const int iqs = elem + block_traits::vdr_mmvq * (sg.get_local_linear_id() % block_elements_per_subgroup);
2596
+ partial_sum += reorder_vec_dot_q_sycl()(vx, bx_offset, d_offset, q8_1_quant_ptr, q8_1_ds_ptr, iqs);
2597
+ }
2598
+ }
2599
+
2600
+ auto sum = sycl::reduce_over_group(sg, partial_sum, std::plus<>());
2601
+ if (sg.leader()) {
2602
+ dst[row] = sum;
2603
+ }
2604
+ }
2605
+
2606
+ template <typename reorder_vec_dot_q_sycl>
2607
+ static void launch_mul_mat_vec_q_moe_reorder(
2608
+ const void * vx_base, const void * vy, const int32_t * ids_dev,
2609
+ float * dst_base, const int ncols, const int nrows, const int n_experts_used,
2610
+ const size_t expert_weight_stride, const size_t dst_row_stride,
2611
+ const size_t src1_row_stride,
2612
+ dpct::queue_ptr stream) {
2613
+ const int block_num_y = (nrows + GGML_SYCL_MMV_Y - 1) / GGML_SYCL_MMV_Y;
2614
+ const sycl::range<3> block_nums(1, (unsigned) n_experts_used, (unsigned) block_num_y);
2615
+ const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, WARP_SIZE);
2616
+ stream->submit([&](sycl::handler & cgh) {
2617
+ cgh.parallel_for(
2618
+ sycl::nd_range<3>(block_nums * block_dims, block_dims),
2619
+ [=](sycl::nd_item<3> item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
2620
+ mul_mat_vec_q_moe_reorder<reorder_vec_dot_q_sycl>(
2621
+ vx_base, vy, dst_base, ids_dev, ncols, nrows,
2622
+ expert_weight_stride, dst_row_stride, src1_row_stride, item);
2623
+ });
2624
+ });
2625
+ }
2626
+
2627
+ bool ggml_sycl_mul_mat_vec_q_id_reorder(
2628
+ enum ggml_type src0_type,
2629
+ const void * vx_base,
2630
+ const void * vy,
2631
+ const int32_t * ids_dev,
2632
+ float * dst_base,
2633
+ int ncols,
2634
+ int nrows,
2635
+ int n_experts_used,
2636
+ size_t expert_weight_stride,
2637
+ size_t dst_row_stride,
2638
+ size_t src1_row_stride,
2639
+ dpct::queue_ptr stream) {
2640
+ switch (src0_type) {
2641
+ case GGML_TYPE_Q4_K:
2642
+ launch_mul_mat_vec_q_moe_reorder<reorder_vec_dot_q_sycl<GGML_TYPE_Q4_K>>(
2643
+ vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
2644
+ expert_weight_stride, dst_row_stride, src1_row_stride, stream);
2645
+ return true;
2646
+ case GGML_TYPE_Q5_K:
2647
+ launch_mul_mat_vec_q_moe_reorder<reorder_vec_dot_q_sycl<GGML_TYPE_Q5_K>>(
2648
+ vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
2649
+ expert_weight_stride, dst_row_stride, src1_row_stride, stream);
2650
+ return true;
2651
+ case GGML_TYPE_Q6_K:
2652
+ launch_mul_mat_vec_q_moe_reorder<reorder_vec_dot_q_sycl<GGML_TYPE_Q6_K>>(
2653
+ vx_base, vy, ids_dev, dst_base, ncols, nrows, n_experts_used,
2654
+ expert_weight_stride, dst_row_stride, src1_row_stride, stream);
2655
+ return true;
2656
+ default:
2657
+ return false;
2658
+ }
2659
+ }