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
@@ -25,6 +25,10 @@ fn store_shmem(val: f16, idx: u32) {
25
25
  }
26
26
  #endif // SCALAR
27
27
 
28
+ #define QUANT_SHMEM shmem
29
+ #define QUANT_OUT_TYPE f16
30
+ #include "quant_inner_loops.tmpl"
31
+
28
32
  #ifdef INIT_SRC0_SHMEM_FLOAT
29
33
  fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
30
34
  for (var elem_idx = thread_id * VEC_SIZE; elem_idx < TILE_SRC0_SHMEM; elem_idx += TOTAL_WORKGROUP_SIZE * VEC_SIZE) {
@@ -42,6 +46,7 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
42
46
  }
43
47
  #endif // INIT_SRC0_SHMEM_FLOAT
44
48
 
49
+ #ifndef MUL_MAT_ID
45
50
  #ifdef INIT_SRC1_SHMEM_FLOAT
46
51
  fn init_shmem_src1(thread_id: u32, batch_offset: u32, offset_n: u32, k_outer: u32) {
47
52
  for (var elem_idx = thread_id * VEC_SIZE; elem_idx < TILE_SRC1_SHMEM; elem_idx += TOTAL_WORKGROUP_SIZE * VEC_SIZE) {
@@ -58,304 +63,265 @@ fn init_shmem_src1(thread_id: u32, batch_offset: u32, offset_n: u32, k_outer: u3
58
63
  }
59
64
  }
60
65
  #endif // INIT_SRC1_SHMEM_FLOAT
66
+ #endif
61
67
 
62
- #ifdef INIT_SRC0_SHMEM_Q4_0
63
- const BLOCK_SIZE = 32u;
64
- // the number of blocks per k-tile. Note that this currently only works if TILE_K is a multiple of BLOCK_SIZE, which may need to be rethought for larger quantized types.
65
- override BLOCKS_K = TILE_K/BLOCK_SIZE;
66
- const NQ = 16u;
67
- const F16_PER_BLOCK = 9u; // 1 scale + 8x4 packed weights
68
- const WEIGHTS_PER_F16 = 4u; // 4 weights per f16
69
- const F16_PER_THREAD = NQ / WEIGHTS_PER_F16;
68
+ #ifdef INIT_SRC0_SHMEM_Q1_0
69
+ const BLOCK_SIZE = 128u;
70
+ const BLOCK_SIZE_BYTES = 18u;
71
+ const NQ = 8u; // 8 weights (1 byte of qs) per thread per iteration
70
72
 
71
73
  fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
72
74
  for (var i = thread_id * NQ; i < TILE_SRC0_SHMEM; i += TOTAL_WORKGROUP_SIZE * NQ) {
73
- let blck_idx = i / BLOCK_SIZE;
74
- let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F16;
75
- let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * 2u;
76
-
77
- let tile_m = blck_idx / BLOCKS_K;
75
+ let tile_m = i / TILE_K;
76
+ let tile_k_start = i % TILE_K;
78
77
  let global_m = offset_m + tile_m;
79
- let block_k = blck_idx % BLOCKS_K;
80
- let global_k = k_outer / BLOCK_SIZE + block_k;
81
-
82
- if (global_m < params.m && global_k < params.k / BLOCK_SIZE) {
83
- let src0_idx = batch_offset + global_m * params.stride_01 + global_k;
84
- let scale_idx = src0_idx * F16_PER_BLOCK;
85
- let d = src0[scale_idx];
78
+ let global_k_start = k_outer + tile_k_start;
86
79
 
87
- for (var j = 0u; j < F16_PER_THREAD; j += 2) {
88
- let q_0 = src0[scale_idx + 1u + block_offset + j];
89
- let q_1 = src0[scale_idx + 1u + block_offset + j + 1];
80
+ if (global_m >= params.m) {
81
+ break;
82
+ }
90
83
 
91
- let q_packed = bitcast<u32>(vec2(q_0, q_1));
92
- for (var k = 0u; k < 4u; k++) {
93
- let q_byte = get_byte(q_packed, k);
94
- let q_hi = (f16((q_byte >> 4) & 0xF) - 8.0) * d;
95
- let q_lo = (f16(q_byte & 0xF) - 8.0) * d;
96
- shmem[shmem_idx + j * 2 + k] = q_lo;
97
- shmem[shmem_idx + j * 2 + k + 16u] = q_hi;
98
- }
84
+ let block_k = global_k_start / BLOCK_SIZE;
85
+ let byte_in_block = (global_k_start % BLOCK_SIZE) / 8u;
86
+ let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
87
+ let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
88
+ let d = load_f16_at_src0(block_byte_base);
89
+ let q_byte = load_u32_at_src0(block_byte_base + 2u + byte_in_block) & 0xFFu;
90
+
91
+ for (var bit = 0u; bit < NQ; bit++) {
92
+ let global_k = global_k_start + bit;
93
+ if (global_k < params.k) {
94
+ shmem[i + bit] = select(-d, d, ((q_byte >> bit) & 1u) != 0u);
99
95
  }
100
96
  }
101
97
  }
102
98
  }
103
- #endif // INIT_SRC0_SHMEM_Q4_0
99
+ #endif // INIT_SRC0_SHMEM_Q1_0
104
100
 
105
- #ifdef INIT_SRC0_SHMEM_Q4_1
101
+ // legacy-quants
102
+ #if defined(INIT_SRC0_SHMEM_Q4_0) || defined(INIT_SRC0_SHMEM_Q4_1) || defined(INIT_SRC0_SHMEM_Q5_0) || defined(INIT_SRC0_SHMEM_Q5_1) || defined(INIT_SRC0_SHMEM_Q8_0) || defined(INIT_SRC0_SHMEM_Q8_1) || defined(INIT_SRC0_SHMEM_MXFP4)
106
103
  const BLOCK_SIZE = 32u;
107
104
  // the number of blocks per k-tile. Note that this currently only works if TILE_K is a multiple of BLOCK_SIZE, which may need to be rethought for larger quantized types.
108
105
  override BLOCKS_K = TILE_K/BLOCK_SIZE;
109
106
  const NQ = 16u;
110
- const F16_PER_BLOCK = 10u; // 1 scale + 8 packed weights + 1 mean
111
- const WEIGHTS_PER_F16 = 4u; // 4 weights per f16
112
- const F16_PER_THREAD = NQ / WEIGHTS_PER_F16;
107
+ #if defined(INIT_SRC0_SHMEM_Q8_0) || defined(INIT_SRC0_SHMEM_Q8_1)
108
+ const BYTES_PER_THREAD = 16u; // NQ(16) weights use 16 bytes of q
109
+ #else
110
+ const BYTES_PER_THREAD = 8u; // NQ(16) weights use 8 bytes of q
111
+ #endif
112
+ const BYTES_PER_INNER_LOOP = 4u; // == sizeof(q_packed)
113
113
 
114
114
  fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
115
115
  for (var i = thread_id * NQ; i < TILE_SRC0_SHMEM; i += TOTAL_WORKGROUP_SIZE * NQ) {
116
- let blck_idx = i / BLOCK_SIZE;
117
- let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F16;
118
- let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * 2u;
116
+ let block_idx = i / BLOCK_SIZE;
117
+ let block_offset = (i % BLOCK_SIZE) / NQ;
118
+ let shmem_idx = block_idx * BLOCK_SIZE + block_offset * BYTES_PER_THREAD;
119
119
 
120
- let tile_m = blck_idx / BLOCKS_K;
120
+ let tile_m = block_idx / BLOCKS_K;
121
121
  let global_m = offset_m + tile_m;
122
- let block_k = blck_idx % BLOCKS_K;
123
- let global_k = k_outer / BLOCK_SIZE + block_k;
122
+ let block_k = block_idx % BLOCKS_K;
123
+ let global_block_k = k_outer / BLOCK_SIZE + block_k;
124
+
125
+ if (global_m < params.m && global_block_k < params.k / BLOCK_SIZE) {
126
+ let src0_idx = batch_offset + global_m * params.stride_01 + global_block_k;
127
+
128
+ #if defined(INIT_SRC0_SHMEM_Q4_0)
129
+ let block_byte_base = src0_idx * 18u; // BLOCK_SIZE_BYTES = 18u;
130
+ let d = load_f16_at_src0(block_byte_base);
124
131
 
125
- if (global_m < params.m && global_k < params.k / BLOCK_SIZE) {
126
- let src0_idx = batch_offset + global_m * params.stride_01 + global_k;
127
- let scale_idx = src0_idx * F16_PER_BLOCK;
128
- let d = src0[scale_idx];
129
- let m = src0[scale_idx + 1u];
132
+ // load NQ(16) weights
133
+ for (var j = 0u; j < BYTES_PER_THREAD / BYTES_PER_INNER_LOOP; j += 1) {
134
+ let q_byte_offset = block_byte_base + 2u + block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP;
135
+ let q_packed = load_u32_at_src0(q_byte_offset);
136
+ dequant_q4_0_packed_to_shmem(q_packed, d, shmem_idx + j * BYTES_PER_INNER_LOOP);
137
+ }
138
+ #endif // INIT_SRC0_SHMEM_Q4_0
130
139
 
131
- for (var j = 0u; j < F16_PER_THREAD; j += 2) {
132
- let q_0 = src0[scale_idx + 2u + block_offset + j];
133
- let q_1 = src0[scale_idx + 2u + block_offset + j + 1];
140
+ #if defined(INIT_SRC0_SHMEM_Q4_1)
141
+ let block_byte_base = src0_idx * 20u; // BLOCK_SIZE_BYTES = 20u;
142
+ let dm = unpack2x16float(load_u32_at_src0_aligned(block_byte_base));
143
+ let d = f16(dm[0]);
144
+ let m = f16(dm[1]);
134
145
 
135
- let q_packed = bitcast<u32>(vec2(q_0, q_1));
136
- for (var k = 0u; k < 4u; k++) {
146
+ // load NQ(16) weights
147
+ for (var j = 0u; j < BYTES_PER_THREAD / BYTES_PER_INNER_LOOP; j += 1) {
148
+ let q_byte_offset = block_byte_base + 4u + block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP;
149
+ let q_packed = load_u32_at_src0(q_byte_offset);
150
+
151
+ for (var k = 0u; k < BYTES_PER_INNER_LOOP; k++) {
137
152
  let q_byte = get_byte(q_packed, k);
138
153
  let q_lo = f16(q_byte & 0xF) * d + m;
139
154
  let q_hi = f16((q_byte >> 4) & 0xF) * d + m;
140
- shmem[shmem_idx + j * 2 + k] = q_lo;
141
- shmem[shmem_idx + j * 2 + k + 16u] = q_hi;
155
+ shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k] = q_lo;
156
+ shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k + 16u] = q_hi;
142
157
  }
143
158
  }
144
- }
145
- }
146
- }
147
159
  #endif // INIT_SRC0_SHMEM_Q4_1
148
160
 
149
- #ifdef INIT_SRC0_SHMEM_Q5_0
150
- // 32 weights per block, each at 4 bits each = 32 * 4 = 128 bits / 16 = 8 f16s per block
151
- const BLOCK_SIZE = 32u;
152
- // the number of blocks per k-tile. Note that this currently only works if TILE_K is a multiple of BLOCK_SIZE, which may need to be rethought for larger quantized types.
153
- // tile_k is defined as 32u, so blocks_k ends up being 1 always
154
- override BLOCKS_K = TILE_K / BLOCK_SIZE;
155
- const NQ = 16u;
156
- const F16_PER_BLOCK = 11u; // 1 scale + 2 qh + 8 packed weights
157
- const WEIGHTS_PER_F16 = 4u; // 4 weights per f16
158
- const F16_PER_THREAD = NQ / WEIGHTS_PER_F16; // 16 / 4 = 4 f16s per thread, each thread should handle 4 f16s * 4 weights per = 16 weights
159
-
160
- fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
161
+ #if defined(INIT_SRC0_SHMEM_Q5_0)
162
+ let block_byte_base = src0_idx * 22u; // BLOCK_SIZE_BYTES = 22u;
161
163
 
162
- for (var i = thread_id * NQ; i < TILE_SRC0_SHMEM; i += TOTAL_WORKGROUP_SIZE * NQ) {
163
- let blck_idx = i / BLOCK_SIZE;
164
- let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F16;
165
- let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * 2u;
166
-
167
- let tile_m = blck_idx / BLOCKS_K;
168
- let global_m = offset_m + tile_m;
169
- let block_k = blck_idx % BLOCKS_K;
170
- let global_k = k_outer / BLOCK_SIZE + block_k;
164
+ let d = load_f16_at_src0(block_byte_base);
165
+ let qh_packed = load_u32_at_src0(block_byte_base + 2u);
171
166
 
172
- if (global_m < params.m && global_k < params.k / BLOCK_SIZE) {
173
- let src0_idx = batch_offset + global_m * params.stride_01 + global_k;
174
- let scale_idx = src0_idx * F16_PER_BLOCK;
167
+ // load NQ(16) weights
168
+ for (var j = 0u; j < BYTES_PER_THREAD / BYTES_PER_INNER_LOOP; j += 1) {
169
+ let q_byte_offset = block_byte_base + 6u + block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP;
170
+ let q_packed = load_u32_at_src0(q_byte_offset);
175
171
 
176
- let d = src0[scale_idx];
177
- let qh0 = src0[scale_idx + 1u];
178
- let qh1 = src0[scale_idx + 2u];
179
- let qh_packed = bitcast<u32>(vec2(qh0, qh1));
180
-
181
- for (var j = 0u; j < 2; j++) {
182
- let q_0 = src0[scale_idx + 3u + block_offset + (j*2)];
183
- let q_1 = src0[scale_idx + 3u + block_offset + (j*2) + 1u];
184
-
185
- let q_packed = bitcast<u32>(vec2(q_0, q_1));
186
-
187
- let j_adjusted = j + (block_offset / 2u);
188
-
189
-
190
- for (var k = 0u; k < 4u; k++) {
172
+ for (var k = 0u; k < BYTES_PER_INNER_LOOP; k++) {
191
173
  let q_byte = get_byte(q_packed, k);
192
174
 
193
- let qh_hi = (qh_packed >> (j_adjusted * 4 + k + 12)) & 0x10;
175
+ let byte_idx = block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP + k;
176
+ let qh_hi = (qh_packed >> (byte_idx + 12u)) & 0x10;
194
177
  let q_hi = (f16(((q_byte >> 4) & 0xF) | qh_hi) - 16.0) * d;
195
- let qh_lo = ((qh_packed >> (j_adjusted * 4 + k)) << 4) & 0x10;
178
+ let qh_lo = ((qh_packed >> byte_idx) << 4) & 0x10;
196
179
  let q_lo = (f16((q_byte & 0xF) | qh_lo) - 16.0) * d;
197
-
198
- shmem[shmem_idx + j * 4u + k] = q_lo; // store first weight
199
- shmem[shmem_idx + j * 4u + k + 16u] = q_hi; // store second weight
180
+ shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k] = q_lo;
181
+ shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k + 16u] = q_hi;
200
182
  }
201
183
  }
202
- }
203
- }
204
- }
205
184
  #endif // INIT_SRC0_SHMEM_Q5_0
206
185
 
207
- #ifdef INIT_SRC0_SHMEM_Q5_1
208
- // 32 weights per block, each at 4 bits each = 32 * 4 = 128 bits / 16 = 8 f16s per block
209
- const BLOCK_SIZE = 32u;
210
- // the number of blocks per k-tile. Note that this currently only works if TILE_K is a multiple of BLOCK_SIZE, which may need to be rethought for larger quantized types.
211
- // tile_k is defined as 32u, so blocks_k ends up being 1 always
212
- override BLOCKS_K = TILE_K / BLOCK_SIZE;
213
- const NQ = 16u;
214
- const F16_PER_BLOCK = 12u; // 1 scale + 2 qh + 8 packed weights + 1 mean
215
- const WEIGHTS_PER_F16 = 4u; // 4 weights per f16
216
- const F16_PER_THREAD = NQ / WEIGHTS_PER_F16; // 16 / 4 = 4 f16s per thread, each thread should handle 4 f16s * 4 weights per = 16 weights
186
+ #if defined(INIT_SRC0_SHMEM_Q5_1)
187
+ let block_byte_base = src0_idx * 24u; // BLOCK_SIZE_BYTES = 24u;
217
188
 
218
- fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
189
+ let dm = unpack2x16float(load_u32_at_src0_aligned(block_byte_base));
190
+ let d = f16(dm[0]);
191
+ let m = f16(dm[1]);
192
+ let qh_packed = load_u32_at_src0_aligned(block_byte_base + 4u);
219
193
 
220
- for (var i = thread_id * NQ; i < TILE_SRC0_SHMEM; i += TOTAL_WORKGROUP_SIZE * NQ) {
221
- let blck_idx = i / BLOCK_SIZE;
222
- let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F16;
223
- let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * 2u;
194
+ // load NQ(16) weights
195
+ for (var j = 0u; j < BYTES_PER_THREAD / BYTES_PER_INNER_LOOP; j += 1) {
196
+ let q_byte_offset = block_byte_base + 8u + block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP;
197
+ let q_packed = load_u32_at_src0_aligned(q_byte_offset);
224
198
 
225
- let tile_m = blck_idx / BLOCKS_K;
226
- let global_m = offset_m + tile_m;
227
- let block_k = blck_idx % BLOCKS_K;
228
- let global_k = k_outer / BLOCK_SIZE + block_k;
229
-
230
- if (global_m < params.m && global_k < params.k / BLOCK_SIZE) {
231
- let src0_idx = batch_offset + global_m * params.stride_01 + global_k;
232
- let scale_idx = src0_idx * F16_PER_BLOCK;
233
-
234
- let d = src0[scale_idx];
235
- let m = src0[scale_idx + 1u];
236
- let qh0 = src0[scale_idx + 2u];
237
- let qh1 = src0[scale_idx + 3u];
238
- let qh_packed = bitcast<u32>(vec2(qh0, qh1));
239
-
240
- for (var j = 0u; j < 2; j++) {
241
-
242
- let q_0 = src0[scale_idx + 4u + block_offset + (j*2)];
243
- let q_1 = src0[scale_idx + 4u + block_offset + (j*2) + 1u];
244
-
245
- let q_packed = bitcast<u32>(vec2(q_0, q_1));
246
-
247
- let j_adjusted = j + (block_offset / 2u);
248
-
249
-
250
- for (var k = 0u; k < 4u; k++) {
199
+ for (var k = 0u; k < BYTES_PER_INNER_LOOP; k++) {
251
200
  let q_byte = get_byte(q_packed, k);
252
201
 
253
- let qh_hi = (qh_packed >> (j_adjusted * 4 + k + 12)) & 0x10;
254
- let q_hi = (f16(((q_byte >> 4) & 0xF) | qh_hi)) * d + m;
255
- let qh_lo = ((qh_packed >> (j_adjusted * 4 + k)) << 4) & 0x10;
256
- let q_lo = (f16((q_byte & 0xF) | qh_lo)) * d + m;
257
-
258
- shmem[shmem_idx + j * 4u + k] = q_lo; // store first weight
259
- shmem[shmem_idx + j * 4u + k + 16u] = q_hi; // store second weight
202
+ let byte_idx = block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP + k;
203
+ let qh_hi = (qh_packed >> (byte_idx + 12u)) & 0x10;
204
+ let q_hi = f16(((q_byte >> 4) & 0xF) | qh_hi) * d + m;
205
+ let qh_lo = ((qh_packed >> byte_idx) << 4) & 0x10;
206
+ let q_lo = f16((q_byte & 0xF) | qh_lo) * d + m;
207
+ shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k] = q_lo;
208
+ shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k + 16u] = q_hi;
260
209
  }
261
210
  }
262
- }
263
- }
264
- }
265
211
  #endif // INIT_SRC0_SHMEM_Q5_1
