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
@@ -0,0 +1,1553 @@
1
+ #ifdef U32_DEQUANT_HELPERS
2
+ #define SRC0_TYPE u32
3
+
4
+ fn byte_of(v: u32, b: u32) -> u32 {
5
+ return (v >> (b * 8u)) & 0xFFu;
6
+ }
7
+
8
+ fn sbyte_of(v: u32, b: u32) -> i32 {
9
+ let raw = i32((v >> (b * 8u)) & 0xFFu);
10
+ return select(raw, raw - 256, raw >= 128);
11
+ }
12
+ #endif
13
+
14
+ #ifdef VEC
15
+ #define VEC_SIZE 4u
16
+ #define SRC0_TYPE vec4<SRC0_INNER_TYPE>
17
+ #define SRC1_TYPE vec4<SRC1_INNER_TYPE>
18
+
19
+ fn inner_dot(src0_val: SRC0_TYPE, src1_val: SRC1_TYPE) -> f32 {
20
+ return f32(dot(SRC1_TYPE(src0_val), src1_val));
21
+ }
22
+ #endif
23
+
24
+ #ifdef SCALAR
25
+ #define VEC_SIZE 1u
26
+ #define SRC0_TYPE SRC0_INNER_TYPE
27
+ #define SRC1_TYPE SRC1_INNER_TYPE
28
+
29
+ fn inner_dot(src0_val: SRC0_TYPE, src1_val: SRC1_TYPE) -> f32 {
30
+ return f32(src0_val) * f32(src1_val);
31
+ }
32
+ #endif
33
+
34
+ #ifdef MUL_ACC_FLOAT
35
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
36
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
37
+
38
+ let k_vec = params.k / VEC_SIZE;
39
+ let src1_idx_base_vec = src1_idx_base / VEC_SIZE;
40
+
41
+ // Each thread walks K, loads from the vector, and updates
42
+ // a small block of output rows held in registers.
43
+ for (var k = thread_id; k < k_vec; k += WG_SIZE) {
44
+ var x_vals: array<SRC1_TYPE, NUM_COLS>;
45
+ for (var col = 0u;col < NUM_COLS;col += 1) {
46
+ x_vals[col] = src1[src1_idx_base_vec + col * (params.stride_11 / VEC_SIZE) + k];
47
+ }
48
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
49
+ let output_row = row_base + row;
50
+ if (output_row < params.m) {
51
+ let src0_idx = (src0_batch_offset + output_row * params.stride_01) / VEC_SIZE + k;
52
+ let w = src0[src0_idx];
53
+ for (var col = 0u;col < NUM_COLS;col += 1) {
54
+ acc[col][row] += inner_dot(w, x_vals[col]);
55
+ }
56
+ }
57
+ }
58
+ }
59
+
60
+ return acc;
61
+ }
62
+ #endif
63
+
64
+ #ifdef MUL_ACC_Q1_0
65
+ #define BLOCK_SIZE 128
66
+ #define BLOCK_SIZE_BYTES 18
67
+ #define THREADS_PER_BLOCK 16
68
+ #define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
69
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
70
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
71
+
72
+ let num_blocks = params.k / BLOCK_SIZE;
73
+ let thread_within_block = thread_id % THREADS_PER_BLOCK;
74
+ for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
75
+ let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * ELEMS_PER_THREAD;
76
+ var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
77
+ for (var col = 0u; col < NUM_COLS;col += 1) {
78
+ for (var i = 0u; i < ELEMS_PER_THREAD; i++) {
79
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
80
+ }
81
+ }
82
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
83
+ let output_row = row_base + row;
84
+ if (output_row < params.m) {
85
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
86
+ let d = f32(load_f16_at_src0(block_byte_base));
87
+ let q_byte = load_u32_at_src0(block_byte_base + 2u + thread_within_block) & 0xFFu;
88
+ for (var col = 0u;col < NUM_COLS;col += 1) {
89
+ var row_sum = 0.0;
90
+ for (var bit = 0u; bit < 8u; bit++) {
91
+ let w = select(-d, d, ((q_byte >> bit) & 1u) != 0u);
92
+ row_sum += w * x_block[col][bit];
93
+ }
94
+ acc[col][row] += row_sum;
95
+ }
96
+ }
97
+ }
98
+ }
99
+
100
+ return acc;
101
+ }
102
+ #endif
103
+
104
+ #ifdef MUL_ACC_Q4_0
105
+ #define BLOCK_SIZE 32
106
+ #define BLOCK_SIZE_BYTES 18
107
+ #define THREADS_PER_BLOCK 4
108
+ #define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
109
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
110
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
111
+
112
+ let num_blocks = params.k / BLOCK_SIZE;
113
+ let thread_within_block = thread_id % 4;
114
+ for (var block = thread_id/THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE/THREADS_PER_BLOCK) {
115
+ let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * 4;
116
+ var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
117
+ for (var col = 0u; col < NUM_COLS;col += 1) {
118
+ for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
119
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
120
+ x_block[col][i + 4] = f32(src1[x_base + col * params.stride_11 + i + 16]);
121
+ }
122
+ }
123
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
124
+ let output_row = row_base + row;
125
+ if (output_row < params.m) {
126
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
127
+ let d = f32(load_f16_at_src0(block_byte_base));
128
+ let q_packed = load_u32_at_src0(block_byte_base + 2u + 4u * thread_within_block);
129
+ for (var col = 0u;col < NUM_COLS;col += 1) {
130
+ var row_sum = 0.0;
131
+ for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
132
+ let q_byte = get_byte(q_packed, byte_idx);
133
+ let q_lo = (f32(q_byte & 0xFu) - 8.0) * d;
134
+ let q_hi = (f32((q_byte >> 4u) & 0xFu) - 8.0) * d;
135
+ row_sum += q_lo * x_block[col][byte_idx];
136
+ row_sum += q_hi * x_block[col][byte_idx + 4u];
137
+ }
138
+ acc[col][row] += row_sum;
139
+ }
140
+ }
141
+ }
142
+ }
143
+
144
+ return acc;
145
+ }
146
+ #endif
147
+
148
+ #ifdef MUL_ACC_Q4_1
149
+ #define BLOCK_SIZE 32
150
+ #define BLOCK_SIZE_BYTES 20
151
+ #define THREADS_PER_BLOCK 4
152
+ #define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
153
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
154
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
155
+
156
+ let num_blocks = params.k / BLOCK_SIZE;
157
+ let thread_within_block = thread_id % THREADS_PER_BLOCK;
158
+ for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
159
+ let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * 4;
160
+ var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
161
+ for (var col = 0u; col < NUM_COLS;col += 1) {
162
+ for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
163
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
164
+ x_block[col][i + 4] = f32(src1[x_base + col * params.stride_11 + i + 16]);
165
+ }
166
+ }
167
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
168
+ let output_row = row_base + row;
169
+ if (output_row < params.m) {
170
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
171
+ let d = f32(load_f16_at_src0(block_byte_base));
172
+ let m = f32(load_f16_at_src0(block_byte_base + 2u));
173
+ let q_packed = load_u32_at_src0(block_byte_base + 4u + 4u * thread_within_block);
174
+ for (var col = 0u;col < NUM_COLS;col += 1) {
175
+ var row_sum = 0.0;
176
+ for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
177
+ let q_byte = get_byte(q_packed, byte_idx);
178
+ let q_lo = f32(q_byte & 0xFu) * d + m;
179
+ let q_hi = f32((q_byte >> 4u) & 0xFu) * d + m;
180
+ row_sum += q_lo * x_block[col][byte_idx];
181
+ row_sum += q_hi * x_block[col][byte_idx + 4u];
182
+ }
183
+ acc[col][row] += row_sum;
184
+ }
185
+ }
186
+ }
187
+ }
188
+
189
+ return acc;
190
+ }
191
+ #endif
192
+
193
+ #ifdef MUL_ACC_Q5_0
194
+ #define BLOCK_SIZE 32
195
+ #define BLOCK_SIZE_BYTES 22
196
+ #define THREADS_PER_BLOCK 4
197
+ #define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
198
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
199
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
200
+
201
+ let num_blocks = params.k / BLOCK_SIZE;
202
+ let thread_within_block = thread_id % THREADS_PER_BLOCK;
203
+ for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
204
+ let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * 4;
205
+ var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
206
+ for (var col = 0u; col < NUM_COLS;col += 1) {
207
+ for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
208
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
209
+ x_block[col][i + 4] = f32(src1[x_base + col * params.stride_11 + i + 16]);
210
+ }
211
+ }
212
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
213
+ let output_row = row_base + row;
214
+ if (output_row < params.m) {
215
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
216
+ let d = f32(load_f16_at_src0(block_byte_base));
217
+ let qh_packed = load_u32_at_src0(block_byte_base + 2u);
218
+ let q_packed = load_u32_at_src0(block_byte_base + 6u + 4u * thread_within_block);
219
+ let qh_shift = thread_within_block * 4u;
220
+ for (var col = 0u;col < NUM_COLS;col += 1) {
221
+ var row_sum = 0.0;
222
+ for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
223
+ let q_byte = get_byte(q_packed, byte_idx);
224
+ let qh_lo = ((qh_packed >> (qh_shift + byte_idx)) << 4u) & 0x10u;
225
+ let qh_hi = (qh_packed >> (qh_shift + byte_idx + 12u)) & 0x10u;
226
+ let q_lo = (f32((q_byte & 0xFu) | qh_lo) - 16.0) * d;
227
+ let q_hi = (f32(((q_byte >> 4u) & 0xFu) | qh_hi) - 16.0) * d;
228
+ row_sum += q_lo * x_block[col][byte_idx];
229
+ row_sum += q_hi * x_block[col][byte_idx + 4u];
230
+ }
231
+ acc[col][row] += row_sum;
232
+ }
233
+ }
234
+ }
235
+ }
236
+
237
+ return acc;
238
+ }
239
+ #endif
240
+
241
+ #ifdef MUL_ACC_Q5_1
242
+ #define BLOCK_SIZE 32
243
+ #define BLOCK_SIZE_BYTES 24
244
+ #define THREADS_PER_BLOCK 4
245
+ #define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
246
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
247
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
248
+
249
+ let num_blocks = params.k / BLOCK_SIZE;
250
+ let thread_within_block = thread_id % THREADS_PER_BLOCK;
251
+ for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
252
+ let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * 4;
253
+ var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
254
+ for (var col = 0u; col < NUM_COLS;col += 1) {
255
+ for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
256
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
257
+ x_block[col][i + 4] = f32(src1[x_base + col * params.stride_11 + i + 16]);
258
+ }
259
+ }
260
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
261
+ let output_row = row_base + row;
262
+ if (output_row < params.m) {
263
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
264
+ let d = f32(load_f16_at_src0(block_byte_base));
265
+ let m = f32(load_f16_at_src0(block_byte_base + 2u));
266
+ let qh_packed = load_u32_at_src0(block_byte_base + 4u);
267
+ let q_packed = load_u32_at_src0(block_byte_base + 8u + 4u * thread_within_block);
268
+ let qh_shift = thread_within_block * 4u;
269
+ for (var col = 0u;col < NUM_COLS;col += 1) {
270
+ var row_sum = 0.0;
271
+ for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
272
+ let q_byte = get_byte(q_packed, byte_idx);
273
+ let qh_lo = ((qh_packed >> (qh_shift + byte_idx)) << 4u) & 0x10u;
274
+ let qh_hi = (qh_packed >> (qh_shift + byte_idx + 12u)) & 0x10u;
275
+ let q_lo = f32((q_byte & 0xFu) | qh_lo) * d + m;
276
+ let q_hi = f32(((q_byte >> 4u) & 0xFu) | qh_hi) * d + m;
277
+ row_sum += q_lo * x_block[col][byte_idx];
278
+ row_sum += q_hi * x_block[col][byte_idx + 4u];
279
+ }
280
+ acc[col][row] += row_sum;
281
+ }
282
+ }
283
+ }
284
+ }
285
+
286
+ return acc;
287
+ }
288
+ #endif
289
+
290
+ #ifdef MUL_ACC_Q8_0
291
+ #define BLOCK_SIZE 32
292
+ #define BLOCK_SIZE_BYTES 34
293
+ #define THREADS_PER_BLOCK 4
294
+ #define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
295
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
296
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
297
+
298
+ let num_blocks = params.k / BLOCK_SIZE;
299
+ let thread_within_block = thread_id % THREADS_PER_BLOCK;
300
+ for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
301
+ let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * ELEMS_PER_THREAD;
302
+ var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
303
+ for (var col = 0u; col < NUM_COLS;col += 1) {
304
+ for (var i = 0u; i < ELEMS_PER_THREAD; i++) {
305
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
306
+ }
307
+ }
308
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
309
+ let output_row = row_base + row;
310
+ if (output_row < params.m) {
311
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
312
+ let d = f32(load_f16_at_src0(block_byte_base));
313
+ var q_packed: array<u32, ELEMS_PER_THREAD / 4u>;
314
+ for (var packed_idx = 0u; packed_idx < ELEMS_PER_THREAD / 4u; packed_idx++) {
315
+ q_packed[packed_idx] = load_u32_at_src0(block_byte_base + 2u + 4u * (thread_within_block * 2u + packed_idx));
316
+ }
317
+ for (var col = 0u;col < NUM_COLS;col += 1) {
318
+ var row_sum = 0.0;
319
+ for (var packed_idx = 0u; packed_idx < ELEMS_PER_THREAD / 4u; packed_idx++) {
320
+ for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
321
+ let q_val = f32(get_byte_i32(q_packed[packed_idx], byte_idx)) * d;
322
+ row_sum += q_val * x_block[col][packed_idx * 4u + byte_idx];
323
+ }
324
+ }
325
+ acc[col][row] += row_sum;
326
+ }
327
+ }
328
+ }
329
+ }
330
+
331
+ return acc;
332
+ }
333
+ #endif
334
+
335
+ #ifdef MUL_ACC_Q8_1
336
+ #define BLOCK_SIZE 32
337
+ #define BLOCK_SIZE_BYTES 36
338
+ #define THREADS_PER_BLOCK 4
339
+ #define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
340
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
341
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
342
+
343
+ let num_blocks = params.k / BLOCK_SIZE;
344
+ let thread_within_block = thread_id % THREADS_PER_BLOCK;
345
+ for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
346
+ let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * ELEMS_PER_THREAD;
347
+ var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
348
+ for (var col = 0u; col < NUM_COLS;col += 1) {
349
+ for (var i = 0u; i < ELEMS_PER_THREAD; i++) {
350
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
351
+ }
352
+ }
353
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
354
+ let output_row = row_base + row;
355
+ if (output_row < params.m) {
356
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
357
+ let d = f32(load_f16_at_src0(block_byte_base));
358
+ let m = f32(load_f16_at_src0(block_byte_base + 2u));
359
+ var q_packed: array<u32, ELEMS_PER_THREAD / 4u>;
360
+ for (var packed_idx = 0u; packed_idx < ELEMS_PER_THREAD / 4u; packed_idx++) {
361
+ q_packed[packed_idx] = load_u32_at_src0(block_byte_base + 4u + 4u * (thread_within_block * 2u + packed_idx));
362
+ }
363
+ for (var col = 0u;col < NUM_COLS;col += 1) {
364
+ var row_sum = 0.0;
365
+ for (var packed_idx = 0u; packed_idx < ELEMS_PER_THREAD / 4u; packed_idx++) {
366
+ for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
367
+ let q_val = f32(get_byte_i32(q_packed[packed_idx], byte_idx)) * d + m;
368
+ row_sum += q_val * x_block[col][packed_idx * 4u + byte_idx];
369
+ }
370
+ }
371
+ acc[col][row] += row_sum;
372
+ }
373
+ }
374
+ }
375
+ }
376
+
377
+ return acc;
378
+ }
379
+ #endif
380
+
381
+ #ifdef MUL_ACC_Q2_K
382
+ #define BLOCK_SIZE 256
383
+ #define BLOCK_SIZE_BYTES 84
384
+ #define THREADS_PER_BLOCK 16
385
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
386
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
387
+
388
+ let tid = thread_id % THREADS_PER_BLOCK;
389
+ let block_group = thread_id / THREADS_PER_BLOCK;
390
+ let num_block_groups: u32 = WG_SIZE / THREADS_PER_BLOCK;
391
+
392
+ let lane = tid / 2u;
393
+ let phase = tid % 2u;
394
+ let iq = lane / 4u;
395
+ let ir = lane % 4u;
396
+ let is = ir / 2u;
397
+
398
+ let y_offset = 128u * iq + 8u * ir + 4u * phase;
399
+ let sc0_byte = 8u * iq + is;
400
+ let sc2_byte = 8u * iq + is + 2u;
401
+ let sc4_byte = 8u * iq + is + 4u;
402
+ let sc6_byte = 8u * iq + is + 6u;
403
+ let qs_byte = 16u + (16u * iq + 4u * ir) * 2u + 4u * phase;
404
+
405
+ let num_blocks = params.k / BLOCK_SIZE;
406
+
407
+ for (var block = block_group; block < num_blocks; block += num_block_groups) {
408
+ let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
409
+ var x_block: array<array<f32, 16>, NUM_COLS>;
410
+ for (var col = 0u; col < NUM_COLS;col += 1) {
411
+ for (var i = 0u; i < 4u; i++) {
412
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
413
+ x_block[col][i + 4u] = f32(src1[x_base + col * params.stride_11 + 32u + i]);
414
+ x_block[col][i + 8u] = f32(src1[x_base + col * params.stride_11 + 64u + i]);
415
+ x_block[col][i + 12u] = f32(src1[x_base + col * params.stride_11 + 96u + i]);
416
+ }
417
+ }
418
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
419
+ let output_row = row_base + row;
420
+ if (output_row < params.m) {
421
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
422
+
423
+ let dall = f32(load_f16_at_src0(block_byte_base + 80u));
424
+ let dmin = f32(load_f16_at_src0(block_byte_base + 82u)) * (1.0 / 16.0);
425
+
426
+ let sc0 = byte_of(load_u32_at_src0_aligned(block_byte_base + sc0_byte), sc0_byte & 3u);
427
+ let sc2 = byte_of(load_u32_at_src0_aligned(block_byte_base + sc2_byte), sc2_byte & 3u);
428
+ let sc4 = byte_of(load_u32_at_src0_aligned(block_byte_base + sc4_byte), sc4_byte & 3u);
429
+ let sc6 = byte_of(load_u32_at_src0_aligned(block_byte_base + sc6_byte), sc6_byte & 3u);
430
+
431
+ let q_u32 = load_u32_at_src0_aligned(block_byte_base + qs_byte);
432
+ let qs0 = q_u32 & 0xFFFFu;
433
+ let qs1 = q_u32 >> 16u;
434
+
435
+ for (var col = 0u;col < NUM_COLS;col += 1) {
436
+ var sumy = vec4<f32>(0.0, 0.0, 0.0, 0.0);
437
+ var acc1 = vec4<f32>(0.0, 0.0, 0.0, 0.0);
438
+ var acc2 = vec4<f32>(0.0, 0.0, 0.0, 0.0);
439
+
440
+ sumy[0] = x_block[col][0] + x_block[col][1] + x_block[col][2] + x_block[col][3];
441
+ sumy[1] = x_block[col][4] + x_block[col][5] + x_block[col][6] + x_block[col][7];
442
+ sumy[2] = x_block[col][8] + x_block[col][9] + x_block[col][10] + x_block[col][11];
443
+ sumy[3] = x_block[col][12] + x_block[col][13] + x_block[col][14] + x_block[col][15];
444
+
445
+ acc1[0] = x_block[col][0] * f32(qs0 & 0x0003u) + x_block[col][2] * f32(qs1 & 0x0003u);
446
+ acc2[0] = x_block[col][1] * f32(qs0 & 0x0300u) + x_block[col][3] * f32(qs1 & 0x0300u);
447
+ acc1[1] = x_block[col][4] * f32(qs0 & 0x000Cu) + x_block[col][6] * f32(qs1 & 0x000Cu);
448
+ acc2[1] = x_block[col][5] * f32(qs0 & 0x0C00u) + x_block[col][7] * f32(qs1 & 0x0C00u);
449
+ acc1[2] = x_block[col][8] * f32(qs0 & 0x0030u) + x_block[col][10] * f32(qs1 & 0x0030u);
450
+ acc2[2] = x_block[col][9] * f32(qs0 & 0x3000u) + x_block[col][11] * f32(qs1 & 0x3000u);
451
+ acc1[3] = x_block[col][12] * f32(qs0 & 0x00C0u) + x_block[col][14] * f32(qs1 & 0x00C0u);
452
+ acc2[3] = x_block[col][13] * f32(qs0 & 0xC000u) + x_block[col][15] * f32(qs1 & 0xC000u);
453
+
454
+ acc[col][row] += dall * ((acc1[0] + (1.0/256.0) * acc2[0]) * f32(sc0 & 0xFu) +
455
+ (acc1[1] + (1.0/256.0) * acc2[1]) * f32(sc2 & 0xFu) / 4.0 +
456
+ (acc1[2] + (1.0/256.0) * acc2[2]) * f32(sc4 & 0xFu) / 16.0 +
457
+ (acc1[3] + (1.0/256.0) * acc2[3]) * f32(sc6 & 0xFu) / 64.0)
458
+ - dmin * (sumy[0] * f32(sc0 & 0xF0u) + sumy[1] * f32(sc2 & 0xF0u) +
459
+ sumy[2] * f32(sc4 & 0xF0u) + sumy[3] * f32(sc6 & 0xF0u));
460
+ }
461
+ }
462
+ }
463
+ }
464
+
465
+ return acc;
466
+ }
467
+ #endif
468
+
469
+ #ifdef MUL_ACC_Q3_K
470
+ #define BLOCK_SIZE 256
471
+ #define BLOCK_SIZE_BYTES 110
472
+ #define THREADS_PER_BLOCK 16
473
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
474
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
475
+
476
+ let tid = thread_id % THREADS_PER_BLOCK;
477
+ let block_group = thread_id / THREADS_PER_BLOCK;
478
+ let num_block_groups: u32 = WG_SIZE / THREADS_PER_BLOCK;
479
+
480
+ let lane = tid / 2u;
481
+ let phase = tid % 2u;
482
+ let ip = lane / 4u;
483
+ let il = 2u * ((lane % 4u) / 2u);
484
+ let ir = lane % 2u;
485
+ let l0 = 8u * ir;
486
+
487
+ let q_byte = 32u + 32u * ip + l0 + 16u * phase;
488
+ let h_byte = l0 + 16u * phase;
489
+ let y_offset = 128u * ip + 32u * il + l0 + 16u * phase;
490
+
491
+ let s_shift1 = 4u * ip;
492
+ let s_shift2 = s_shift1 + il;
493
+
494
+ let v1 = select(64.0, 4.0, il == 0u);
495
+ let v2 = 4.0 * v1;
496
+ let shift = 2u * il;
497
+
498
+ var qm0: u32; var qm1: u32; var qm2: u32; var qm3: u32;
499
+ if (il == 0u) {
500
+ qm0 = 0x0003u; qm1 = 0x0300u; qm2 = 0x000Cu; qm3 = 0x0C00u;
501
+ } else {
502
+ qm0 = 0x0030u; qm1 = 0x3000u; qm2 = 0x00C0u; qm3 = 0xC000u;
503
+ }
504
+
505
+ let mm_idx = 2u * ip + il / 2u;
506
+ var hm0: u32; var hm1: u32; var hm2: u32; var hm3: u32;
507
+ switch (mm_idx) {
508
+ case 0u: { hm0=0x0001u; hm1=0x0100u; hm2=0x0002u; hm3=0x0200u; }
509
+ case 1u: { hm0=0x0004u; hm1=0x0400u; hm2=0x0008u; hm3=0x0800u; }
510
+ case 2u: { hm0=0x0010u; hm1=0x1000u; hm2=0x0020u; hm3=0x2000u; }
511
+ default: { hm0=0x0040u; hm1=0x4000u; hm2=0x0080u; hm3=0x8000u; }
512
+ }
513
+
514
+ let num_blocks = params.k / BLOCK_SIZE;
515
+
516
+ for (var block = block_group; block < num_blocks; block += num_block_groups) {
517
+ let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
518
+ var x_block: array<array<f32, 16>, NUM_COLS>;
519
+ for (var col = 0u; col < NUM_COLS;col += 1) {
520
+ for (var i = 0u; i < 8u; i++) {
521
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
522
+ x_block[col][i + 8u] = f32(src1[x_base + col * params.stride_11 + 32u + i]);
523
+ }
524
+ }
525
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
526
+ let output_row = row_base + row;
527
+ if (output_row < params.m) {
528
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
529
+
530
+ let d = f32(load_f16_at_src0(block_byte_base + 108u));
531
+ let a_base = 96u;
532
+ let a_il0 = load_u16_at_src0(block_byte_base + a_base + il * 2u);
533
+ let a_il1 = load_u16_at_src0(block_byte_base + a_base + (il + 1u) * 2u);
534
+ let a_4 = load_u16_at_src0(block_byte_base + a_base + 8u);
535
+ let a_5 = load_u16_at_src0(block_byte_base + a_base + 10u);
536
+
537
+ var scales32 = a_4 | (a_5 << 16u);
538
+ let aux32 = ((scales32 >> s_shift2) << 4u) & 0x30303030u;
539
+ scales32 = a_il0 | (a_il1 << 16u);
540
+ scales32 = ((scales32 >> s_shift1) & 0x0F0F0F0Fu) | aux32;
541
+
542
+ let scale0 = f32(i32(byte_of(scales32, phase + 0u)) - 32);
543
+ let scale1 = f32(i32(byte_of(scales32, phase + 2u)) - 32);
544
+
545
+ let q_u32_0 = load_u32_at_src0(block_byte_base + q_byte + 0u);
546
+ let q_u32_1 = load_u32_at_src0(block_byte_base + q_byte + 4u);
547
+ let h_u32_0 = load_u32_at_src0(block_byte_base + h_byte + 0u);
548
+ let h_u32_1 = load_u32_at_src0(block_byte_base + h_byte + 4u);
549
+
550
+ for (var col = 0u;col < NUM_COLS;col += 1) {
551
+ var s1 = 0.0; var s2 = 0.0; var s3 = 0.0;
552
+ var s4 = 0.0; var s5 = 0.0; var s6 = 0.0;
553
+
554
+ for (var l = 0u; l < 8u; l += 2u) {
555
+ let q_u32 = select(q_u32_0, q_u32_1, l >= 4u);
556
+ let qs = select(q_u32 & 0xFFFFu, q_u32 >> 16u, (l & 2u) != 0u);
557
+ let h_u32 = select(h_u32_0, h_u32_1, l >= 4u);
558
+ let hv = select(h_u32 & 0xFFFFu, h_u32 >> 16u, (l & 2u) != 0u);
559
+
560
+ s1 += x_block[col][l + 0u] * f32(qs & qm0);
561
+ s2 += x_block[col][l + 1u] * f32(qs & qm1);
562
+ s3 += select(0.0, x_block[col][l + 0u], (hv & hm0) == 0u) +
563
+ select(0.0, x_block[col][l + 1u], (hv & hm1) == 0u);
564
+ s4 += x_block[col][l + 8u] * f32(qs & qm2);
565
+ s5 += x_block[col][l + 9u] * f32(qs & qm3);
566
+ s6 += select(0.0, x_block[col][l + 8u], (hv & hm2) == 0u) +
567
+ select(0.0, x_block[col][l + 9u], (hv & hm3) == 0u);
568
+ }
569
+
570
+ let d1 = d * (s1 + (1.0/256.0) * s2 - s3 * v1);
571
+ let d2 = d * (s4 + (1.0/256.0) * s5 - s6 * v2);
572
+ acc[col][row] += (d1 * scale0 + 0.25 * d2 * scale1) / f32(1u << shift);
573
+ }
574
+ }
575
+ }
576
+ }
577
+
578
+ return acc;
579
+ }
580
+ #endif
581
+
582
+ #ifdef MUL_ACC_Q4_K
583
+ #define BLOCK_SIZE 256
584
+ #define BLOCK_SIZE_BYTES 144
585
+ #define THREADS_PER_BLOCK 16
586
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
587
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
588
+
589
+ let tid = thread_id % THREADS_PER_BLOCK;
590
+ let block_group = thread_id / THREADS_PER_BLOCK;
591
+ let num_block_groups: u32 = WG_SIZE / THREADS_PER_BLOCK;
592
+
593
+ let il = tid / 4u;
594
+ let ir = tid % 4u;
595
+ let im = il / 2u;
596
+ let in = il % 2u;
597
+ let l0 = 4u * (2u * ir + in);
598
+
599
+ let y_offset = 64u * im + l0;
600
+ let q_offset = 32u * im + l0;
601
+ let sc0_byte = 4u + im * 2u;
602
+ let sc2_byte = 4u + (im + 2u) * 2u;
603
+ let sc4_byte = 4u + (im + 4u) * 2u;
604
+
605
+ let num_blocks = params.k / BLOCK_SIZE;
606
+
607
+ for (var block = block_group; block < num_blocks; block += num_block_groups) {
608
+ let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
609
+ var x_block: array<array<f32, 16>, NUM_COLS>;
610
+ for (var col = 0u; col < NUM_COLS;col += 1) {
611
+ let col_base = x_base + col * params.stride_11;
612
+ for (var i = 0u; i < 4u; i++) {
613
+ x_block[col][i] = f32(src1[col_base + i]);
614
+ x_block[col][i + 4u] = f32(src1[col_base + 32u + i]);
615
+ x_block[col][i + 8u] = f32(src1[col_base + 128u + i]);
616
+ x_block[col][i + 12u] = f32(src1[col_base + 160u + i]);
617
+ }
618
+ }
619
+
620
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
621
+ let output_row = row_base + row;
622
+ if (output_row < params.m) {
623
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
624
+
625
+ let d = f32(load_f16_at_src0(block_byte_base + 0u));
626
+ let dmin = f32(load_f16_at_src0(block_byte_base + 2u));
627
+
628
+ let sc0_u32 = load_u32_at_src0_aligned(block_byte_base + sc0_byte);
629
+ let sc0 = select(sc0_u32 & 0xFFFFu, sc0_u32 >> 16u, (sc0_byte & 2u) != 0u);
630
+ let sc2_u32 = load_u32_at_src0_aligned(block_byte_base + sc2_byte);
631
+ let sc2 = select(sc2_u32 & 0xFFFFu, sc2_u32 >> 16u, (sc2_byte & 2u) != 0u);
632
+ let sc4_u32 = load_u32_at_src0_aligned(block_byte_base + sc4_byte);
633
+ let sc4 = select(sc4_u32 & 0xFFFFu, sc4_u32 >> 16u, (sc4_byte & 2u) != 0u);
634
+
635
+ let sc16_0 = sc0 & 0x3F3Fu;
636
+ let sc16_1 = sc2 & 0x3F3Fu;
637
+ let sc16_2 = (sc4 & 0x0F0Fu) | ((sc0 & 0xC0C0u) >> 2u);
638
+ let sc16_3 = ((sc4 >> 4u) & 0x0F0Fu) | ((sc2 & 0xC0C0u) >> 2u);
639
+
640
+ let scale0 = f32(sc16_0 & 0xFFu);
641
+ let scale1 = f32((sc16_0 >> 8u) & 0xFFu);
642
+ let min0 = f32(sc16_1 & 0xFFu);
643
+ let min1 = f32((sc16_1 >> 8u) & 0xFFu);
644
+ let scale2 = f32(sc16_2 & 0xFFu);
645
+ let scale3 = f32((sc16_2 >> 8u) & 0xFFu);
646
+ let min2 = f32(sc16_3 & 0xFFu);
647
+ let min3 = f32((sc16_3 >> 8u) & 0xFFu);
648
+
649
+ let q1_u32 = load_u32_at_src0_aligned(block_byte_base + 16u + q_offset);
650
+ let q2_u32 = load_u32_at_src0_aligned(block_byte_base + 80u + q_offset);
651
+
652
+ for (var col = 0u;col < NUM_COLS;col += 1) {
653
+ var dot = vec4<f32>(0.0, 0.0, 0.0, 0.0);
654
+ var sumx = vec4<f32>(0.0, 0.0, 0.0, 0.0);
655
+ for (var i = 0u; i < 4u; i++) {
656
+ let q1b = byte_of(q1_u32, i);
657
+ let q2b = byte_of(q2_u32, i);
658
+ dot[0] += x_block[col][i] * f32(q1b & 0x0Fu);
659
+ dot[1] += x_block[col][i + 4u] * f32(q1b >> 4u);
660
+ dot[2] += x_block[col][i + 8u] * f32(q2b & 0x0Fu);
661
+ dot[3] += x_block[col][i + 12u] * f32(q2b >> 4u);
662
+ sumx[0] += x_block[col][i];
663
+ sumx[1] += x_block[col][i + 4u];
664
+ sumx[2] += x_block[col][i + 8u];
665
+ sumx[3] += x_block[col][i + 12u];
666
+ }
667
+
668
+ acc[col][row] += d * (dot[0] * scale0 + dot[1] * scale1 + dot[2] * scale2 + dot[3] * scale3)
669
+ - dmin * (sumx[0] * min0 + sumx[1] * min1 + sumx[2] * min2 + sumx[3] * min3);
670
+ }
671
+ }
672
+ }
673
+ }
674
+
675
+ return acc;
676
+ }
677
+ #endif
678
+
679
+ #ifdef MUL_ACC_Q5_K
680
+ #define BLOCK_SIZE 256
681
+ #define BLOCK_SIZE_BYTES 176
682
+ #define THREADS_PER_BLOCK 16
683
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
684
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
685
+
686
+ let tid = thread_id % THREADS_PER_BLOCK;
687
+ let block_group = thread_id / THREADS_PER_BLOCK;
688
+ let num_block_groups: u32 = WG_SIZE / THREADS_PER_BLOCK;
689
+
690
+ let il = tid / 4u;
691
+ let ir = tid % 4u;
692
+ let im = il / 2u;
693
+ let in = il % 2u;
694
+ let l0 = 4u * (2u * ir + in);
695
+
696
+ let y_offset = 64u * im + l0;
697
+ let q_offset = 48u + 32u * im + l0;
698
+ let qh_offset = 16u + 8u * ir + 4u * in;
699
+ let sc0_byte = 4u + im * 2u;
700
+ let sc2_byte = 4u + (im + 2u) * 2u;
701
+ let sc4_byte = 4u + (im + 4u) * 2u;
702
+
703
+ let hm1 = 1u << (2u * im);
704
+ let hm2 = hm1 << 1u;
705
+ let hm3 = hm1 << 4u;
706
+ let hm4 = hm2 << 4u;
707
+
708
+ let num_blocks = params.k / BLOCK_SIZE;
709
+
710
+ for (var block = block_group; block < num_blocks; block += num_block_groups) {
711
+ let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
712
+ var x_block: array<array<f32, 16>, NUM_COLS>;
713
+ for (var col = 0u; col < NUM_COLS;col += 1) {
714
+ let col_base = x_base + col * params.stride_11;
715
+ for (var i = 0u; i < 4u; i++) {
716
+ x_block[col][i] = f32(src1[col_base + i]);
717
+ x_block[col][i + 4u] = f32(src1[col_base + 32u + i]);
718
+ x_block[col][i + 8u] = f32(src1[col_base + 128u + i]);
719
+ x_block[col][i + 12u] = f32(src1[col_base + 160u + i]);
720
+ }
721
+ }
722
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
723
+ let output_row = row_base + row;
724
+ if (output_row < params.m) {
725
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
726
+
727
+ let d = f32(load_f16_at_src0(block_byte_base + 0u));
728
+ let dmin = f32(load_f16_at_src0(block_byte_base + 2u));
729
+
730
+ let sc0_u32 = load_u32_at_src0_aligned(block_byte_base + sc0_byte);
731
+ let sc0 = select(sc0_u32 & 0xFFFFu, sc0_u32 >> 16u, (sc0_byte & 2u) != 0u);
732
+ let sc2_u32 = load_u32_at_src0_aligned(block_byte_base + sc2_byte);
733
+ let sc2 = select(sc2_u32 & 0xFFFFu, sc2_u32 >> 16u, (sc2_byte & 2u) != 0u);
734
+ let sc4_u32 = load_u32_at_src0_aligned(block_byte_base + sc4_byte);
735
+ let sc4 = select(sc4_u32 & 0xFFFFu, sc4_u32 >> 16u, (sc4_byte & 2u) != 0u);
736
+
737
+ let sc16_0 = sc0 & 0x3F3Fu;
738
+ let sc16_1 = sc2 & 0x3F3Fu;
739
+ let sc16_2 = (sc4 & 0x0F0Fu) | ((sc0 & 0xC0C0u) >> 2u);
740
+ let sc16_3 = ((sc4 >> 4u) & 0x0F0Fu) | ((sc2 & 0xC0C0u) >> 2u);
741
+
742
+ let f0 = f32(sc16_0 & 0xFFu);
743
+ let f1 = f32((sc16_0 >> 8u) & 0xFFu);
744
+ let m0 = f32(sc16_1 & 0xFFu);
745
+ let m1 = f32((sc16_1 >> 8u) & 0xFFu);
746
+ let f4 = f32(sc16_2 & 0xFFu);
747
+ let f5 = f32((sc16_2 >> 8u) & 0xFFu);
748
+ let m4 = f32(sc16_3 & 0xFFu);
749
+ let m5 = f32((sc16_3 >> 8u) & 0xFFu);
750
+
751
+ let q1_u32 = load_u32_at_src0_aligned(block_byte_base + q_offset);
752
+ let q2_u32 = load_u32_at_src0_aligned(block_byte_base + q_offset + 64u);
753
+ let qh_u32 = load_u32_at_src0_aligned(block_byte_base + qh_offset);
754
+
755
+ for (var col = 0u;col < NUM_COLS;col += 1) {
756
+ var vals = vec4<f32>(0.0, 0.0, 0.0, 0.0);
757
+ var sumy = vec4<f32>(0.0, 0.0, 0.0, 0.0);
758
+ for (var i = 0u; i < 4u; i++) {
759
+ let q1b = byte_of(q1_u32, i);
760
+ let q2b = byte_of(q2_u32, i);
761
+ let qhb = byte_of(qh_u32, i);
762
+
763
+ let yl0 = x_block[col][i];
764
+ let yl8 = x_block[col][i + 4u];
765
+ let yh0 = x_block[col][i + 8u];
766
+ let yh8 = x_block[col][i + 12u];
767
+
768
+ sumy[0] += yl0;
769
+ sumy[1] += yl8;
770
+ sumy[2] += yh0;
771
+ sumy[3] += yh8;
772
+
773
+ let q0 = f32((q1b & 0x0Fu) | select(0u, 0x10u, (qhb & hm1) != 0u));
774
+ let q1 = f32((q1b >> 4u) | select(0u, 0x10u, (qhb & hm2) != 0u));
775
+ let q2 = f32((q2b & 0x0Fu) | select(0u, 0x10u, (qhb & hm3) != 0u));
776
+ let q3 = f32((q2b >> 4u) | select(0u, 0x10u, (qhb & hm4) != 0u));
777
+
778
+ vals[0] += yl0 * q0;
779
+ vals[1] += yl8 * q1;
780
+ vals[2] += yh0 * q2;
781
+ vals[3] += yh8 * q3;
782
+ }
783
+
784
+ acc[col][row] += d * (f0 * vals[0] + f1 * vals[1] + f4 * vals[2] + f5 * vals[3])
785
+ - dmin * (sumy[0] * m0 + sumy[1] * m1 +
786
+ sumy[2] * m4 + sumy[3] * m5);
787
+ }
788
+ }
789
+ }
790
+ }
791
+
792
+ return acc;
793
+ }
794
+ #endif
795
+
796
+ #ifdef MUL_ACC_Q6_K
797
+ #define BLOCK_SIZE 256
798
+ #define BLOCK_SIZE_BYTES 210
799
+ #define THREADS_PER_BLOCK 16
800
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
801
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
802
+
803
+ let tid = thread_id % THREADS_PER_BLOCK;
804
+ let block_group = thread_id / THREADS_PER_BLOCK;
805
+ let num_block_groups: u32 = WG_SIZE / THREADS_PER_BLOCK;
806
+
807
+ let ip = tid / 8u;
808
+ let il = tid % 8u;
809
+ let l0 = 4u * il;
810
+ let is = 8u * ip + l0 / 16u;
811
+
812
+ let y_offset = 128u * ip + l0;
813
+ let q_offset_l = 64u * ip + l0;
814
+ let q_offset_h = 32u * ip + l0;
815
+
816
+ let num_blocks = params.k / BLOCK_SIZE;
817
+ let sc_base_byte = 192u + (is & ~3u);
818
+ let sc_byte_pos = is & 3u;
819
+
820
+ for (var block = block_group; block < num_blocks; block += num_block_groups) {
821
+ let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
822
+ var x_block: array<array<f32, 16>, NUM_COLS>;
823
+ for (var col = 0u; col < NUM_COLS;col += 1) {
824
+ let col_base = x_base + col * params.stride_11;
825
+ for (var l = 0u; l < 4u; l++) {
826
+ x_block[col][l] = f32(src1[col_base + l]);
827
+ x_block[col][l + 4u] = f32(src1[col_base + 32u + l]);
828
+ x_block[col][l + 8u] = f32(src1[col_base + 64u + l]);
829
+ x_block[col][l + 12u] = f32(src1[col_base + 96u + l]);
830
+ }
831
+ }
832
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
833
+ let output_row = row_base + row;
834
+ if (output_row < params.m) {
835
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
836
+
837
+ let d = f32(load_f16_at_src0(block_byte_base + 208u));
838
+ let ql1_u32 = load_u32_at_src0(block_byte_base + q_offset_l);
839
+ let ql2_u32 = load_u32_at_src0(block_byte_base + q_offset_l + 32u);
840
+ let qh_u32 = load_u32_at_src0(block_byte_base + 128u + q_offset_h);
841
+ let sc_u32_0 = load_u32_at_src0(block_byte_base + sc_base_byte);
842
+ let sc_u32_1 = load_u32_at_src0(block_byte_base + sc_base_byte + 4u);
843
+
844
+ let sc0 = sbyte_of(sc_u32_0, sc_byte_pos);
845
+ let sc2 = sbyte_of(sc_u32_0, sc_byte_pos + 2u);
846
+ let sc4 = sbyte_of(sc_u32_1, sc_byte_pos);
847
+ let sc6 = sbyte_of(sc_u32_1, sc_byte_pos + 2u);
848
+
849
+ for (var col = 0u;col < NUM_COLS;col += 1) {
850
+ var sums = vec4<f32>(0.0, 0.0, 0.0, 0.0);
851
+
852
+ for (var l = 0u; l < 4u; l++) {
853
+ let q1b = byte_of(ql1_u32, l);
854
+ let q2b = byte_of(ql2_u32, l);
855
+ let qhb = byte_of(qh_u32, l);
856
+
857
+ let dq0 = f32(i32((q1b & 0x0Fu) | ((qhb & 0x03u) << 4u)) - 32);
858
+ let dq1 = f32(i32((q2b & 0x0Fu) | ((qhb & 0x0Cu) << 2u)) - 32);
859
+ let dq2 = f32(i32((q1b >> 4u) | (qhb & 0x30u)) - 32);
860
+ let dq3 = f32(i32((q2b >> 4u) | ((qhb & 0xC0u) >> 2u)) - 32);
861
+
862
+ sums[0] += x_block[col][l] * dq0;
863
+ sums[1] += x_block[col][l + 4u] * dq1;
864
+ sums[2] += x_block[col][l + 8u] * dq2;
865
+ sums[3] += x_block[col][l + 12u] * dq3;
866
+ }
867
+
868
+ acc[col][row] += d * (sums[0] * f32(sc0) + sums[1] * f32(sc2) +
869
+ sums[2] * f32(sc4) + sums[3] * f32(sc6));
870
+ }
871
+ }
872
+ }
873
+ }
874
+
875
+ return acc;
876
+ }
877
+ #endif
878
+
879
+ #ifdef MUL_ACC_IQ1_S
880
+ #define BLOCK_SIZE 256
881
+ #define BLOCK_SIZE_BYTES 50
882
+ #define THREADS_PER_BLOCK 16
883
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
884
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
885
+
886
+ let tid = thread_id % THREADS_PER_BLOCK;
887
+ let block_group = thread_id / THREADS_PER_BLOCK;
888
+ let num_block_groups: u32 = WG_SIZE / THREADS_PER_BLOCK;
889
+
890
+ let sub_blk = tid / 2u;
891
+ let half = tid % 2u;
892
+ let slot0 = half * 2u;
893
+ let y_offset = sub_blk * 32u + slot0 * 8u;
894
+
895
+ let num_blocks = params.k / BLOCK_SIZE;
896
+
897
+ for (var block = block_group; block < num_blocks; block += num_block_groups) {
898
+ let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
899
+ var x_block: array<array<f32, 16>, NUM_COLS>;
900
+ for (var col = 0u; col < NUM_COLS;col += 1) {
901
+ for (var i = 0u; i < 16u; i++) {
902
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
903
+ }
904
+ }
905
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
906
+ let output_row = row_base + row;
907
+ if (output_row < params.m) {
908
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
909
+
910
+ let d = f32(load_f16_at_src0(block_byte_base));
911
+ let qh = load_u32_at_src0(block_byte_base + 34u + sub_blk * 2u) & 0xFFFFu;
912
+ let dl = d * f32(2u * ((qh >> 12u) & 7u) + 1u);
913
+ let delta = select(IQ1_DELTA, -IQ1_DELTA, (qh & 0x8000u) != 0u);
914
+ let qs_w = load_u32_at_src0(block_byte_base + 2u + sub_blk * 4u);
915
+
916
+ for (var col = 0u;col < NUM_COLS;col += 1) {
917
+ var row_sum = 0.0;
918
+ for (var ll = 0u; ll < 2u; ll++) {
919
+ let l = slot0 + ll;
920
+ let qs_byte = get_byte(qs_w, l);
921
+ let ig = (qs_byte | (((qh >> (3u * l)) & 7u) << 8u)) * 8u;
922
+ let gw = iq1_grid[ig / 16u];
923
+ let bit_base = (ig % 16u) * 2u;
924
+ for (var j = 0u; j < 8u; j++) {
925
+ let g = (gw >> (bit_base + j * 2u)) & 3u;
926
+ let gs = select(f32(g), f32(g) - 4.0, (g & 2u) != 0u);
927
+ row_sum += dl * (gs + delta) * x_block[col][ll * 8u + j];
928
+ }
929
+ }
930
+ acc[col][row] += row_sum;
931
+ }
932
+ }
933
+ }
934
+ }
935
+
936
+ return acc;
937
+ }
938
+ #endif
939
+
940
+ #ifdef MUL_ACC_IQ1_M
941
+ #define BLOCK_SIZE 256
942
+ #define BLOCK_SIZE_BYTES 56
943
+ #define THREADS_PER_BLOCK 16
944
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
945
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
946
+
947
+ let tid = thread_id % THREADS_PER_BLOCK;
948
+ let block_group = thread_id / THREADS_PER_BLOCK;
949
+ let num_block_groups: u32 = WG_SIZE / THREADS_PER_BLOCK;
950
+
951
+ let sub_blk = tid / 2u;
952
+ let half = tid % 2u;
953
+ let slot0 = half * 2u;
954
+ let y_offset = sub_blk * 32u + slot0 * 8u;
955
+
956
+ let num_blocks = params.k / BLOCK_SIZE;
957
+
958
+ for (var block = block_group; block < num_blocks; block += num_block_groups) {
959
+ let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
960
+ var x_block: array<array<f32, 16>, NUM_COLS>;
961
+ for (var col = 0u; col < NUM_COLS;col += 1) {
962
+ for (var i = 0u; i < 16u; i++) {
963
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
964
+ }
965
+ }
966
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
967
+ let output_row = row_base + row;
968
+ if (output_row < params.m) {
969
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
970
+
971
+ let sc_lo = load_u32_at_src0(block_byte_base + 48u);
972
+ let sc_hi = load_u32_at_src0(block_byte_base + 52u);
973
+ let sc0 = sc_lo & 0xFFFFu;
974
+ let sc1 = (sc_lo >> 16u) & 0xFFFFu;
975
+ let sc2 = sc_hi & 0xFFFFu;
976
+ let sc3 = (sc_hi >> 16u) & 0xFFFFu;
977
+ let d_bits = (sc0 >> 12u) | ((sc1 >> 8u) & 0xF0u) | ((sc2 >> 4u) & 0xF00u) | (sc3 & 0xF000u);
978
+ let d = f32(bitcast<vec2<f16>>(d_bits)[0]);
979
+
980
+ let sc_u16 = select(select(sc2, sc3, sub_blk >= 6u),
981
+ select(sc0, sc1, sub_blk >= 2u),
982
+ sub_blk < 4u);
983
+
984
+ let qs_w = load_u32_at_src0(block_byte_base + sub_blk * 4u);
985
+ let qh = load_u32_at_src0(block_byte_base + 32u + sub_blk * 2u) & 0xFFFFu;
986
+ let qh_lo = qh & 0xFFu;
987
+ let qh_hi = (qh >> 8u) & 0xFFu;
988
+
989
+ for (var col = 0u;col < NUM_COLS;col += 1) {
990
+ var row_sum = 0.0;
991
+ for (var ll = 0u; ll < 2u; ll++) {
992
+ let l = slot0 + ll;
993
+ let bit_off = 6u * (sub_blk % 2u) + 3u * (l / 2u);
994
+ let sub_scale = (sc_u16 >> bit_off) & 0x7u;
995
+ let dl = d * f32(2u * sub_scale + 1u);
996
+ let qh_byte = select(qh_lo, qh_hi, l >= 2u);
997
+ let ll2 = l % 2u;
998
+ let grid_idx = get_byte(qs_w, l) | (((qh_byte >> (4u * ll2)) & 7u) << 8u);
999
+ let delta = select(IQ1_DELTA, -IQ1_DELTA, ((qh_byte >> (3u + 4u * ll2)) & 1u) != 0u);
1000
+ let ig = grid_idx * 8u;
1001
+ let gw = iq1_grid[ig / 16u];
1002
+ let bit_base = (ig % 16u) * 2u;
1003
+ for (var j = 0u; j < 8u; j++) {
1004
+ let g = (gw >> (bit_base + j * 2u)) & 3u;
1005
+ let gs = select(f32(g), f32(g) - 4.0, (g & 2u) != 0u);
1006
+ row_sum += dl * (gs + delta) * x_block[col][ll * 8u + j];
1007
+ }
1008
+ }
1009
+ acc[col][row] += row_sum;
1010
+ }
1011
+ }
1012
+ }
1013
+ }
1014
+
1015
+ return acc;
1016
+ }
1017
+ #endif
1018
+
1019
+ #ifdef MUL_ACC_IQ2_XXS
1020
+ #define BLOCK_SIZE 256
1021
+ #define BLOCK_SIZE_BYTES 66
1022
+ #define THREADS_PER_BLOCK 16
1023
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
1024
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
1025
+
1026
+ let tid = thread_id % THREADS_PER_BLOCK;
1027
+ let block_group = thread_id / THREADS_PER_BLOCK;
1028
+ let num_block_groups: u32 = WG_SIZE / THREADS_PER_BLOCK;
1029
+
1030
+ let sub_blk = tid / 2u;
1031
+ let half = tid % 2u;
1032
+ let slot0 = half * 2u;
1033
+ let y_offset = sub_blk * 32u + slot0 * 8u;
1034
+
1035
+ let num_blocks = params.k / BLOCK_SIZE;
1036
+
1037
+ for (var block = block_group; block < num_blocks; block += num_block_groups) {
1038
+ let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
1039
+ var x_block: array<array<f32, 16>, NUM_COLS>;
1040
+ for (var col = 0u; col < NUM_COLS;col += 1) {
1041
+ for (var i = 0u; i < 16u; i++) {
1042
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
1043
+ }
1044
+ }
1045
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
1046
+ let output_row = row_base + row;
1047
+ if (output_row < params.m) {
1048
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
1049
+ let d = f32(load_f16_at_src0(block_byte_base));
1050
+ let aux_lo = load_u32_at_src0(block_byte_base + 2u + sub_blk * 8u);
1051
+ let aux_hi = load_u32_at_src0(block_byte_base + 2u + sub_blk * 8u + 4u);
1052
+ let ls = aux_hi >> 28u;
1053
+ let db = d * (0.5 + f32(ls)) * 0.25;
1054
+
1055
+ for (var col = 0u;col < NUM_COLS;col += 1) {
1056
+ var row_sum = 0.0;
1057
+ for (var ll = 0u; ll < 2u; ll++) {
1058
+ let l = slot0 + ll;
1059
+ let grid_idx = (aux_lo >> (8u * l)) & 0xFFu;
1060
+ let signs_idx = (aux_hi >> (7u * l)) & 0x7Fu;
1061
+ let signs = (ksigns_iq2xs[signs_idx / 4u] >> ((signs_idx % 4u) * 8u)) & 0xFFu;
1062
+ let gw_lo = iq2xxs_grid[grid_idx * 2u];
1063
+ let gw_hi = iq2xxs_grid[grid_idx * 2u + 1u];
1064
+ for (var j = 0u; j < 8u; j++) {
1065
+ let gw = select(gw_hi, gw_lo, j < 4u);
1066
+ let b = f32((gw >> ((j & 3u) * 8u)) & 0xFFu);
1067
+ let s = select(1.0, -1.0, ((signs >> j) & 1u) != 0u);
1068
+ row_sum += db * b * s * x_block[col][ll * 8u + j];
1069
+ }
1070
+ }
1071
+ acc[col][row] += row_sum;
1072
+ }
1073
+ }
1074
+ }
1075
+ }
1076
+
1077
+ return acc;
1078
+ }
1079
+ #endif
1080
+
1081
+ #ifdef MUL_ACC_IQ2_XS
1082
+ #define BLOCK_SIZE 256
1083
+ #define BLOCK_SIZE_BYTES 74
1084
+ #define THREADS_PER_BLOCK 16
1085
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
1086
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
1087
+
1088
+ let tid = thread_id % THREADS_PER_BLOCK;
1089
+ let block_group = thread_id / THREADS_PER_BLOCK;
1090
+ let num_block_groups: u32 = WG_SIZE / THREADS_PER_BLOCK;
1091
+
1092
+ let sub_blk = tid / 2u;
1093
+ let half = tid % 2u;
1094
+ let slot0 = half * 2u;
1095
+ let y_offset = sub_blk * 32u + slot0 * 8u;
1096
+
1097
+ let num_blocks = params.k / BLOCK_SIZE;
1098
+
1099
+ for (var block = block_group; block < num_blocks; block += num_block_groups) {
1100
+ let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
1101
+ var x_block: array<array<f32, 16>, NUM_COLS>;
1102
+ for (var col = 0u; col < NUM_COLS;col += 1) {
1103
+ for (var i = 0u; i < 16u; i++) {
1104
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
1105
+ }
1106
+ }
1107
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
1108
+ let output_row = row_base + row;
1109
+ if (output_row < params.m) {
1110
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
1111
+ let d = f32(load_f16_at_src0(block_byte_base));
1112
+ let qs_lo = load_u32_at_src0(block_byte_base + 2u + sub_blk * 8u);
1113
+ let qs_hi = load_u32_at_src0(block_byte_base + 2u + sub_blk * 8u + 4u);
1114
+ let scales_word = load_u32_at_src0(block_byte_base + 66u + (sub_blk / 4u) * 4u);
1115
+ let scales_byte = get_byte(scales_word, sub_blk % 4u);
1116
+
1117
+ for (var col = 0u;col < NUM_COLS;col += 1) {
1118
+ var row_sum = 0.0;
1119
+ for (var ll = 0u; ll < 2u; ll++) {
1120
+ let l = slot0 + ll;
1121
+ let qs_word = select(qs_hi, qs_lo, l < 2u);
1122
+ let half2 = (l % 2u) * 16u;
1123
+ let qs_val = (qs_word >> half2) & 0xFFFFu;
1124
+ let grid_idx = qs_val & 0x1FFu;
1125
+ let signs_idx = (qs_val >> 9u) & 0x7Fu;
1126
+ let sub_scale = (scales_byte >> (4u * (l / 2u))) & 0xFu;
1127
+ let db = d * (0.5 + f32(sub_scale)) * 0.25;
1128
+ let signs = (ksigns_iq2xs[signs_idx / 4u] >> ((signs_idx % 4u) * 8u)) & 0xFFu;
1129
+ let gw_lo = iq2xs_grid[grid_idx * 2u];
1130
+ let gw_hi = iq2xs_grid[grid_idx * 2u + 1u];
1131
+ for (var j = 0u; j < 8u; j++) {
1132
+ let gw = select(gw_hi, gw_lo, j < 4u);
1133
+ let b = f32((gw >> ((j & 3u) * 8u)) & 0xFFu);
1134
+ let s = select(1.0, -1.0, ((signs >> j) & 1u) != 0u);
1135
+ row_sum += db * b * s * x_block[col][ll * 8u + j];
1136
+ }
1137
+ }
1138
+ acc[col][row] += row_sum;
1139
+ }
1140
+ }
1141
+ }
1142
+ }
1143
+
1144
+ return acc;
1145
+ }
1146
+ #endif
1147
+
1148
+ #ifdef MUL_ACC_IQ2_S
1149
+ #define BLOCK_SIZE 256
1150
+ #define BLOCK_SIZE_BYTES 82
1151
+ #define THREADS_PER_BLOCK 16
1152
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
1153
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
1154
+
1155
+ let tid = thread_id % THREADS_PER_BLOCK;
1156
+ let block_group = thread_id / THREADS_PER_BLOCK;
1157
+ let num_block_groups: u32 = WG_SIZE / THREADS_PER_BLOCK;
1158
+
1159
+ let sub_blk = tid / 2u;
1160
+ let half = tid % 2u;
1161
+ let slot0 = half * 2u;
1162
+ let y_offset = sub_blk * 32u + slot0 * 8u;
1163
+
1164
+ let num_blocks = params.k / BLOCK_SIZE;
1165
+
1166
+ for (var block = block_group; block < num_blocks; block += num_block_groups) {
1167
+ let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
1168
+ var x_block: array<array<f32, 16>, NUM_COLS>;
1169
+ for (var col = 0u; col < NUM_COLS;col += 1) {
1170
+ for (var i = 0u; i < 16u; i++) {
1171
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
1172
+ }
1173
+ }
1174
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
1175
+ let output_row = row_base + row;
1176
+ if (output_row < params.m) {
1177
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
1178
+ let d = f32(load_f16_at_src0(block_byte_base));
1179
+ let qs_w = load_u32_at_src0(block_byte_base + 2u + sub_blk * 4u);
1180
+ let sg_w = load_u32_at_src0(block_byte_base + 34u + sub_blk * 4u);
1181
+ let qh_word = load_u32_at_src0(block_byte_base + 66u + (sub_blk / 4u) * 4u);
1182
+ let qh_byte = get_byte(qh_word, sub_blk % 4u);
1183
+ let sc_word = load_u32_at_src0(block_byte_base + 74u + (sub_blk / 4u) * 4u);
1184
+ let scales_byte = get_byte(sc_word, sub_blk % 4u);
1185
+
1186
+ for (var col = 0u;col < NUM_COLS;col += 1) {
1187
+ var row_sum = 0.0;
1188
+ for (var ll = 0u; ll < 2u; ll++) {
1189
+ let l = slot0 + ll;
1190
+ let qs_byte = get_byte(qs_w, l);
1191
+ let sign_byte = get_byte(sg_w, l);
1192
+ let grid_idx = qs_byte | (((qh_byte >> (2u * l)) & 3u) << 8u);
1193
+ let sub_scale = (scales_byte >> (4u * (l / 2u))) & 0xFu;
1194
+ let db = d * (0.5 + f32(sub_scale)) * 0.25;
1195
+ let gw_lo = iq2s_grid[grid_idx * 2u];
1196
+ let gw_hi = iq2s_grid[grid_idx * 2u + 1u];
1197
+ for (var j = 0u; j < 8u; j++) {
1198
+ let gw = select(gw_hi, gw_lo, j < 4u);
1199
+ let b = f32((gw >> ((j & 3u) * 8u)) & 0xFFu);
1200
+ let s = select(1.0, -1.0, ((sign_byte >> j) & 1u) != 0u);
1201
+ row_sum += db * b * s * x_block[col][ll * 8u + j];
1202
+ }
1203
+ }
1204
+ acc[col][row] += row_sum;
1205
+ }
1206
+ }
1207
+ }
1208
+ }
1209
+
1210
+ return acc;
1211
+ }
1212
+ #endif
1213
+
1214
+ #ifdef MUL_ACC_IQ3_XXS
1215
+ #define BLOCK_SIZE 256
1216
+ #define BLOCK_SIZE_BYTES 98
1217
+ #define THREADS_PER_BLOCK 16
1218
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
1219
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
1220
+
1221
+ let tid = thread_id % THREADS_PER_BLOCK;
1222
+ let block_group = thread_id / THREADS_PER_BLOCK;
1223
+ let num_block_groups: u32 = WG_SIZE / THREADS_PER_BLOCK;
1224
+
1225
+ let sub_blk = tid / 2u;
1226
+ let half = tid % 2u;
1227
+ let slot0 = half * 2u;
1228
+ let y_offset = sub_blk * 32u + slot0 * 8u;
1229
+
1230
+ let num_blocks = params.k / BLOCK_SIZE;
1231
+
1232
+ for (var block = block_group; block < num_blocks; block += num_block_groups) {
1233
+ let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
1234
+ var x_block: array<array<f32, 16>, NUM_COLS>;
1235
+ for (var col = 0u; col < NUM_COLS;col += 1) {
1236
+ for (var i = 0u; i < 16u; i++) {
1237
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
1238
+ }
1239
+ }
1240
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
1241
+ let output_row = row_base + row;
1242
+ if (output_row < params.m) {
1243
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
1244
+ let d = f32(load_f16_at_src0(block_byte_base));
1245
+ let qs_lo = load_u32_at_src0(block_byte_base + 2u + sub_blk * 8u);
1246
+ let qs_hi = load_u32_at_src0(block_byte_base + 2u + sub_blk * 8u + 4u);
1247
+ let aux = load_u32_at_src0(block_byte_base + 66u + sub_blk * 4u);
1248
+ let ls = aux >> 28u;
1249
+ let db = d * (0.5 + f32(ls)) * 0.5;
1250
+
1251
+ for (var col = 0u;col < NUM_COLS;col += 1) {
1252
+ var row_sum = 0.0;
1253
+ for (var ll = 0u; ll < 2u; ll++) {
1254
+ let l = slot0 + ll;
1255
+ let qs_word = select(qs_hi, qs_lo, l < 2u);
1256
+ let byte_pos = (l % 2u) * 2u;
1257
+ let grid_idx_0 = (qs_word >> (byte_pos * 8u)) & 0xFFu;
1258
+ let grid_idx_1 = (qs_word >> ((byte_pos + 1u) * 8u)) & 0xFFu;
1259
+ let signs_idx = (aux >> (7u * l)) & 0x7Fu;
1260
+ let signs = (ksigns_iq2xs[signs_idx / 4u] >> ((signs_idx % 4u) * 8u)) & 0xFFu;
1261
+ let grid1 = iq3xxs_grid[grid_idx_0];
1262
+ let grid2 = iq3xxs_grid[grid_idx_1];
1263
+ for (var j = 0u; j < 4u; j++) {
1264
+ let b1 = f32((grid1 >> (j * 8u)) & 0xFFu);
1265
+ let b2 = f32((grid2 >> (j * 8u)) & 0xFFu);
1266
+ let s1 = select(1.0, -1.0, ((signs >> j) & 1u) != 0u);
1267
+ let s2 = select(1.0, -1.0, ((signs >> (j + 4u)) & 1u) != 0u);
1268
+ row_sum += db * b1 * s1 * x_block[col][ll * 8u + j];
1269
+ row_sum += db * b2 * s2 * x_block[col][ll * 8u + j + 4u];
1270
+ }
1271
+ }
1272
+ acc[col][row] += row_sum;
1273
+ }
1274
+ }
1275
+ }
1276
+ }
1277
+
1278
+ return acc;
1279
+ }
1280
+ #endif
1281
+
1282
+ #ifdef MUL_ACC_IQ3_S
1283
+ #define BLOCK_SIZE 256
1284
+ #define BLOCK_SIZE_BYTES 110
1285
+ #define THREADS_PER_BLOCK 16
1286
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
1287
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
1288
+
1289
+ let tid = thread_id % THREADS_PER_BLOCK;
1290
+ let block_group = thread_id / THREADS_PER_BLOCK;
1291
+ let num_block_groups: u32 = WG_SIZE / THREADS_PER_BLOCK;
1292
+
1293
+ let sub_blk = tid / 2u;
1294
+ let half = tid % 2u;
1295
+ let slot0 = half * 2u;
1296
+ let y_offset = sub_blk * 32u + slot0 * 8u;
1297
+
1298
+ let num_blocks = params.k / BLOCK_SIZE;
1299
+
1300
+ for (var block = block_group; block < num_blocks; block += num_block_groups) {
1301
+ let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
1302
+ var x_block: array<array<f32, 16>, NUM_COLS>;
1303
+ for (var col = 0u; col < NUM_COLS;col += 1) {
1304
+ for (var i = 0u; i < 16u; i++) {
1305
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
1306
+ }
1307
+ }
1308
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
1309
+ let output_row = row_base + row;
1310
+ if (output_row < params.m) {
1311
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
1312
+ let d = f32(load_f16_at_src0(block_byte_base));
1313
+ let qs_lo = load_u32_at_src0(block_byte_base + 2u + sub_blk * 8u);
1314
+ let qs_hi = load_u32_at_src0(block_byte_base + 2u + sub_blk * 8u + 4u);
1315
+ let qh_word = load_u32_at_src0(block_byte_base + 66u + (sub_blk / 4u) * 4u);
1316
+ let qh_byte = get_byte(qh_word, sub_blk % 4u);
1317
+ let sg_w = load_u32_at_src0(block_byte_base + 74u + sub_blk * 4u);
1318
+ let sc_word = load_u32_at_src0(block_byte_base + 106u);
1319
+ let scales_byte = get_byte(sc_word, sub_blk / 2u);
1320
+ let sub_scale = (scales_byte >> (4u * (sub_blk % 2u))) & 0xFu;
1321
+ let db = d * (1.0 + 2.0 * f32(sub_scale));
1322
+
1323
+ for (var col = 0u;col < NUM_COLS;col += 1) {
1324
+ var row_sum = 0.0;
1325
+ for (var ll = 0u; ll < 2u; ll++) {
1326
+ let l = slot0 + ll;
1327
+ let qs_word = select(qs_hi, qs_lo, l < 2u);
1328
+ let byte_pos = (l % 2u) * 2u;
1329
+ let qs0 = (qs_word >> (byte_pos * 8u)) & 0xFFu;
1330
+ let qs1 = (qs_word >> ((byte_pos + 1u) * 8u)) & 0xFFu;
1331
+ let grid_idx_1 = qs0 | (((qh_byte >> (2u * l)) & 1u) << 8u);
1332
+ let grid_idx_2 = qs1 | (((qh_byte >> (2u * l + 1u)) & 1u) << 8u);
1333
+ let sign_byte = get_byte(sg_w, l);
1334
+ let grid1 = iq3s_grid[grid_idx_1];
1335
+ let grid2 = iq3s_grid[grid_idx_2];
1336
+ for (var j = 0u; j < 4u; j++) {
1337
+ let b1 = f32((grid1 >> (j * 8u)) & 0xFFu);
1338
+ let b2 = f32((grid2 >> (j * 8u)) & 0xFFu);
1339
+ let s1 = select(1.0, -1.0, ((sign_byte >> j) & 1u) != 0u);
1340
+ let s2 = select(1.0, -1.0, ((sign_byte >> (j + 4u)) & 1u) != 0u);
1341
+ row_sum += db * b1 * s1 * x_block[col][ll * 8u + j];
1342
+ row_sum += db * b2 * s2 * x_block[col][ll * 8u + j + 4u];
1343
+ }
1344
+ }
1345
+ acc[col][row] += row_sum;
1346
+ }
1347
+ }
1348
+ }
1349
+ }
1350
+
1351
+ return acc;
1352
+ }
1353
+ #endif
1354
+
1355
+ #ifdef MUL_ACC_IQ4_NL
1356
+ #define BLOCK_SIZE 32
1357
+ #define BLOCK_SIZE_BYTES 18
1358
+ #define THREADS_PER_BLOCK 4
1359
+ #define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
1360
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
1361
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
1362
+
1363
+ let num_blocks = params.k / BLOCK_SIZE;
1364
+ let thread_within_block = thread_id % THREADS_PER_BLOCK;
1365
+ for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
1366
+ let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * 4u;
1367
+ var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
1368
+ for (var col = 0u; col < NUM_COLS;col += 1) {
1369
+ for (var i = 0u; i < ELEMS_PER_THREAD / 2u; i++) {
1370
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
1371
+ x_block[col][i + 4u] = f32(src1[x_base + col * params.stride_11 + i + 16u]);
1372
+ }
1373
+ }
1374
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
1375
+ let output_row = row_base + row;
1376
+ if (output_row < params.m) {
1377
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
1378
+ let d = f32(load_f16_at_src0(block_byte_base));
1379
+ let q_packed = load_u32_at_src0(block_byte_base + 2u + 4u * thread_within_block);
1380
+ for (var col = 0u;col < NUM_COLS;col += 1) {
1381
+ var row_sum = 0.0;
1382
+ for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
1383
+ let q_byte = get_byte(q_packed, byte_idx);
1384
+ let q_lo = f32(kvalues_iq4nl[q_byte & 0xFu]) * d;
1385
+ let q_hi = f32(kvalues_iq4nl[(q_byte >> 4u) & 0xFu]) * d;
1386
+ row_sum += q_lo * x_block[col][byte_idx];
1387
+ row_sum += q_hi * x_block[col][byte_idx + 4u];
1388
+ }
1389
+ acc[col][row] += row_sum;
1390
+ }
1391
+ }
1392
+ }
1393
+ }
1394
+
1395
+ return acc;
1396
+ }
1397
+ #endif
1398
+
1399
+ #ifdef MUL_ACC_IQ4_XS
1400
+ #define BLOCK_SIZE 256
1401
+ #define BLOCK_SIZE_BYTES 136
1402
+ #define THREADS_PER_BLOCK 16
1403
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
1404
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
1405
+
1406
+ let tid = thread_id % THREADS_PER_BLOCK;
1407
+ let block_group = thread_id / THREADS_PER_BLOCK;
1408
+ let num_block_groups: u32 = WG_SIZE / THREADS_PER_BLOCK;
1409
+
1410
+ let sub_blk = tid / 2u;
1411
+ let half = tid % 2u;
1412
+ let y_offset = sub_blk * 32u + half * 16u;
1413
+
1414
+ let num_blocks = params.k / BLOCK_SIZE;
1415
+
1416
+ for (var block = block_group; block < num_blocks; block += num_block_groups) {
1417
+ let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
1418
+ var x_block: array<array<f32, 16>, NUM_COLS>;
1419
+ for (var col = 0u; col < NUM_COLS;col += 1) {
1420
+ for (var i = 0u; i < 16u; i++) {
1421
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
1422
+ }
1423
+ }
1424
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
1425
+ let output_row = row_base + row;
1426
+ if (output_row < params.m) {
1427
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
1428
+ let d = f32(load_f16_at_src0(block_byte_base));
1429
+ let scales_h = load_u16_at_src0(block_byte_base + 2u);
1430
+ let scales_l_word = load_u32_at_src0(block_byte_base + 4u);
1431
+ let sl_byte = get_byte(scales_l_word, sub_blk / 2u);
1432
+ let sl = (sl_byte >> (4u * (sub_blk % 2u))) & 0xFu;
1433
+ let sh_bits = (scales_h >> (2u * sub_blk)) & 3u;
1434
+ let ls = i32(sl | (sh_bits << 4u));
1435
+ let dl = d * f32(ls - 32);
1436
+
1437
+ let qs_byte_off = 8u + sub_blk * 16u;
1438
+ let q_w0 = load_u32_at_src0(block_byte_base + qs_byte_off);
1439
+ let q_w1 = load_u32_at_src0(block_byte_base + qs_byte_off + 4u);
1440
+ let q_w2 = load_u32_at_src0(block_byte_base + qs_byte_off + 8u);
1441
+ let q_w3 = load_u32_at_src0(block_byte_base + qs_byte_off + 12u);
1442
+
1443
+ for (var col = 0u;col < NUM_COLS;col += 1) {
1444
+ var row_sum = 0.0;
1445
+ for (var i = 0u; i < 16u; i++) {
1446
+ let q_word = select(
1447
+ select(q_w0, q_w1, i >= 4u),
1448
+ select(q_w2, q_w3, i >= 12u),
1449
+ i >= 8u);
1450
+ let q_byte = get_byte(q_word, i % 4u);
1451
+ let nib = select(q_byte & 0xFu, (q_byte >> 4u) & 0xFu, half == 1u);
1452
+ row_sum += f32(kvalues_iq4nl[nib]) * dl * x_block[col][i];
1453
+ }
1454
+ acc[col][row] += row_sum;
1455
+ }
1456
+ }
1457
+ }
1458
+ }
1459
+
1460
+ return acc;
1461
+ }
1462
+ #endif
1463
+
1464
+ #ifdef MUL_ACC_MXFP4
1465
+ #define BLOCK_SIZE 32
1466
+ #define BLOCK_SIZE_BYTES 17
1467
+ #define THREADS_PER_BLOCK 4
1468
+ #define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
1469
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
1470
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
1471
+
1472
+ let num_blocks = params.k / BLOCK_SIZE;
1473
+ let thread_within_block = thread_id % 4;
1474
+ for (var block = thread_id/THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE/THREADS_PER_BLOCK) {
1475
+ let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * 4;
1476
+ var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
1477
+ for (var col = 0u; col < NUM_COLS;col += 1) {
1478
+ for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
1479
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
1480
+ x_block[col][i + 4] = f32(src1[x_base + col * params.stride_11 + i + 16]);
1481
+ }
1482
+ }
1483
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
1484
+ let output_row = row_base + row;
1485
+ if (output_row < params.m) {
1486
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
1487
+ let eu8 = get_byte(load_u32_at_src0(block_byte_base), 0);
1488
+ let e = ldexp(1.0, i32(eu8) - 128);
1489
+ let q_packed = load_u32_at_src0(block_byte_base + 1u + 4u * thread_within_block);
1490
+ for (var col = 0u;col < NUM_COLS;col += 1) {
1491
+ var row_sum = 0.0;
1492
+ for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
1493
+ let q_byte = get_byte(q_packed, byte_idx);
1494
+ let q_lo = f32(kvalues_mxfp4[q_byte & 0xFu]) * e;
1495
+ let q_hi = f32(kvalues_mxfp4[(q_byte >> 4u) & 0xFu]) * e;
1496
+ row_sum += q_lo * x_block[col][byte_idx];
1497
+ row_sum += q_hi * x_block[col][byte_idx + 4u];
1498
+ }
1499
+ acc[col][row] += row_sum;
1500
+ }
1501
+ }
1502
+ }
1503
+ }
1504
+
1505
+ return acc;
1506
+ }
1507
+ #endif
1508
+
1509
+ #ifdef MUL_ACC_NVFP4
1510
+ #define BLOCK_SIZE 64
1511
+ #define BLOCK_SIZE_BYTES 36
1512
+ #define THREADS_PER_BLOCK 4
1513
+ #define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
1514
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
1515
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
1516
+
1517
+ let num_blocks = params.k / BLOCK_SIZE;
1518
+ let sub = thread_id % THREADS_PER_BLOCK;
1519
+ for (var block = thread_id/THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE/THREADS_PER_BLOCK) {
1520
+ let x_base = src1_idx_base + block * BLOCK_SIZE + sub * ELEMS_PER_THREAD;
1521
+ var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
1522
+ for (var col = 0u; col < NUM_COLS;col += 1) {
1523
+ for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
1524
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
1525
+ x_block[col][i + 8] = f32(src1[x_base + col * params.stride_11 + i + 8]);
1526
+ }
1527
+ }
1528
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
1529
+ let output_row = row_base + row;
1530
+ if (output_row < params.m) {
1531
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
1532
+ let d = ue4m3_to_fp32(get_byte(load_u32_at_src0_aligned(block_byte_base), sub)) * 0.5;
1533
+ let q_w0 = load_u32_at_src0_aligned(block_byte_base + 4u + 8u * sub);
1534
+ let q_w1 = load_u32_at_src0_aligned(block_byte_base + 8u + 8u * sub);
1535
+ for (var col = 0u;col < NUM_COLS;col += 1) {
1536
+ var row_sum = 0.0;
1537
+ for (var l = 0u; l < 8u; l++) {
1538
+ let q_word = select(q_w0, q_w1, l >= 4u);
1539
+ let q_byte = get_byte(q_word, l % 4u);
1540
+ let q_lo = f32(kvalues_mxfp4[q_byte & 0xFu]) * d;
1541
+ let q_hi = f32(kvalues_mxfp4[(q_byte >> 4u) & 0xFu]) * d;
1542
+ row_sum += q_lo * x_block[col][l];
1543
+ row_sum += q_hi * x_block[col][l + 8u];
1544
+ }
1545
+ acc[col][row] += row_sum;
1546
+ }
1547
+ }
1548
+ }
1549
+ }
1550
+
1551
+ return acc;
1552
+ }
1553
+ #endif