266
212
 
267
- #ifdef INIT_SRC0_SHMEM_Q8_0
268
- const BLOCK_SIZE = 32u;
269
- // the number of blocks per k-tile. Note that this currently only works if TILE_K is a multiple of BLOCK_SIZE, which may need to be rethought for larger quantized types.
270
- override BLOCKS_K = TILE_K/BLOCK_SIZE;
271
- const NQ = 16u;
272
- const F16_PER_BLOCK = 17u; // 1 scale + 16 in array of weights
273
- const WEIGHTS_PER_F16 = 2u; // 2 8-bit weights per f16
274
- const F16_PER_THREAD = NQ / WEIGHTS_PER_F16; // 8 f16s per thread
275
-
276
- fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
277
- for (var i = thread_id * NQ; i < TILE_SRC0_SHMEM; i += TOTAL_WORKGROUP_SIZE * NQ) {
278
- let blck_idx = i / BLOCK_SIZE;
279
- let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F16;
280
- let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * 2u;
213
+ #if defined(INIT_SRC0_SHMEM_Q8_0)
214
+ let block_byte_base = src0_idx * 34u; // BLOCK_SIZE_BYTES = 34u;
215
+ let d = load_f16_at_src0(block_byte_base);
281
216
 
282
- let tile_m = blck_idx / BLOCKS_K;
283
- let global_m = offset_m + tile_m;
284
- let block_k = blck_idx % BLOCKS_K;
285
- let global_k = k_outer / BLOCK_SIZE + block_k;
286
-
287
- if (global_m < params.m && global_k < params.k / BLOCK_SIZE) {
288
- let src0_idx = batch_offset + global_m * params.stride_01 + global_k;
289
- let scale_idx = src0_idx * F16_PER_BLOCK;
290
- let d = src0[scale_idx];
291
-
292
- for (var j = 0u; j < F16_PER_THREAD; j+=2) {
293
- let q_0 = src0[scale_idx + 1u + block_offset + j];
294
- let q_1 = src0[scale_idx + 1u + block_offset + j + 1];
217
+ // load NQ(16) weights
218
+ for (var j = 0u; j < BYTES_PER_THREAD / BYTES_PER_INNER_LOOP; j += 1) {
219
+ let q_byte_offset = block_byte_base + 2u + block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP;
220
+ let q_packed = load_u32_at_src0(q_byte_offset);
221
+ dequant_q8_0_packed_to_shmem(q_packed, d, shmem_idx + j * BYTES_PER_INNER_LOOP);
222
+ }
223
+ #endif // INIT_SRC0_SHMEM_Q8_0
295
224
 
296
- let q_packed = bitcast<u32>(vec2(q_0, q_1));
297
- for (var k = 0u; k < 4u; k++) {
225
+ #if defined(INIT_SRC0_SHMEM_Q8_1)
226
+ let block_byte_base = src0_idx * 36u; // BLOCK_SIZE_BYTES = 36u;
227
+ let dm = unpack2x16float(load_u32_at_src0_aligned(block_byte_base));
228
+ let d = f16(dm[0]);
229
+ let m = f16(dm[1]);
230
+
231
+ // load NQ(16) weights
232
+ for (var j = 0u; j < BYTES_PER_THREAD / BYTES_PER_INNER_LOOP; j += 1) {
233
+ let q_byte_offset = block_byte_base + 4u + block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP;
234
+ let q_packed = load_u32_at_src0(q_byte_offset);
235
+ for (var k = 0u; k < BYTES_PER_INNER_LOOP; k++) {
298
236
  let q_byte = get_byte_i32(q_packed, k);
237
+ let q_val = f16(q_byte) * d + m;
238
+ shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k] = q_val;
239
+ }
240
+ }
241
+ #endif // INIT_SRC0_SHMEM_Q8_1
299
242
 
300
- let q_val = f16(q_byte) * d;
301
- shmem[shmem_idx + j * 2 + k] = q_val;
243
+ #if defined(INIT_SRC0_SHMEM_MXFP4)
244
+ let block_byte_base = src0_idx * 17u; // BLOCK_SIZE_BYTES = 17u;
245
+ let eu8 = get_byte(load_u32_at_src0_aligned(block_byte_base), block_byte_base & 3u);
246
+ let e = ldexp(1.0, i32(eu8) - 128);
247
+
248
+ // load NQ(16) weights
249
+ for (var j = 0u; j < BYTES_PER_THREAD / BYTES_PER_INNER_LOOP; j += 1) {
250
+ let q_byte_offset = block_byte_base + 1u + block_offset * BYTES_PER_THREAD + j * BYTES_PER_INNER_LOOP;
251
+ let q_packed = load_u32_at_src0(q_byte_offset);
252
+ for (var k = 0u; k < BYTES_PER_INNER_LOOP; k++) {
253
+ let q_byte = get_byte(q_packed, k);
254
+ let q_hi = f32(kvalues_mxfp4[(q_byte >> 4) & 0xF]) * e;
255
+ let q_lo = f32(kvalues_mxfp4[q_byte & 0xF]) * e;
256
+ shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k] = f16(q_lo);
257
+ shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k + 16u] = f16(q_hi);
302
258
  }
303
259
  }
260
+ #endif // INIT_SRC0_SHMEM_MXFP4
304
261
  }
305
262
  }
306
263
  }
307
- #endif // INIT_SRC0_SHMEM_Q8_0
264
+ #endif // legacy-quants
308
265
 
309
- #ifdef INIT_SRC0_SHMEM_Q8_1
310
- const BLOCK_SIZE = 32u;
311
- // the number of blocks per k-tile. Note that this currently only works if TILE_K is a multiple of BLOCK_SIZE, which may need to be rethought for larger quantized types.
312
- override BLOCKS_K = TILE_K/BLOCK_SIZE;
266
+ #if defined(INIT_SRC0_SHMEM_NVFP4)
267
+ const BLOCK_SIZE = 64u;
268
+ const BLOCK_SIZE_BYTES = 36u;
269
+ const SUB_BLOCK_SIZE = 16u; // elements sharing one UE4M3 scale
313
270
  const NQ = 16u;
314
- const F16_PER_BLOCK = 18u; // 1 scale + 1 mean + 8 32-bit values in array of weights
315
- const WEIGHTS_PER_F16 = 2u; // 2 8-bit weights per f16
316
- const F16_PER_THREAD = NQ / WEIGHTS_PER_F16; // 8 f16s per thread, 2 threads per block
271
+ const BYTES_PER_THREAD = 8u;
272
+ const BYTES_PER_INNER_LOOP = 4u;
317
273
 
318
274
  fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
319
275
  for (var i = thread_id * NQ; i < TILE_SRC0_SHMEM; i += TOTAL_WORKGROUP_SIZE * NQ) {
320
- let blck_idx = i / BLOCK_SIZE;
321
- let block_offset = (i % BLOCK_SIZE) / WEIGHTS_PER_F16;
322
- let shmem_idx = blck_idx * BLOCK_SIZE + block_offset * 2u;
323
-
324
- let tile_m = blck_idx / BLOCKS_K;
276
+ let tile_m = i / TILE_K;
277
+ let tile_k_start = i % TILE_K;
325
278
  let global_m = offset_m + tile_m;
326
- let block_k = blck_idx % BLOCKS_K;
327
- let global_k = k_outer / BLOCK_SIZE + block_k;
279
+ let global_k_start = k_outer + tile_k_start;
328
280
 
329
- if (global_m < params.m && global_k < params.k / BLOCK_SIZE) {
330
- let src0_idx = batch_offset + global_m * params.stride_01 + global_k;
331
- let scale_idx = src0_idx * F16_PER_BLOCK;
332
- let d = src0[scale_idx];
333
- let m = src0[scale_idx + 1u];
281
+ if (global_m >= params.m) {
282
+ break;
283
+ }
334
284
 
335
- for (var j = 0u; j < F16_PER_THREAD; j+=2) {
336
- let q_0 = src0[scale_idx + 2u + block_offset + j];
337
- let q_1 = src0[scale_idx + 2u + block_offset + j + 1];
285
+ let block_k = global_k_start / BLOCK_SIZE;
286
+ let sub_block = (global_k_start % BLOCK_SIZE) / SUB_BLOCK_SIZE;
287
+ let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
338
288
 
339
- let q_packed = bitcast<u32>(vec2(q_0, q_1));
340
- for (var k = 0u; k < 4u; k++) {
341
- let q_byte = get_byte_i32(q_packed, k);
289
+ let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
290
+ let d_byte_base = block_byte_base;
291
+ let qs_byte_base = block_byte_base + 4u;
342
292
 
343
- let q_val = f16(q_byte) * d + m;
344
- shmem[shmem_idx + j * 2 + k] = q_val;
345
- }
293
+ let d = ue4m3_to_fp32(get_byte(load_u32_at_src0_aligned(d_byte_base), sub_block)) * 0.5;
294
+
295
+ for (var j = 0u; j < BYTES_PER_THREAD / BYTES_PER_INNER_LOOP; j++) {
296
+ let q_packed = load_u32_at_src0_aligned(qs_byte_base + sub_block * 8u + j * 4u);
297
+ for (var k = 0u; k < BYTES_PER_INNER_LOOP; k++) {
298
+ let q_byte = get_byte(q_packed, k);
299
+ shmem[i + j * BYTES_PER_INNER_LOOP + k] = f16(f32(kvalues_mxfp4[q_byte & 0xF]) * d);
300
+ shmem[i + j * BYTES_PER_INNER_LOOP + k + 8u] = f16(f32(kvalues_mxfp4[(q_byte >> 4) & 0xF]) * d);
346
301
  }
347
302
  }
348
303
  }
349
304
  }
350
- #endif // INIT_SRC0_SHMEM_Q8_1
305
+ #endif // INIT_SRC0_SHMEM_NVFP4
351
306
 
352
- #ifdef INIT_SRC0_SHMEM_Q2_K
307
+ // k-quants
308
+ #if defined(INIT_SRC0_SHMEM_Q2_K) || defined(INIT_SRC0_SHMEM_Q3_K) || defined(INIT_SRC0_SHMEM_Q4_K) || defined(INIT_SRC0_SHMEM_Q5_K) || defined(INIT_SRC0_SHMEM_Q6_K)
353
309
  const BLOCK_SIZE = 256u;
354
- const F16_PER_BLOCK = 42u;
310
+ const NQ = 4u;
311
+
312
+ fn store_shmem_kquants(val: vec4<f16>, idx: u32) {
313
+ shmem[idx] = val.x;
314
+ shmem[idx + 1] = val.y;
315
+ shmem[idx + 2] = val.z;
316
+ shmem[idx + 3] = val.w;
317
+ }
318
+
319
+ fn load_byte_at_src0_aligned(byte_offset: u32) -> u32 {
320
+ return get_byte(load_u32_at_src0_aligned(byte_offset), byte_offset % 4u);
321
+ }
355
322
 
356
323
  fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
357
- // Use standard thread layout instead of lane/row_group
358
- for (var elem_idx = thread_id; elem_idx < TILE_SRC0_SHMEM; elem_idx += TOTAL_WORKGROUP_SIZE) {
324
+ for (var elem_idx = thread_id * NQ; elem_idx < TILE_SRC0_SHMEM; elem_idx += TOTAL_WORKGROUP_SIZE * NQ) {
359
325
  let tile_m = elem_idx / TILE_K;
360
326
  let tile_k = elem_idx % TILE_K;
361
327
 
@@ -363,60 +329,250 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
363
329
  let global_k = k_outer + tile_k;
364
330
 
365
331
  if (global_m >= params.m || global_k >= params.k) {
366
- shmem[elem_idx] = f16(0.0);
332
+ store_shmem_kquants(vec4<f16>(f16(0.0), f16(0.0), f16(0.0), f16(0.0)), elem_idx);
367
333
  continue;
368
334
  }
369
335
 
370
- let block_k = global_k / BLOCK_SIZE;
371
- let k_in_block = global_k % BLOCK_SIZE;
336
+ let block_k = global_k / BLOCK_SIZE;
337
+ let k_in_block = global_k % BLOCK_SIZE; // k_in_block % 4 == 0;
372
338
 
373
339
  let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
374
- let scale_idx = src0_idx * F16_PER_BLOCK;
375
340
 
376
- let d = src0[scale_idx + 40u];
377
- let dmin = src0[scale_idx + 41u];
341
+ #if defined(INIT_SRC0_SHMEM_Q2_K)
342
+ let block_byte_base = src0_idx * 84u; // BLOCK_SIZE_BYTES = 84u;
343
+ let scales_byte_base = block_byte_base;
344
+ let qs_byte_base = block_byte_base + 16u;
345
+ let dm_byte_base = block_byte_base + 80u;
346
+
347
+ let d_packed = unpack2x16float(load_u32_at_src0_aligned(dm_byte_base));
348
+ let d = f16(d_packed[0]);
349
+ let dmin = f16(d_packed[1]);
350
+
351
+ let chunk = k_in_block / 128u;
352
+ let pos_in_chunk = k_in_block % 32u;
353
+ let sub_block = k_in_block / 16u;
354
+ let shift_phase = (k_in_block % 128u) / 32u;
355
+
356
+ // whole 2 bits (4 elems)
357
+ let qs_word = load_u32_at_src0_aligned(qs_byte_base + 32u * chunk + 1u * pos_in_chunk);
358
+ let qs_vec4 = vec4<f16>(
359
+ f16((qs_word >> (2u * shift_phase + 0u)) & 0x3u),
360
+ f16((qs_word >> (2u * shift_phase + 8u)) & 0x3u),
361
+ f16((qs_word >> (2u * shift_phase + 16u)) & 0x3u),
362
+ f16((qs_word >> (2u * shift_phase + 24u)) & 0x3u),
363
+ );
364
+
365
+ let scale = load_byte_at_src0_aligned(scales_byte_base + sub_block);
366
+
367
+ let dl = d * f16(scale & 0xFu);
368
+ let ml = dmin * f16(scale >> 4u);
369
+
370
+ store_shmem_kquants(qs_vec4 * dl - ml, elem_idx);
371
+ #endif // INIT_SRC0_SHMEM_Q2_K
372
+
373
+ #if defined(INIT_SRC0_SHMEM_Q3_K)
374
+ let block_byte_base = src0_idx * 110u; // BLOCK_SIZE_BYTES = 110u;
375
+ let hmask_byte_base = block_byte_base + 0u;
376
+ let qs_byte_base = block_byte_base + 32u;
377
+ let scales_byte_base = block_byte_base + 96u;
378
+
379
+ let d_all = load_f16_at_src0(block_byte_base + 108u);
380
+
381
+ let chunk = k_in_block / 128u;
382
+ let pos_in_chunk = k_in_block % 32u;
383
+ let sub_block = k_in_block / 16u;
384
+ let shift_phase = (k_in_block % 128u) / 32u;
385
+
386
+ let hmask_block = pos_in_chunk;
387
+ let hmask_shift_phase = k_in_block / 32u;
388
+
389
+ // low 2 bits (4 elems)
390
+ let q_lo2_word = load_u32_at_src0(qs_byte_base + 32u * chunk + 1u * hmask_block);
391
+ let q_lo2_vec4 = vec4<f16>(
392
+ f16((q_lo2_word >> (2u * shift_phase + 0u)) & 3u),
393
+ f16((q_lo2_word >> (2u * shift_phase + 8u)) & 3u),
394
+ f16((q_lo2_word >> (2u * shift_phase + 16u)) & 3u),
395
+ f16((q_lo2_word >> (2u * shift_phase + 24u)) & 3u)
396
+ );
397
+
398
+ // high 1 bit (4 elems)
399
+ let q_hi1_word = load_u32_at_src0(hmask_byte_base + pos_in_chunk);
400
+ let q_hi1_vec4 = vec4<f16>(
401
+ f16(select(4.0, 0.0, ((q_hi1_word >> (1u * hmask_shift_phase + 0u)) & 1u) == 1u)),
402
+ f16(select(4.0, 0.0, ((q_hi1_word >> (1u * hmask_shift_phase + 8u)) & 1u) == 1u)),
403
+ f16(select(4.0, 0.0, ((q_hi1_word >> (1u * hmask_shift_phase + 16u)) & 1u) == 1u)),
404
+ f16(select(4.0, 0.0, ((q_hi1_word >> (1u * hmask_shift_phase + 24u)) & 1u) == 1u))
405
+ );
406
+
407
+ let q_vec4 = q_lo2_vec4 - q_hi1_vec4;
408
+
409
+ let scale_low4 = (load_byte_at_src0_aligned(scales_byte_base + (sub_block % 8u)) >> (4u * (sub_block / 8u))) & 0xFu;
410
+ let scale_hi2 = (load_byte_at_src0_aligned(scales_byte_base + 8u + (sub_block % 4u)) >> (2u * (sub_block / 4u))) & 3u;
411
+ let dl = d_all * (f16((scale_hi2 << 4u) | scale_low4) - 32.0);
412
+
413
+ store_shmem_kquants(dl * q_vec4, elem_idx);
414
+ #endif // INIT_SRC0_SHMEM_Q3_K
415
+
416
+ #if defined(INIT_SRC0_SHMEM_Q4_K)
417
+ let block_byte_base = src0_idx * 144u; // BLOCK_SIZE_BYTES = 144u;
418
+ let dm_byte_base = block_byte_base + 0u;
419
+ let scale_byte_base = block_byte_base + 4u;
420
+ let qs_byte_base = block_byte_base + 16u;
421
+
422
+ let dm = unpack2x16float(load_u32_at_src0_aligned(dm_byte_base));
423
+ let d = f16(dm[0]);
424
+ let dmin = f16(dm[1]);
425
+
426
+ let chunk = k_in_block / 64u;
427
+ let pos_in_chunk = (k_in_block % 64u) % 32u;
428
+ let sub_block = k_in_block / 32u;
429
+ let shift_phase = sub_block & 1u;
430
+
431
+ // whole 4 bits (4 elems)
432
+ let qs_word = load_u32_at_src0_aligned(qs_byte_base + 32u * chunk + 1u * pos_in_chunk);
433
+ let qs_vec4 = vec4<f16>(
434
+ f16((qs_word >> (4u * shift_phase + 0u)) & 0xFu),
435
+ f16((qs_word >> (4u * shift_phase + 8u)) & 0xFu),
436
+ f16((qs_word >> (4u * shift_phase + 16u)) & 0xFu),
437
+ f16((qs_word >> (4u * shift_phase + 24u)) & 0xFu)
438
+ );
439
+
440
+ var sc: u32;
441
+ var mn: u32;
442
+
443
+ if (sub_block < 4u) {
444
+ let sc_byte = get_byte(load_u32_at_src0_aligned(scale_byte_base), sub_block % 4u);
445
+ let min_byte = get_byte(load_u32_at_src0_aligned(scale_byte_base + 4), sub_block % 4u);
446
+ sc = sc_byte & 63u;
447
+ mn = min_byte & 63u;
448
+ } else {
449
+ let sc_min_lo = get_byte(load_u32_at_src0_aligned(scale_byte_base + 8), (sub_block + 4u) % 4u);
450
+ let sc_hi = get_byte(load_u32_at_src0_aligned(scale_byte_base), (sub_block - 4u) % 4u);
451
+ let min_hi = get_byte(load_u32_at_src0_aligned(scale_byte_base + 4), sub_block % 4u);
452
+ sc = (sc_min_lo & 0xFu) | ((sc_hi >> 6u) << 4u);
453
+ mn = (sc_min_lo >> 4u) | ((min_hi >> 6u) << 4u);
454
+ }
455
+
456
+ let dl = d * f16(sc);
457
+ let ml = dmin * f16(mn);
378
458
 
379
- // Decode the element at position k_in_block
380
- let block_of_32 = k_in_block / 32u;
381
- let pos_in_32 = k_in_block % 32u;
459
+ store_shmem_kquants(dl * qs_vec4 - vec4(ml, ml, ml, ml), elem_idx);
460
+ #endif // INIT_SRC0_SHMEM_Q4_K
382
461
 
383
- let q_b_idx = (block_of_32 / 4u) * 32u;
384
- let shift = (block_of_32 % 4u) * 2u;
385
- let k = (pos_in_32 / 16u) * 16u;
386
- let l = pos_in_32 % 16u;
462
+ #if defined(INIT_SRC0_SHMEM_Q5_K)
463
+ let block_byte_base = src0_idx * 176u; // BLOCK_SIZE_BYTES = 176u;
464
+ let dm_byte_base = block_byte_base + 0u;
465
+ let scale_byte_base = block_byte_base + 4u;
466
+ let qh_byte_base = block_byte_base + 16u;
467
+ let qs_byte_base = block_byte_base + 48u;
468
+
469
+ let dm = unpack2x16float(load_u32_at_src0_aligned(dm_byte_base));
470
+ let d = f16(dm[0]);
471
+ let dmin = f16(dm[1]);
472
+
473
+ let chunk = k_in_block / 64u;
474
+ let pos_in_chunk = (k_in_block % 64u) % 32u;
475
+ let sub_block = k_in_block / 32u;
476
+ let shift_phase = sub_block & 1u;
477
+
478
+ let qh_block = k_in_block % 32u;
479
+ let qh_shift_phase = sub_block;
480
+
481
+ // low 4 bits (4 elems)
482
+ let qs_word = load_u32_at_src0_aligned(qs_byte_base + 32u * chunk + 1u * pos_in_chunk);
483
+ let qs_lo4_vec4 = vec4<f16>(
484
+ f16((qs_word >> (4u * shift_phase + 0u)) & 0xFu),
485
+ f16((qs_word >> (4u * shift_phase + 8u)) & 0xFu),
486
+ f16((qs_word >> (4u * shift_phase + 16u)) & 0xFu),
487
+ f16((qs_word >> (4u * shift_phase + 24u)) & 0xFu)
488
+ );
489
+
490
+ // high 1 bit (4 elems)
491
+ let qh_word = load_u32_at_src0_aligned(qh_byte_base + qh_block);
492
+ let qh_vec4 = vec4<f16>(
493
+ f16(select(0.0, 16.0, ((qh_word >> (1u * qh_shift_phase + 0u)) & 1u) == 1u)),
494
+ f16(select(0.0, 16.0, ((qh_word >> (1u * qh_shift_phase + 8u)) & 1u) == 1u)),
495
+ f16(select(0.0, 16.0, ((qh_word >> (1u * qh_shift_phase + 16u)) & 1u) == 1u)),
496
+ f16(select(0.0, 16.0, ((qh_word >> (1u * qh_shift_phase + 24u)) & 1u) == 1u))
497
+ );
387
498
 
388
- let is = k_in_block / 16u;
499
+ var sc: u32;
500
+ var mn: u32;
389
501
 
390
- let sc_0 = src0[scale_idx + 2u * (is / 4u)];
391
- let sc_1 = src0[scale_idx + 2u * (is / 4u) + 1u];
392
- let sc_packed = bitcast<u32>(vec2(sc_0, sc_1));
393
- let sc = get_byte(sc_packed, is % 4u);
502
+ if (sub_block < 4u) {
503
+ let sc_byte = get_byte(load_u32_at_src0_aligned(scale_byte_base), sub_block % 4u);
504
+ let min_byte = get_byte(load_u32_at_src0_aligned(scale_byte_base + 4), sub_block % 4u);
505
+ sc = sc_byte & 63u;
506
+ mn = min_byte & 63u;
507
+ } else {
508
+ let sc_min_lo = get_byte(load_u32_at_src0_aligned(scale_byte_base + 8), (sub_block + 4u) % 4u);
509
+ let sc_hi = get_byte(load_u32_at_src0_aligned(scale_byte_base), (sub_block - 4u) % 4u);
510
+ let min_hi = get_byte(load_u32_at_src0_aligned(scale_byte_base + 4), sub_block % 4u);
511
+ sc = (sc_min_lo & 0xFu) | ((sc_hi >> 6u) << 4u);
512
+ mn = (sc_min_lo >> 4u) | ((min_hi >> 6u) << 4u);
513
+ }
394
514
 
395
- let dl = d * f16(sc & 0xFu);
396
- let ml = dmin * f16(sc >> 4u);
515
+ let dl = d * f16(sc);
516
+ let ml = dmin * f16(mn);
397
517
 
398
- let q_idx = q_b_idx + k + l;
399
- let q_0 = src0[scale_idx + 8u + 2u * (q_idx / 4u)];
400
- let q_1 = src0[scale_idx + 8u + 2u * (q_idx / 4u) + 1u];
401
- let q_packed = bitcast<u32>(vec2(q_0, q_1));
402
- let q_byte = get_byte(q_packed, q_idx % 4u);
403
- let qs_val = (q_byte >> shift) & 3u;
518
+ store_shmem_kquants((qh_vec4 + qs_lo4_vec4) * dl - vec4<f16>(ml, ml, ml, ml), elem_idx);
519
+ #endif // INIT_SRC0_SHMEM_Q5_K
404
520
 
405
- let q_val = f16(qs_val) * dl - ml;
406
- shmem[elem_idx] = q_val;
521
+ #if defined(INIT_SRC0_SHMEM_Q6_K)
522
+ let block_byte_base = src0_idx * 210u; // BLOCK_SIZE_BYTES = 210u;
523
+ let ql_byte_base = block_byte_base;
524
+ let qh_byte_base = block_byte_base + 128u;
525
+ let scales_byte_base = block_byte_base + 192u;
526
+ let d_byte_base = block_byte_base + 208u;
527
+
528
+ let d = load_f16_at_src0(d_byte_base);
529
+
530
+ let chunk = k_in_block / 128u;
531
+ let ql_pos_in_chunk = (k_in_block % 128u) % 64u;
532
+ let qh_pos_in_chunk = (k_in_block % 128u) % 32u;
533
+ let sub_block = k_in_block / 16u;
534
+ let ql_shift_phase = (k_in_block % 128u) / 64u;
535
+ let qh_shift_phase = (k_in_block % 128u) / 32u;
536
+
537
+ // low 4 bits (4 elems)
538
+ let ql_word = load_u32_at_src0(ql_byte_base + 64u * chunk + 1u * ql_pos_in_chunk);
539
+ let ql_lo4_vec4 = vec4<u32>(
540
+ (ql_word >> (4u * ql_shift_phase + 0u)) & 0xFu,
541
+ (ql_word >> (4u * ql_shift_phase + 8u)) & 0xFu,
542
+ (ql_word >> (4u * ql_shift_phase + 16u)) & 0xFu,
543
+ (ql_word >> (4u * ql_shift_phase + 24u)) & 0xFu
544
+ );
545
+
546
+ // hi 2 bits (4 elems)
547
+ let qh_word = load_u32_at_src0(qh_byte_base + 32u * chunk + 1u * qh_pos_in_chunk);
548
+ let qh_hi2_vec4 = vec4<u32>(
549
+ ((qh_word >> (2u * qh_shift_phase + 0u)) & 0x3u) << 4u,
550
+ ((qh_word >> (2u * qh_shift_phase + 8u)) & 0x3u) << 4u,
551
+ ((qh_word >> (2u * qh_shift_phase + 16u)) & 0x3u) << 4u,
552
+ ((qh_word >> (2u * qh_shift_phase + 24u)) & 0x3u) << 4u,
553
+ );
554
+
555
+ let q_vec4 = vec4<f16>(qh_hi2_vec4 | ql_lo4_vec4) - vec4<f16>(32.0, 32.0, 32.0, 32.0);
556
+
557
+ let scale_byte = scales_byte_base + 1u * sub_block;
558
+ let scale_word = load_u32_at_src0_aligned(scale_byte);
559
+ let scale = get_byte_i32(scale_word, scale_byte & 3u);
560
+
561
+ store_shmem_kquants(d * q_vec4 * f16(scale), elem_idx);
562
+ #endif // INIT_SRC0_SHMEM_Q6_K
407
563
  }
408
564
  }
409
- #endif // INIT_SRC0_SHMEM_Q2_K
565
+ #endif // k-quants
410
566
 
411
- #ifdef INIT_SRC0_SHMEM_Q3_K
412
- const BLOCK_SIZE = 256u;
413
- const F16_PER_BLOCK = 55u;
567
+ #if defined(INIT_SRC0_SHMEM_IQ4_NL)
568
+ const BLOCK_SIZE = 32u;
569
+ const BLOCK_SIZE_BYTES = 18u;
570
+ const NQ = 4u;
414
571
 
415
572
  fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
416
- for (var elem_idx = thread_id; elem_idx < TILE_SRC0_SHMEM; elem_idx += TOTAL_WORKGROUP_SIZE) {
573
+ for (var elem_idx = thread_id * NQ; elem_idx < TILE_SRC0_SHMEM; elem_idx += NQ * TOTAL_WORKGROUP_SIZE) {
417
574
  let tile_m = elem_idx / TILE_K;
418
575
  let tile_k = elem_idx % TILE_K;
419
-
420
576
  let global_m = offset_m + tile_m;
421
577
  let global_k = k_outer + tile_k;
422
578
 
@@ -425,342 +581,465 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
425
581
  continue;
426
582
  }
427
583
 
428
- let block_k = global_k / BLOCK_SIZE;
429
- let k_in_block = global_k % BLOCK_SIZE;
584
+ let block_k = global_k / BLOCK_SIZE;
585
+ let k_in_block = global_k % BLOCK_SIZE; // k_in_block % 4 == 0;
430
586
 
431
587
  let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
432
- let scale_idx = src0_idx * F16_PER_BLOCK;
433
-
434
- let d = src0[scale_idx + 54u];
435
-
436
- // Load and unpack scales
437
- let kmask1: u32 = 0x03030303u;
438
- let kmask2: u32 = 0x0f0f0f0fu;
439
-
440
- var scale_vals: array<u32, 4>;
441
- for (var i: u32 = 0u; i < 4u; i++) {
442
- let scale_0 = src0[scale_idx + 48u + (2u*i)];
443
- let scale_1 = src0[scale_idx + 48u + (2u*i) + 1u];
444
- scale_vals[i] = bitcast<u32>(vec2(scale_0, scale_1));
445
- }
446
588
 
447
- var tmp: u32 = scale_vals[2];
448
- scale_vals[2] = ((scale_vals[0] >> 4u) & kmask2) | (((tmp >> 4u) & kmask1) << 4u);
449
- scale_vals[3] = ((scale_vals[1] >> 4u) & kmask2) | (((tmp >> 6u) & kmask1) << 4u);
450
- scale_vals[0] = (scale_vals[0] & kmask2) | ((tmp & kmask1) << 4u);
451
- scale_vals[1] = (scale_vals[1] & kmask2) | (((tmp >> 2u) & kmask1) << 4u);
452
-
453
- // Load hmask and qs arrays
454
- var hmask_vals: array<u32, 8>;
455
- for (var i: u32 = 0u; i < 8u; i++) {
456
- let hmask_0 = src0[scale_idx + (2u*i)];
457
- let hmask_1 = src0[scale_idx + (2u*i) + 1u];
458
- hmask_vals[i] = bitcast<u32>(vec2(hmask_0, hmask_1));
459
- }
460
-
461
- var qs_vals: array<u32, 16>;
462
- for (var i: u32 = 0u; i < 16u; i++) {
463
- let qs_0 = src0[scale_idx + 16u + (2u*i)];
464
- let qs_1 = src0[scale_idx + 16u + (2u*i) + 1u];
465
- qs_vals[i] = bitcast<u32>(vec2(qs_0, qs_1));
466
- }
467
-
468
- let half = k_in_block / 128u; // 0 or 1
469
- let pos_in_half = k_in_block % 128u; // 0-127
470
- let shift_group = pos_in_half / 32u; // 0-3
471
- let pos_in_32 = pos_in_half % 32u; // 0-31
472
- let k_group = pos_in_32 / 16u; // 0 or 1
473
- let l = pos_in_32 % 16u; // 0-15
589
+ let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
590
+ let d_byte_base = block_byte_base + 0u;
591
+ let qs_byte_base = block_byte_base + 2u;
474
592
 
475
- let q_b_idx = half * 32u; // 0 or 32
476
- let shift = shift_group * 2u; // 0, 2, 4, 6
477
- let k = k_group * 16u; // 0 or 16
478
- let is = k_in_block / 16u; // 0-15
593
+ let d = load_f16_at_src0(d_byte_base);
479
594
 
480
- // m increments every 32 elements across entire 256 element block
481
- let m_shift = k_in_block / 32u; // 0-7
482
- let m: u32 = 1u << m_shift; // 1,2,4,8,16,32,64,128
595
+ let id_qtr = (k_in_block % 16u) / 4u;
596
+ let shift_phase = k_in_block / 16u;
483
597
 
484
- let sc = get_byte(scale_vals[is / 4u], is % 4u);
485
- let dl = d * (f16(sc) - 32.0);
598
+ let qs_u32 = load_u32_at_src0(qs_byte_base + 4u * id_qtr);
486
599
 
487
- let q_idx = q_b_idx + k + l;
488
- let hm_idx = k + l;
489
-
490
- let q_byte = get_byte(qs_vals[q_idx / 4u], q_idx % 4u);
491
- let hmask_byte = get_byte(hmask_vals[hm_idx / 4u], hm_idx % 4u);
600
+ shmem[elem_idx + 0u] = d * f16(kvalues_iq4nl[(qs_u32 >> ( 0u + 4u * shift_phase)) & 0xFu]);
601
+ shmem[elem_idx + 1u] = d * f16(kvalues_iq4nl[(qs_u32 >> ( 8u + 4u * shift_phase)) & 0xFu]);
602
+ shmem[elem_idx + 2u] = d * f16(kvalues_iq4nl[(qs_u32 >> (16u + 4u * shift_phase)) & 0xFu]);
603
+ shmem[elem_idx + 3u] = d * f16(kvalues_iq4nl[(qs_u32 >> (24u + 4u * shift_phase)) & 0xFu]);
604
+ }
605
+ }
606
+ #endif // INIT_SRC0_SHMEM_IQ4_NL
492
607
 
493
- let hm = select(4.0, 0.0, (hmask_byte & m) != 0);
494
- let qs_val = (q_byte >> shift) & 3u;
608
+ // i-quants (super block size: 256)
609
+ #if defined(INIT_SRC0_SHMEM_IQ4_XS) || defined(INIT_SRC0_SHMEM_IQ1_S) || defined(INIT_SRC0_SHMEM_IQ1_M) || defined(INIT_SRC0_SHMEM_IQ2_XXS) \
610
+ || defined(INIT_SRC0_SHMEM_IQ2_XS) || defined(INIT_SRC0_SHMEM_IQ2_S) || defined(INIT_SRC0_SHMEM_IQ3_XXS) || defined(INIT_SRC0_SHMEM_IQ3_S)
611
+ const BLOCK_SIZE = 256u;
612
+ const NQ = 16u;
495
613
 
496
- let q_val = (f16(qs_val) - f16(hm)) * dl;
497
- shmem[elem_idx] = q_val;
498
- }
614
+ fn store_shmem_iquants(val: vec4<f16>, idx: u32) {
615
+ shmem[idx] = val.x;
616
+ shmem[idx + 1] = val.y;
617
+ shmem[idx + 2] = val.z;
618
+ shmem[idx + 3] = val.w;
499
619
  }
500
620
 
501
- #endif // INIT_SRC0_SHMEM_Q3_K
621
+ fn load_byte_at_src0_aligned(byte_offset: u32) -> u32 {
622
+ return get_byte(load_u32_at_src0_aligned(byte_offset), byte_offset % 4u);
623
+ }
502
624
 
503
- #ifdef INIT_SRC0_SHMEM_Q4_K
504
- const BLOCK_SIZE = 256u;
505
- const F16_PER_BLOCK = 72u;
625
+ #if defined(INIT_SRC0_SHMEM_IQ1_M) || defined(INIT_SRC0_SHMEM_IQ1_S)
626
+ fn create_iq_gw4(dl: f32, gw: u32, shift_base: u32, delta: f32) -> vec4<f16> {
627
+ return vec4<f16>(
628
+ f16(dl * (f32((bitcast<i32>(((gw >> (shift_base + 0u)) & 3u) << 30u) >> 30u)) + delta)),
629
+ f16(dl * (f32((bitcast<i32>(((gw >> (shift_base + 2u)) & 3u) << 30u) >> 30u)) + delta)),
630
+ f16(dl * (f32((bitcast<i32>(((gw >> (shift_base + 4u)) & 3u) << 30u) >> 30u)) + delta)),
631
+ f16(dl * (f32((bitcast<i32>(((gw >> (shift_base + 6u)) & 3u) << 30u) >> 30u)) + delta)),
632
+ );
633
+ }
634
+ #endif
635
+
636
+ #if defined(INIT_SRC0_SHMEM_IQ4_XS)
637
+ fn create_iq_gw4(dl: f16, qs_u32: u32, shift_phase: u32) -> vec4<f16> {
638
+ return vec4<f16>(
639
+ dl * f16(kvalues_iq4nl[(qs_u32 >> (4 * shift_phase + 0u)) & 0xFu]),
640
+ dl * f16(kvalues_iq4nl[(qs_u32 >> (4 * shift_phase + 8u)) & 0xFu]),
641
+ dl * f16(kvalues_iq4nl[(qs_u32 >> (4 * shift_phase + 16u)) & 0xFu]),
642
+ dl * f16(kvalues_iq4nl[(qs_u32 >> (4 * shift_phase + 24u)) & 0xFu]),
643
+ );
644
+ }
645
+ #endif
646
+
647
+ #if defined(INIT_SRC0_SHMEM_IQ2_XXS)
648
+ fn create_iq_gw4(ig: u32, grid_phase: u32) -> vec4<f32> {
649
+ return vec4<f32>(
650
+ f32(get_byte(iq2xxs_grid[(ig + grid_phase + 0u) / 4u], (ig + grid_phase + 0u) % 4u)),
651
+ f32(get_byte(iq2xxs_grid[(ig + grid_phase + 1u) / 4u], (ig + grid_phase + 1u) % 4u)),
652
+ f32(get_byte(iq2xxs_grid[(ig + grid_phase + 2u) / 4u], (ig + grid_phase + 2u) % 4u)),
653
+ f32(get_byte(iq2xxs_grid[(ig + grid_phase + 3u) / 4u], (ig + grid_phase + 3u) % 4u)),
654
+ );
655
+ }
656
+ #endif
657
+
658
+ #if defined(INIT_SRC0_SHMEM_IQ2_XS)
659
+ fn create_iq_gw4(ig: u32, grid_phase: u32) -> vec4<f32> {
660
+ return vec4<f32>(
661
+ f32(get_byte(iq2xs_grid[(ig + grid_phase + 0u) / 4u], (ig + grid_phase + 0u) % 4u)),
662
+ f32(get_byte(iq2xs_grid[(ig + grid_phase + 1u) / 4u], (ig + grid_phase + 1u) % 4u)),
663
+ f32(get_byte(iq2xs_grid[(ig + grid_phase + 2u) / 4u], (ig + grid_phase + 2u) % 4u)),
664
+ f32(get_byte(iq2xs_grid[(ig + grid_phase + 3u) / 4u], (ig + grid_phase + 3u) % 4u)),
665
+ );
666
+ }
667
+ #endif
668
+
669
+ #if defined(INIT_SRC0_SHMEM_IQ2_S)
670
+ fn create_iq_gw4(ig: u32, grid_phase: u32) -> vec4<f32> {
671
+ return vec4<f32>(
672
+ f32(get_byte(iq2s_grid[(ig + grid_phase + 0u) / 4u], (ig + grid_phase + 0u) % 4u)),
673
+ f32(get_byte(iq2s_grid[(ig + grid_phase + 1u) / 4u], (ig + grid_phase + 1u) % 4u)),
674
+ f32(get_byte(iq2s_grid[(ig + grid_phase + 2u) / 4u], (ig + grid_phase + 2u) % 4u)),
675
+ f32(get_byte(iq2s_grid[(ig + grid_phase + 3u) / 4u], (ig + grid_phase + 3u) % 4u)),
676
+ );
677
+ }
678
+ #endif
679
+
680
+ #if defined(INIT_SRC0_SHMEM_IQ3_XXS)
681
+ fn create_iq_gw4(ig: u32) -> vec4<f32> {
682
+ return vec4<f32>(
683
+ f32(get_byte(iq3xxs_grid[ig], 0)),
684
+ f32(get_byte(iq3xxs_grid[ig], 1)),
685
+ f32(get_byte(iq3xxs_grid[ig], 2)),
686
+ f32(get_byte(iq3xxs_grid[ig], 3)),
687
+ );
688
+ }
689
+ #endif
690
+
691
+ #if defined(INIT_SRC0_SHMEM_IQ3_S)
692
+ fn create_iq_gw4(ig: u32) -> vec4<f32> {
693
+ return vec4<f32>(
694
+ f32(get_byte(iq3s_grid[ig], 0)),
695
+ f32(get_byte(iq3s_grid[ig], 1)),
696
+ f32(get_byte(iq3s_grid[ig], 2)),
697
+ f32(get_byte(iq3s_grid[ig], 3)),
698
+ );
699
+ }
700
+ #endif
701
+
702
+ #if defined(INIT_SRC0_SHMEM_IQ2_XXS) || defined(INIT_SRC0_SHMEM_IQ2_XS) || defined(INIT_SRC0_SHMEM_IQ2_S) \
703
+ || defined(INIT_SRC0_SHMEM_IQ3_XXS) || defined(INIT_SRC0_SHMEM_IQ3_S)
704
+ fn create_iq2_m4(signs: u32, mask_phase: u32) -> vec4<f32> {
705
+ return vec4<f32>(
706
+ select(1.0, -1.0, (get_byte(kmask_iq2xs[mask_phase], 0) & signs) != 0u),
707
+ select(1.0, -1.0, (get_byte(kmask_iq2xs[mask_phase], 1) & signs) != 0u),
708
+ select(1.0, -1.0, (get_byte(kmask_iq2xs[mask_phase], 2) & signs) != 0u),
709
+ select(1.0, -1.0, (get_byte(kmask_iq2xs[mask_phase], 3) & signs) != 0u),
710
+ );
711
+ }
712
+ #endif
506
713
 
507
714
  fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
508
- for (var elem_idx = thread_id; elem_idx < TILE_SRC0_SHMEM; elem_idx += TOTAL_WORKGROUP_SIZE) {
715
+ for (var elem_idx = thread_id * NQ; elem_idx < TILE_SRC0_SHMEM; elem_idx += NQ * TOTAL_WORKGROUP_SIZE) {
509
716
  let tile_m = elem_idx / TILE_K;
510
717
  let tile_k = elem_idx % TILE_K;
511
-
512
718
  let global_m = offset_m + tile_m;
513
719
  let global_k = k_outer + tile_k;
514
720
 
515
721
  if (global_m >= params.m || global_k >= params.k) {
516
- shmem[elem_idx] = f16(0.0);
722
+ let zero_vec4 = vec4<f16>(f16(0.0), f16(0.0), f16(0.0), f16(0.0));
723
+ store_shmem_iquants(zero_vec4, elem_idx + 0u);
724
+ store_shmem_iquants(zero_vec4, elem_idx + 4u);
725
+ store_shmem_iquants(zero_vec4, elem_idx + 8u);
726
+ store_shmem_iquants(zero_vec4, elem_idx + 12u);
517
727
  continue;
518
728
  }
519
729
 
520
- let block_k = global_k / BLOCK_SIZE;
521
- let k_in_block = global_k % BLOCK_SIZE;
730
+ let block_k = global_k / BLOCK_SIZE;
731
+ let k_in_block = global_k % BLOCK_SIZE; // k_in_block % 16 == 0;
522
732
 
523
733
  let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
524
- let scale_idx = src0_idx * F16_PER_BLOCK;
525
734
 
526
- let d = src0[scale_idx];
527
- let dmin = src0[scale_idx + 1u];
735
+ #if defined(INIT_SRC0_SHMEM_IQ4_XS)
736
+ let block_byte_base = src0_idx * 136u; // BLOCK_SIZE_BYTES = 136u;
737
+ let d_byte_base = block_byte_base + 0u;
738
+ let scales_l_byte_base = block_byte_base + 4u;
739
+ let qs_byte_base = block_byte_base + 8u;
528
740
 
529
- // Load packed scales
530
- var scale_vals: array<u32, 3>;
531
- for (var i: u32 = 0u; i < 3u; i++) {
532
- let scale_0 = src0[scale_idx + 2u + (2u*i)];
533
- let scale_1 = src0[scale_idx + 2u + (2u*i) + 1u];
534
- scale_vals[i] = bitcast<u32>(vec2(scale_0, scale_1));
535
- }
741
+ let d_scales_h = load_u32_at_src0_aligned(d_byte_base);
742
+ let d = bitcast<vec2<f16>>(d_scales_h).x;
743
+ let scales_h = d_scales_h >> 16u;
536
744
 
537
- // Map k_in_block to loop structure:
538
- // Outer loop over 64-element groups (alternating q_b_idx)
539
- // Inner loop over 2 shifts per group
540
- let group_of_64 = k_in_block / 64u; // 0-3 (maps to q_b_idx)
541
- let pos_in_64 = k_in_block % 64u; // 0-63
542
- let shift_group = pos_in_64 / 32u; // 0 or 1
543
- let l = pos_in_64 % 32u; // 0-31
745
+ let sub_block = k_in_block / 32u;
746
+ let phase = (k_in_block / NQ) % 2u;
544
747
 
545
- let q_b_idx = group_of_64 * 32u; // 0, 32, 64, 96
546
- let shift = shift_group * 4u; // 0 or 4
547
- let is = k_in_block / 32u; // 0-7
748
+ let scales_l_u32 = load_u32_at_src0_aligned(scales_l_byte_base);
749
+ let ls_lo = (get_byte(scales_l_u32, sub_block / 2u) >> (4u * (sub_block % 2u))) & 0xFu;
750
+ let ls_hi = ((scales_h >> (2u * sub_block)) & 3u) << 4u;
751
+ let dl = d * f16(i32(ls_lo | ls_hi) - 32);
548
752
 
549
- var sc: u32;
550
- var mn: u32;
753
+ let qs_0_3_u32 = load_u32_at_src0_aligned(qs_byte_base + 16u * sub_block + 0u);
754
+ let qs_4_7_u32 = load_u32_at_src0_aligned(qs_byte_base + 16u * sub_block + 4u);
755
+ let qs_8_11_u32 = load_u32_at_src0_aligned(qs_byte_base + 16u * sub_block + 8u);
756
+ let qs_12_15_u32 = load_u32_at_src0_aligned(qs_byte_base + 16u * sub_block + 12u);
551
757
 
552
- if (is < 4u) {
553
- let sc_byte = get_byte(scale_vals[is / 4u], is % 4u);
554
- let min_byte = get_byte(scale_vals[(is + 4u) / 4u], is % 4u);
555
- sc = sc_byte & 63u;
556
- mn = min_byte & 63u;
557
- } else {
558
- let sc_min_lo = get_byte(scale_vals[(is + 4u) / 4u], (is + 4u) % 4u);
559
- let sc_hi = get_byte(scale_vals[(is - 4u) / 4u], (is - 4u) % 4u);
560
- let min_hi = get_byte(scale_vals[is / 4u], is % 4u);
758
+ store_shmem_iquants(create_iq_gw4(dl, qs_0_3_u32, phase), elem_idx + 0u);
759
+ store_shmem_iquants(create_iq_gw4(dl, qs_4_7_u32, phase), elem_idx + 4u);
760
+ store_shmem_iquants(create_iq_gw4(dl, qs_8_11_u32, phase), elem_idx + 8u);
761
+ store_shmem_iquants(create_iq_gw4(dl, qs_12_15_u32, phase), elem_idx + 12u);
762
+ #endif // INIT_SRC0_SHMEM_IQ4_XS
561
763
 
562
- sc = (sc_min_lo & 0xFu) | ((sc_hi >> 6u) << 4u);
563
- mn = (sc_min_lo >> 4u) | ((min_hi >> 6u) << 4u);
564
- }
764
+ #if defined(INIT_SRC0_SHMEM_IQ1_S)
765
+ let block_byte_base = src0_idx * 50u; // BLOCK_SIZE_BYTES = 50u;
766
+ let d_byte_base = block_byte_base + 0u;
767
+ let qs_byte_base = block_byte_base + 2u;
768
+ let qh_byte_base = block_byte_base + 34u;
565
769
 
566
- let dl = d * f16(sc);
567
- let ml = dmin * f16(mn);
770
+ let d = load_f16_as_f32_at_src0(d_byte_base);
568
771
 
569
- let q_idx = q_b_idx + l;
570
- let q_0 = src0[scale_idx + 8u + 2u * (q_idx / 4u)];
571
- let q_1 = src0[scale_idx + 8u + 2u * (q_idx / 4u) + 1u];
572
- let q_packed = bitcast<u32>(vec2(q_0, q_1));
772
+ let sub_block = k_in_block / 32u;
773
+ let phase = (k_in_block / NQ) % 2u;
573
774
 
574
- let q_byte = get_byte(q_packed, q_idx % 4u);
575
- let qs_val = (q_byte >> shift) & 0xFu;
775
+ let qh_u16 = load_u32_at_src0(qh_byte_base + sub_block * 2u) & 0xFFFFu;
776
+ let qs_u16 = load_u32_at_src0(qs_byte_base + sub_block * 4u + phase * 2u) & 0xFFFFu;
576
777
 
577
- let q_val = f16(qs_val) * dl - ml;
578
- shmem[elem_idx] = q_val;
579
- }
580
- }
581
- #endif // INIT_SRC0_SHMEM_Q4_K
778
+ let dl = d * (2.0 * f32((qh_u16 >> 12u) & 7u) + 1.0);
779
+ let delta = select(IQ1_DELTA, -IQ1_DELTA, (qh_u16 & 0x8000u) != 0u);
582
780
 
583
- #ifdef INIT_SRC0_SHMEM_Q5_K
584
- const BLOCK_SIZE = 256u;
585
- const F16_PER_BLOCK = 88u;
781
+ let gp0_grid_id = ((qs_u16 & 0xFFu) | (((qh_u16 >> (phase * 6u)) & 7u) << 8u)) * 8u;
782
+ let gp1_grid_id = (((qs_u16 >> 8) & 0xFFu) | (((qh_u16 >> (phase * 6u + 3u)) & 7u) << 8u)) * 8u;
586
783
 
587
- fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
588
- for (var elem_idx = thread_id; elem_idx < TILE_SRC0_SHMEM; elem_idx += TOTAL_WORKGROUP_SIZE) {
589
- let tile_m = elem_idx / TILE_K;
590
- let tile_k = elem_idx % TILE_K;
784
+ let gp0_gw = iq1_grid[(gp0_grid_id) / 16u];
785
+ let gp1_gw = iq1_grid[(gp1_grid_id) / 16u];
591
786
 
592
- let global_m = offset_m + tile_m;
593
- let global_k = k_outer + tile_k;
787
+ let gp0_shift_base = (gp0_grid_id % 16u) * 2u;
788
+ let gp1_shift_base = (gp1_grid_id % 16u) * 2u;
594
789
 
595
- if (global_m >= params.m || global_k >= params.k) {
596
- shmem[elem_idx] = f16(0.0);
597
- continue;
598
- }
790
+ store_shmem_iquants(create_iq_gw4(dl, gp0_gw, gp0_shift_base + 0u, delta), elem_idx + 0u);
791
+ store_shmem_iquants(create_iq_gw4(dl, gp0_gw, gp0_shift_base + 8u, delta), elem_idx + 4u);
792
+ store_shmem_iquants(create_iq_gw4(dl, gp1_gw, gp1_shift_base + 0u, delta), elem_idx + 8u);
793
+ store_shmem_iquants(create_iq_gw4(dl, gp1_gw, gp1_shift_base + 8u, delta), elem_idx + 12u);
794
+ #endif // INIT_SRC0_SHMEM_IQ1_S
599
795
 
600
- let block_k = global_k / BLOCK_SIZE;
601
- let k_in_block = global_k % BLOCK_SIZE;
796
+ #if defined(INIT_SRC0_SHMEM_IQ1_M)
797
+ let block_byte_base = src0_idx * 56u; // BLOCK_SIZE_BYTES = 56u;
798
+ let qs_byte_base = block_byte_base + 0u;
799
+ let qh_byte_base = block_byte_base + 32u;
800
+ let scales_byte_base = block_byte_base + 48u;
602
801
 
603
- let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
604
- let scale_idx = src0_idx * F16_PER_BLOCK;
802
+ let scales0 = load_u32_at_src0_aligned(scales_byte_base);
803
+ let scales1 = load_u32_at_src0_aligned(scales_byte_base + 4u);
804
+ let scale_packed = ((scales0 >> 12u) & 0xFu) |
805
+ ((scales0 >> 24u) & 0x00F0u) |
806
+ ((scales1 >> 4u) & 0x0F00u) |
807
+ ((scales1 >> 16u) & 0xF000u);
808
+ let d = f32(bitcast<vec2<f16>>(scale_packed).x);
605
809
 
606
- let d = src0[scale_idx];
607
- let dmin = src0[scale_idx + 1u];
810
+ let sub_block = k_in_block / 32u;
811
+ let phase = (k_in_block / NQ) % 2u;
608
812
 
609
- // Load packed scales
610
- var scale_vals: array<u32, 3>;
611
- for (var i: u32 = 0u; i < 3u; i++) {
612
- let scale_0 = src0[scale_idx + 2u + (2u*i)];
613
- let scale_1 = src0[scale_idx + 2u + (2u*i) + 1u];
614
- scale_vals[i] = bitcast<u32>(vec2(scale_0, scale_1));
615
- }
813
+ let scale_u32 = select(scales0, scales1, sub_block >= 4u);
814
+ let scale_u3 = (scale_u32 >> (16u * ((sub_block / 2u) % 2u) + 6u * (sub_block % 2u) + 3u * phase)) & 0x7u;
815
+ let dl = d * f32(2u * scale_u3 + 1u);
616
816
 
617
- // The original loop processes elements in groups of 64
618
- // Each group of 64: q_b_idx cycles through [0,32,64,96], shift cycles [0,4]
619
- // But u increments EVERY 32 elements (after each l loop)
620
- let group_of_64 = k_in_block / 64u; // 0-3
621
- let pos_in_64 = k_in_block % 64u; // 0-63
622
- let shift_group = pos_in_64 / 32u; // 0 or 1
623
- let l = pos_in_64 % 32u; // 0-31
817
+ let qh_u8 = (load_u32_at_src0_aligned(qh_byte_base + 4u * (sub_block / 2u)) >> (16u * (sub_block % 2u) + 8u * phase)) & 0xFFu;
818
+ let qs_u16 = (load_u32_at_src0_aligned(qs_byte_base + 4u * sub_block) >> (16u * phase)) & 0xFFFFu;
624
819
 
625
- let q_b_idx = group_of_64 * 32u; // 0, 32, 64, 96
626
- let shift = shift_group * 4u; // 0 or 4
627
- let is = k_in_block / 32u; // 0-7
820
+ let gp0_grid_id = ((qs_u16 & 0xFFu) | ((qh_u8 & 7u) << 8u)) * 8u;
821
+ let gp0_delta = select(IQ1_DELTA, -IQ1_DELTA, (qh_u8 & 0x8u) != 0u);
628
822
 
629
- // u increments every 32 elements (0->1, 1->2, 2->4, 3->8, 4->16, 5->32, 6->64, 7->128)
630
- let u_shift = k_in_block / 32u; // 0-7
631
- let u: u32 = 1u << u_shift;
823
+ let gp1_grid_id = (((qs_u16 >> 8u) & 0xFFu) | (((qh_u8 >> 4u) & 7u) << 8u)) * 8u;
824
+ let gp1_delta = select(IQ1_DELTA, -IQ1_DELTA, (qh_u8 & 0x80u) != 0u);
632
825
 
633
- var sc: u32;
634
- var mn: u32;
826
+ let gp0_gw = iq1_grid[(gp0_grid_id) / 16u];
827
+ let gp1_gw = iq1_grid[(gp1_grid_id) / 16u];
635
828
 
636
- if (is < 4u) {
637
- let sc_byte = get_byte(scale_vals[is / 4u], is % 4u);
638
- let min_byte = get_byte(scale_vals[(is + 4u) / 4u], is % 4u);
639
- sc = sc_byte & 63u;
640
- mn = min_byte & 63u;
641
- } else {
642
- let sc_min_lo = get_byte(scale_vals[(is + 4u) / 4u], (is + 4u) % 4u);
643
- let sc_hi = get_byte(scale_vals[(is - 4u) / 4u], (is - 4u) % 4u);
644
- let min_hi = get_byte(scale_vals[is / 4u], is % 4u);
829
+ let gp0_shift_base = (gp0_grid_id % 16u) * 2u;
830
+ let gp1_shift_base = (gp1_grid_id % 16u) * 2u;
645
831
 
646
- sc = (sc_min_lo & 0xFu) | ((sc_hi >> 6u) << 4u);
647
- mn = (sc_min_lo >> 4u) | ((min_hi >> 6u) << 4u);
648
- }
832
+ store_shmem_iquants(create_iq_gw4(dl, gp0_gw, gp0_shift_base + 0u, gp0_delta), elem_idx + 0u);
833
+ store_shmem_iquants(create_iq_gw4(dl, gp0_gw, gp0_shift_base + 8u, gp0_delta), elem_idx + 4u);
834
+ store_shmem_iquants(create_iq_gw4(dl, gp1_gw, gp1_shift_base + 0u, gp1_delta), elem_idx + 8u);
835
+ store_shmem_iquants(create_iq_gw4(dl, gp1_gw, gp1_shift_base + 8u, gp1_delta), elem_idx + 12u);
836
+ #endif // INIT_SRC0_SHMEM_IQ1_M
649
837
 
650
- let dl = d * f16(sc);
651
- let ml = dmin * f16(mn);
838
+ #if defined(INIT_SRC0_SHMEM_IQ2_XXS)
839
+ let block_byte_base = src0_idx * 66u; // BLOCK_SIZE_BYTES = 66u;
840
+ let d_byte_base = block_byte_base + 0u;
841
+ let qs_byte_base = block_byte_base + 2u;
652
842
 
653
- let q_idx = q_b_idx + l;
654
- let q_0 = src0[scale_idx + 24u + 2u * (q_idx / 4u)];
655
- let q_1 = src0[scale_idx + 24u + 2u * (q_idx / 4u) + 1u];
656
- let q_packed = bitcast<u32>(vec2(q_0, q_1));
843
+ let d = load_f16_as_f32_at_src0(d_byte_base);
657
844
 
658
- let q_byte = get_byte(q_packed, q_idx % 4u);
845
+ let sub_block = k_in_block / 32u;
846
+ let phase = (k_in_block / NQ) % 2u;
659
847
 
660
- let qh_0 = src0[scale_idx + 8u + 2u * (l / 4u)];
661
- let qh_1 = src0[scale_idx + 8u + 2u * (l / 4u) + 1u];
662
- let qh_packed = bitcast<u32>(vec2(qh_0, qh_1));
848
+ let aux0 = load_u32_at_src0(qs_byte_base + 8u * sub_block + 0u);
849
+ let aux1 = load_u32_at_src0(qs_byte_base + 8u * sub_block + 4u);
850
+ let db = d * (0.5 + f32(aux1 >> 28u)) * 0.25;
663
851
 
664
- let qh_byte = get_byte(qh_packed, l % 4u);
852
+ let gp0_ig = get_byte(aux0, 2u * phase + 0u) * 8u;
853
+ let gp1_ig = get_byte(aux0, 2u * phase + 1u) * 8u;
665
854
 
666
- let qs_val = (q_byte >> shift) & 0xFu;
667
- let qh_val = select(0.0, 16.0, (qh_byte & u) != 0);
855
+ let gp0_is = (aux1 >> (14u * phase + 0u)) & 127u;
856
+ let gp1_is = (aux1 >> (14u * phase + 7u)) & 127u;
668
857
 
669
- let q_val = (f16(qs_val) + f16(qh_val)) * dl - ml;
670
- shmem[elem_idx] = q_val;
671
- }
672
- }
858
+ let gp0_signs = get_byte(ksigns_iq2xs[gp0_is / 4u], gp0_is % 4u);
859
+ let gp1_signs = get_byte(ksigns_iq2xs[gp1_is / 4u], gp1_is % 4u);
673
860
 
674
- #endif // INIT_SRC0_SHMEM_Q5_K
861
+ let m_0_3_val4 = create_iq2_m4(gp0_signs, 0);
862
+ let m_4_7_val4 = create_iq2_m4(gp0_signs, 1);
863
+ let m_8_11_val4 = create_iq2_m4(gp1_signs, 0);
864
+ let m_12_15_val4 = create_iq2_m4(gp1_signs, 1);
675
865
 
676
- #ifdef INIT_SRC0_SHMEM_Q6_K
677
- const BLOCK_SIZE = 256u;
678
- const F16_PER_BLOCK = 105u;
866
+ let gw_0_3_val4 = create_iq_gw4(gp0_ig, 0);
867
+ let gw_4_7_val4 = create_iq_gw4(gp0_ig, 4);
868
+ let gw_8_11_val4 = create_iq_gw4(gp1_ig, 0);
869
+ let gw_12_15_val4 = create_iq_gw4(gp1_ig, 4);
679
870
 
680
- fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
681
- for (var elem_idx = thread_id; elem_idx < TILE_SRC0_SHMEM; elem_idx += TOTAL_WORKGROUP_SIZE) {
682
- let tile_m = elem_idx / TILE_K;
683
- let tile_k = elem_idx % TILE_K;
871
+ store_shmem_iquants(vec4<f16>(db * m_0_3_val4 * gw_0_3_val4), elem_idx + 0u);
872
+ store_shmem_iquants(vec4<f16>(db * m_4_7_val4 * gw_4_7_val4), elem_idx + 4u);
873
+ store_shmem_iquants(vec4<f16>(db * m_8_11_val4 * gw_8_11_val4), elem_idx + 8u);
874
+ store_shmem_iquants(vec4<f16>(db * m_12_15_val4 * gw_12_15_val4), elem_idx + 12u);
875
+ #endif // INIT_SRC0_SHMEM_IQ2_XXS
684
876
 
685
- let global_m = offset_m + tile_m;
686
- let global_k = k_outer + tile_k;
877
+ #if defined(INIT_SRC0_SHMEM_IQ2_XS)
878
+ let block_byte_base = src0_idx * 74u; // BLOCK_SIZE_BYTES = 74u;
879
+ let d_byte_base = block_byte_base + 0u;
880
+ let qs_byte_base = block_byte_base + 2u;
881
+ let scales_byte_base = block_byte_base + 66u;
687
882
 
688
- if (global_m >= params.m || global_k >= params.k) {
689
- shmem[elem_idx] = f16(0.0);
690
- continue;
691
- }
883
+ let d = load_f16_as_f32_at_src0(d_byte_base);
692
884
 
693
- let block_k = global_k / BLOCK_SIZE;
694
- let k_in_block = global_k % BLOCK_SIZE;
885
+ let sub_block = k_in_block / 32u;
886
+ let phase = (k_in_block / NQ) % 2u;
695
887
 
696
- let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
697
- let scale_idx = src0_idx * F16_PER_BLOCK;
698
-
699
- let half = k_in_block / 128u;
700
- let pos_in_half = k_in_block % 128u;
701
- let quarter = pos_in_half / 32u;
702
- let l = pos_in_half % 32u;
703
-
704
- let ql_b_idx = half * 64u;
705
- let qh_b_idx = half * 32u;
706
- let sc_b_idx = half * 8u;
707
-
708
- // Load only ql13 word needed
709
- let ql13_flat = ql_b_idx + l;
710
- let ql13_word = ql13_flat / 4u;
711
- let ql13 = bitcast<u32>(vec2(
712
- src0[scale_idx + 2u * ql13_word],
713
- src0[scale_idx + 2u * ql13_word + 1u]
714
- ));
715
- let ql13_b = get_byte(ql13, ql13_flat % 4u);
716
-
717
- // Load only ql24 word needed
718
- let ql24_flat = ql_b_idx + l + 32u;
719
- let ql24_word = ql24_flat / 4u;
720
- let ql24 = bitcast<u32>(vec2(
721
- src0[scale_idx + 2u * ql24_word],
722
- src0[scale_idx + 2u * ql24_word + 1u]
723
- ));
724
- let ql24_b = get_byte(ql24, ql24_flat % 4u);
725
-
726
- // Load only qh word needed
727
- let qh_flat = qh_b_idx + l;
728
- let qh_word = qh_flat / 4u;
729
- let qh = bitcast<u32>(vec2(
730
- src0[scale_idx + 64u + 2u * qh_word],
731
- src0[scale_idx + 64u + 2u * qh_word + 1u]
732
- ));
733
- let qh_b = get_byte(qh, qh_flat % 4u);
734
-
735
- let q1 = f16((ql13_b & 0xFu) | ((qh_b & 3u) << 4u)) - f16(32.0);
736
- let q2 = f16((ql24_b & 0xFu) | (((qh_b >> 2u) & 3u) << 4u)) - f16(32.0);
737
- let q3 = f16((ql13_b >> 4u) | (((qh_b >> 4u) & 3u) << 4u)) - f16(32.0);
738
- let q4 = f16((ql24_b >> 4u) | (((qh_b >> 6u) & 3u) << 4u)) - f16(32.0);
739
-
740
- // Load only the scale word needed
741
- let is = l / 16u;
742
- let sc_idx = sc_b_idx + is + quarter * 2u;
743
- let sc_word = sc_idx / 4u;
744
- let sc = bitcast<u32>(vec2(
745
- src0[scale_idx + 96u + 2u * sc_word],
746
- src0[scale_idx + 96u + 2u * sc_word + 1u]
747
- ));
748
- let sc_val = get_byte_i32(sc, sc_idx % 4u);
749
-
750
- let d = src0[scale_idx + 104u];
751
-
752
- var q_val: f16;
753
- if (quarter == 0u) {
754
- q_val = q1;
755
- } else if (quarter == 1u) {
756
- q_val = q2;
757
- } else if (quarter == 2u) {
758
- q_val = q3;
759
- } else {
760
- q_val = q4;
761
- }
888
+ let scale = (load_byte_at_src0_aligned(scales_byte_base + 1u * sub_block) >> (4u * phase)) & 0xFu;
889
+ let db = d * (0.5 + f32(scale)) * 0.25;
890
+
891
+ let qs_u32 = load_u32_at_src0(qs_byte_base + 8u * sub_block + 4u * phase);
892
+
893
+ let gp0_ig = (qs_u32 & 0x1FFu) * 8u;
894
+ let gp1_ig = ((qs_u32 >> 16u) & 0x1FFu) * 8u;
895
+
896
+ let gp0_is = (qs_u32 >> 9u) & 0x7Fu;
897
+ let gp1_is = (qs_u32 >> 25u) & 0x7Fu;
898
+
899
+ let gp0_signs = get_byte(ksigns_iq2xs[gp0_is / 4u], gp0_is % 4u);
900
+ let gp1_signs = get_byte(ksigns_iq2xs[gp1_is / 4u], gp1_is % 4u);
762
901
 
763
- shmem[elem_idx] = d * f16(sc_val) * q_val;
902
+ let m_0_3_val4 = create_iq2_m4(gp0_signs, 0);
903
+ let m_4_7_val4 = create_iq2_m4(gp0_signs, 1);
904
+ let m_8_11_val4 = create_iq2_m4(gp1_signs, 0);
905
+ let m_12_15_val4 = create_iq2_m4(gp1_signs, 1);
906
+
907
+ let gw_0_3_val4 = create_iq_gw4(gp0_ig, 0);
908
+ let gw_4_7_val4 = create_iq_gw4(gp0_ig, 4);
909
+ let gw_8_11_val4 = create_iq_gw4(gp1_ig, 0);
910
+ let gw_12_15_val4 = create_iq_gw4(gp1_ig, 4);
911
+
912
+ store_shmem_iquants(vec4<f16>(db * m_0_3_val4 * gw_0_3_val4), elem_idx + 0u);
913
+ store_shmem_iquants(vec4<f16>(db * m_4_7_val4 * gw_4_7_val4), elem_idx + 4u);
914
+ store_shmem_iquants(vec4<f16>(db * m_8_11_val4 * gw_8_11_val4), elem_idx + 8u);
915
+ store_shmem_iquants(vec4<f16>(db * m_12_15_val4 * gw_12_15_val4), elem_idx + 12u);
916
+ #endif // INIT_SRC0_SHMEM_IQ2_XS
917
+
918
+ #if defined(INIT_SRC0_SHMEM_IQ2_S)
919
+ let block_byte_base = src0_idx * 82u; // BLOCK_SIZE_BYTES = 82u;
920
+ let d_byte_base = block_byte_base + 0u;
921
+ let qs_byte_base = block_byte_base + 2u;
922
+ let qh_byte_base = block_byte_base + 66u;
923
+ let scales_byte_base = block_byte_base + 74u;
924
+
925
+ let d = load_f16_as_f32_at_src0(d_byte_base);
926
+
927
+ let sub_block = k_in_block / 32u;
928
+ let phase = (k_in_block / NQ) % 2u;
929
+
930
+ let scale = (load_byte_at_src0_aligned(scales_byte_base + 1u * sub_block) >> (4u * phase)) & 0xFu;
931
+ let db = d * (0.5 + f32(scale)) * 0.25;
932
+
933
+ let qs_u16 = load_u32_at_src0(qs_byte_base + 4u * sub_block + 2u * phase) & 0xFFFFu;
934
+ let signs_u16 = load_u32_at_src0(qs_byte_base + 32u + 4u * sub_block + 2u * phase) & 0xFFFFu;
935
+ let qh_u4 = (load_byte_at_src0_aligned(qh_byte_base + 1u * sub_block) >> (4u * phase)) & 0xFu;
936
+
937
+ let gp0_ig = ((qs_u16 & 0xFFu) | ((qh_u4 & 0x3u) << 8u)) * 8u;
938
+ let gp1_ig = (((qs_u16 >> 8u) & 0xFFu) | ((qh_u4 & 0xCu) << 6u)) * 8u;
939
+
940
+ let gp0_signs = get_byte(signs_u16, 0);
941
+ let gp1_signs = get_byte(signs_u16, 1);
942
+
943
+ let m_0_3_val4 = create_iq2_m4(gp0_signs, 0);
944
+ let m_4_7_val4 = create_iq2_m4(gp0_signs, 1);
945
+ let m_8_11_val4 = create_iq2_m4(gp1_signs, 0);
946
+ let m_12_15_val4 = create_iq2_m4(gp1_signs, 1);
947
+
948
+ let gw_0_3_val4 = create_iq_gw4(gp0_ig, 0);
949
+ let gw_4_7_val4 = create_iq_gw4(gp0_ig, 4);
950
+ let gw_8_11_val4 = create_iq_gw4(gp1_ig, 0);
951
+ let gw_12_15_val4 = create_iq_gw4(gp1_ig, 4);
952
+
953
+ store_shmem_iquants(vec4<f16>(db * m_0_3_val4 * gw_0_3_val4), elem_idx + 0u);
954
+ store_shmem_iquants(vec4<f16>(db * m_4_7_val4 * gw_4_7_val4), elem_idx + 4u);
955
+ store_shmem_iquants(vec4<f16>(db * m_8_11_val4 * gw_8_11_val4), elem_idx + 8u);
956
+ store_shmem_iquants(vec4<f16>(db * m_12_15_val4 * gw_12_15_val4), elem_idx + 12u);
957
+ #endif // INIT_SRC0_SHMEM_IQ2_S
958
+
959
+ #if defined(INIT_SRC0_SHMEM_IQ3_XXS)
960
+ let block_byte_base = src0_idx * 98u; // BLOCK_SIZE_BYTES = 98u;
961
+ let d_byte_base = block_byte_base + 0u;
962
+ let qs_byte_base = block_byte_base + 2u;
963
+
964
+ let d = load_f16_as_f32_at_src0(d_byte_base);
965
+
966
+ let sub_block = k_in_block / 32u;
967
+ let phase = (k_in_block / NQ) % 2u;
968
+
969
+ let qs_u32 = load_u32_at_src0(qs_byte_base + 8u * sub_block + 4u * phase);
970
+ let sign_u32 = load_u32_at_src0(qs_byte_base + 64u + 4u * sub_block);
971
+ let db = d * (0.5 + f32(sign_u32 >> 28u)) * 0.5;
972
+
973
+ let ig_0_3 = get_byte(qs_u32, 0);
974
+ let ig_4_7 = get_byte(qs_u32, 1);
975
+ let ig_8_11 = get_byte(qs_u32, 2);
976
+ let ig_12_15 = get_byte(qs_u32, 3);
977
+
978
+ let gp0_is = (sign_u32 >> (14u * phase + 0u)) & 0x7Fu;
979
+ let gp1_is = (sign_u32 >> (14u * phase + 7u)) & 0x7Fu;
980
+
981
+ let gp0_signs = get_byte(ksigns_iq2xs[gp0_is / 4u], gp0_is % 4u);
982
+ let gp1_signs = get_byte(ksigns_iq2xs[gp1_is / 4u], gp1_is % 4u);
983
+
984
+ let m_0_3_val4 = create_iq2_m4(gp0_signs, 0);
985
+ let m_4_7_val4 = create_iq2_m4(gp0_signs, 1);
986
+ let m_8_11_val4 = create_iq2_m4(gp1_signs, 0);
987
+ let m_12_15_val4 = create_iq2_m4(gp1_signs, 1);
988
+
989
+ let gw_0_3_val4 = create_iq_gw4(ig_0_3);
990
+ let gw_4_7_val4 = create_iq_gw4(ig_4_7);
991
+ let gw_8_11_val4 = create_iq_gw4(ig_8_11);
992
+ let gw_12_15_val4 = create_iq_gw4(ig_12_15);
993
+
994
+ store_shmem_iquants(vec4<f16>(db * m_0_3_val4 * gw_0_3_val4), elem_idx + 0u);
995
+ store_shmem_iquants(vec4<f16>(db * m_4_7_val4 * gw_4_7_val4), elem_idx + 4u);
996
+ store_shmem_iquants(vec4<f16>(db * m_8_11_val4 * gw_8_11_val4), elem_idx + 8u);
997
+ store_shmem_iquants(vec4<f16>(db * m_12_15_val4 * gw_12_15_val4), elem_idx + 12u);
998
+ #endif // INIT_SRC0_SHMEM_IQ3_XXS
999
+
1000
+ #if defined(INIT_SRC0_SHMEM_IQ3_S)
1001
+ let block_byte_base = src0_idx * 110u; // BLOCK_SIZE_BYTES = 110u;
1002
+ let d_byte_base = block_byte_base + 0u;
1003
+ let qs_byte_base = block_byte_base + 2u;
1004
+ let qh_byte_base = block_byte_base + 66u;
1005
+ let signs_byte_base = block_byte_base + 74u;
1006
+ let scales_byte_base = block_byte_base + 106u;
1007
+
1008
+ let d = load_f16_as_f32_at_src0(d_byte_base);
1009
+
1010
+ let sub_block = k_in_block / 32u;
1011
+ let phase = (k_in_block / NQ) % 2u;
1012
+
1013
+ let scale = (load_byte_at_src0_aligned(scales_byte_base + 1u * (sub_block / 2u)) >> (4u * (sub_block % 2u))) & 0xFu;
1014
+ let db = d * (1.0 + 2.0 * f32(scale));
1015
+
1016
+ let qs_u32 = load_u32_at_src0(qs_byte_base + 8u * sub_block + 4u * phase);
1017
+ let qh_u4 = (load_byte_at_src0_aligned(qh_byte_base + 1u * sub_block) >> (4u * phase)) & 0xFu;
1018
+ let signs_u16 = (load_u32_at_src0(signs_byte_base + 4u * sub_block + 2u * phase)) & 0xFFFFu;
1019
+
1020
+ let ig_0_3 = ((qs_u32 >> 0u) & 0xFFu) | ((qh_u4 & 0x1u) << 8u);
1021
+ let ig_4_7 = ((qs_u32 >> 8u) & 0xFFu) | ((qh_u4 & 0x2u) << 7u);
1022
+ let ig_8_11 = ((qs_u32 >> 16u) & 0xFFu) | ((qh_u4 & 0x4u) << 6u);
1023
+ let ig_12_15 = ((qs_u32 >> 24u) & 0xFFu) | ((qh_u4 & 0x8u) << 5u);
1024
+
1025
+ let gp0_signs = get_byte(signs_u16, 0);
1026
+ let gp1_signs = get_byte(signs_u16, 1);
1027
+
1028
+ let m_0_3_val4 = create_iq2_m4(gp0_signs, 0);
1029
+ let m_4_7_val4 = create_iq2_m4(gp0_signs, 1);
1030
+ let m_8_11_val4 = create_iq2_m4(gp1_signs, 0);
1031
+ let m_12_15_val4 = create_iq2_m4(gp1_signs, 1);
1032
+
1033
+ let gw_0_3_val4 = create_iq_gw4(ig_0_3);
1034
+ let gw_4_7_val4 = create_iq_gw4(ig_4_7);
1035
+ let gw_8_11_val4 = create_iq_gw4(ig_8_11);
1036
+ let gw_12_15_val4 = create_iq_gw4(ig_12_15);
1037
+
1038
+ store_shmem_iquants(vec4<f16>(db * m_0_3_val4 * gw_0_3_val4), elem_idx + 0u);
1039
+ store_shmem_iquants(vec4<f16>(db * m_4_7_val4 * gw_4_7_val4), elem_idx + 4u);
1040
+ store_shmem_iquants(vec4<f16>(db * m_8_11_val4 * gw_8_11_val4), elem_idx + 8u);
1041
+ store_shmem_iquants(vec4<f16>(db * m_12_15_val4 * gw_12_15_val4), elem_idx + 12u);
1042
+ #endif // INIT_SRC0_SHMEM_IQ3_S
764
1043
  }
765
1044
  }
766
- #endif // INIT_SRC0_SHMEM_Q6_K
1045
+ #endif // i-quants (super block size: 256)