whispercpp 1.3.7 → 1.3.8

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (308) hide show
  1. checksums.yaml +4 -4
  2. data/README.md +5 -4
  3. data/ext/options.rb +1 -1
  4. data/ext/ruby_whisper.c +0 -1
  5. data/ext/ruby_whisper.h +7 -1
  6. data/ext/ruby_whisper_context.c +50 -1
  7. data/ext/ruby_whisper_log_settable.h +1 -2
  8. data/ext/ruby_whisper_params.c +9 -8
  9. data/ext/ruby_whisper_transcribe.cpp +0 -19
  10. data/ext/ruby_whisper_vad_context.c +30 -10
  11. data/ext/ruby_whisper_vad_context_detect.cpp +8 -9
  12. data/ext/ruby_whisper_vad_params.c +4 -4
  13. data/ext/ruby_whisper_vad_segment.c +2 -2
  14. data/ext/sources/CMakeLists.txt +2 -1
  15. data/ext/sources/cmake/parakeet.pc.in +2 -2
  16. data/ext/sources/cmake/whisper.pc.in +2 -2
  17. data/ext/sources/examples/cli/cli.cpp +9 -1
  18. data/ext/sources/examples/common-ggml.cpp +2 -0
  19. data/ext/sources/examples/vad-speech-segments/speech.cpp +3 -2
  20. data/ext/sources/ggml/CMakeLists.txt +3 -4
  21. data/ext/sources/ggml/include/ggml-cuda.h +0 -3
  22. data/ext/sources/ggml/include/ggml-sycl.h +8 -0
  23. data/ext/sources/ggml/include/ggml.h +3 -1
  24. data/ext/sources/ggml/src/CMakeLists.txt +8 -1
  25. data/ext/sources/ggml/src/ggml-backend-meta.cpp +7 -4
  26. data/ext/sources/ggml/src/ggml-common.h +13 -2
  27. data/ext/sources/ggml/src/ggml-cpu/CMakeLists.txt +1 -1
  28. data/ext/sources/ggml/src/ggml-cpu/amx/mmq.cpp +5 -6
  29. data/ext/sources/ggml/src/ggml-cpu/arch/arm/quants.c +78 -4
  30. data/ext/sources/ggml/src/ggml-cpu/arch/x86/quants.c +142 -4
  31. data/ext/sources/ggml/src/ggml-cpu/arch-fallback.h +7 -2
  32. data/ext/sources/ggml/src/ggml-cpu/ggml-cpu.c +14 -0
  33. data/ext/sources/ggml/src/ggml-cpu/llamafile/sgemm.cpp +26 -19
  34. data/ext/sources/ggml/src/ggml-cpu/ops.cpp +129 -46
  35. data/ext/sources/ggml/src/ggml-cpu/quants.c +51 -0
  36. data/ext/sources/ggml/src/ggml-cpu/quants.h +3 -0
  37. data/ext/sources/ggml/src/ggml-cpu/simd-gemm.h +1 -1
  38. data/ext/sources/ggml/src/ggml-cpu/simd-mappings.h +11 -0
  39. data/ext/sources/ggml/src/ggml-cpu/vec.cpp +2 -2
  40. data/ext/sources/ggml/src/ggml-cuda/binbcast.cu +90 -46
  41. data/ext/sources/ggml/src/ggml-cuda/col2im-1d.cu +81 -0
  42. data/ext/sources/ggml/src/ggml-cuda/col2im-1d.cuh +3 -0
  43. data/ext/sources/ggml/src/ggml-cuda/common.cuh +4 -0
  44. data/ext/sources/ggml/src/ggml-cuda/concat.cu +33 -21
  45. data/ext/sources/ggml/src/ggml-cuda/conv-transpose-1d.cu +14 -12
  46. data/ext/sources/ggml/src/ggml-cuda/convert.cu +86 -34
  47. data/ext/sources/ggml/src/ggml-cuda/cpy.cu +80 -29
  48. data/ext/sources/ggml/src/ggml-cuda/fattn-common.cuh +9 -5
  49. data/ext/sources/ggml/src/ggml-cuda/fattn-mma-f16.cuh +4 -0
  50. data/ext/sources/ggml/src/ggml-cuda/fattn-tile.cuh +9 -5
  51. data/ext/sources/ggml/src/ggml-cuda/fattn.cu +27 -21
  52. data/ext/sources/ggml/src/ggml-cuda/gated_delta_net.cu +40 -25
  53. data/ext/sources/ggml/src/ggml-cuda/gated_delta_net.cuh +10 -0
  54. data/ext/sources/ggml/src/ggml-cuda/getrows.cu +15 -12
  55. data/ext/sources/ggml/src/ggml-cuda/ggml-cuda.cu +718 -1248
  56. data/ext/sources/ggml/src/ggml-cuda/mmq.cu +7 -0
  57. data/ext/sources/ggml/src/ggml-cuda/mmvq.cu +77 -40
  58. data/ext/sources/ggml/src/ggml-cuda/out-prod.cu +55 -12
  59. data/ext/sources/ggml/src/ggml-cuda/set-rows.cu +64 -4
  60. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_16-ncols2_2.cu +1 -0
  61. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_32-ncols2_2.cu +1 -0
  62. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_2.cu +1 -0
  63. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_8-ncols2_2.cu +1 -0
  64. data/ext/sources/ggml/src/ggml-cuda/topk-moe.cu +7 -1
  65. data/ext/sources/ggml/src/ggml-cuda/vendors/hip.h +1 -0
  66. data/ext/sources/ggml/src/ggml-cuda/vendors/musa.h +1 -0
  67. data/ext/sources/ggml/src/ggml-hexagon/CMakeLists.txt +0 -5
  68. data/ext/sources/ggml/src/ggml-hexagon/ggml-hexagon.cpp +1634 -1293
  69. data/ext/sources/ggml/src/ggml-hexagon/htp/CMakeLists.txt +11 -40
  70. data/ext/sources/ggml/src/ggml-hexagon/htp/cmake-toolchain.cmake +13 -15
  71. data/ext/sources/ggml/src/ggml-hexagon/htp/concat-ops.c +1 -1
  72. data/ext/sources/ggml/src/ggml-hexagon/htp/flash-attn-ops.c +1749 -399
  73. data/ext/sources/ggml/src/ggml-hexagon/htp/flash-attn-ops.h +303 -0
  74. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-common.h +80 -0
  75. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-dma.h +26 -23
  76. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-profile.h +64 -0
  77. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-utils.h +1 -83
  78. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-fa-kernels.h +555 -0
  79. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-mm-kernels-tiled.h +1303 -0
  80. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-queue.c +9 -0
  81. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-queue.h +27 -4
  82. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-utils.h +59 -37
  83. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-ctx.h +11 -3
  84. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-ops.h +52 -12
  85. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-vtcm.h +19 -0
  86. data/ext/sources/ggml/src/ggml-hexagon/htp/htp_iface.idl +2 -1
  87. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-base.h +14 -30
  88. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-exp.h +39 -0
  89. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-fa-kernels.h +232 -0
  90. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-flat.h +1511 -0
  91. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-tiled.h +1200 -0
  92. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-sigmoid.h +39 -0
  93. data/ext/sources/ggml/src/ggml-hexagon/htp/main.c +127 -32
  94. data/ext/sources/ggml/src/ggml-hexagon/htp/matmul-ops.c +3023 -4425
  95. data/ext/sources/ggml/src/ggml-hexagon/htp/matmul-ops.h +650 -0
  96. data/ext/sources/ggml/src/ggml-hexagon/htp/rope-ops.c +48 -13
  97. data/ext/sources/ggml/src/ggml-hexagon/htp/ssm-conv.c +10 -9
  98. data/ext/sources/ggml/src/ggml-hexagon/htp/worker-pool.c +15 -3
  99. data/ext/sources/ggml/src/ggml-hexagon/htp/worker-pool.h +8 -0
  100. data/ext/sources/ggml/src/ggml-hexagon/htp-opnode.h +168 -50
  101. data/ext/sources/ggml/src/ggml-hexagon/libggml-htp.inf +0 -4
  102. data/ext/sources/ggml/src/ggml-hip/CMakeLists.txt +5 -0
  103. data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.cpp +69 -5
  104. data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.h +4 -1
  105. data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.m +27 -6
  106. data/ext/sources/ggml/src/ggml-metal/ggml-metal-impl.h +38 -0
  107. data/ext/sources/ggml/src/ggml-metal/ggml-metal-ops.cpp +132 -2
  108. data/ext/sources/ggml/src/ggml-metal/ggml-metal-ops.h +2 -0
  109. data/ext/sources/ggml/src/ggml-metal/ggml-metal.metal +345 -87
  110. data/ext/sources/ggml/src/ggml-opencl/CMakeLists.txt +13 -0
  111. data/ext/sources/ggml/src/ggml-opencl/fa_tune.h +92 -0
  112. data/ext/sources/ggml/src/ggml-opencl/ggml-opencl.cpp +4060 -357
  113. data/ext/sources/ggml/src/ggml-opencl/kernels/cvt.cl +198 -0
  114. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f16.cl +81 -41
  115. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32.cl +88 -39
  116. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_f16.cl +1995 -96
  117. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_q4_0.cl +1615 -0
  118. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_q8_0.cl +1486 -0
  119. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_pre_f16.cl +156 -0
  120. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_mxfp4_f32_ns.cl +74 -6
  121. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_0_f32_ns.cl +74 -6
  122. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_1_f32_ns.cl +74 -6
  123. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_k_f32_ns.cl +71 -6
  124. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_0_f32_ns.cl +74 -6
  125. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_1_f32_ns.cl +74 -6
  126. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_k_f32_ns.cl +74 -6
  127. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q6_k_f32_ns.cl +74 -6
  128. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q1_0_f32.cl +94 -0
  129. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q1_0_f32.cl +121 -0
  130. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q8_0_f32.cl +1 -1
  131. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q1_0_f32_l4_lm.cl +156 -0
  132. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_f16_f32_l4.cl +1149 -0
  133. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q1_0_f32.cl +141 -0
  134. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q1_0_f32_flat.cl +190 -0
  135. data/ext/sources/ggml/src/ggml-opencl/kernels/norm.cl +5 -2
  136. data/ext/sources/ggml/src/ggml-opencl/kernels/set_rows.cl +500 -0
  137. data/ext/sources/ggml/src/ggml-opencl/libdl.h +79 -0
  138. data/ext/sources/ggml/src/ggml-openvino/.clang-format +0 -5
  139. data/ext/sources/ggml/src/ggml-openvino/CMakeLists.txt +2 -4
  140. data/ext/sources/ggml/src/ggml-openvino/ggml-decoder.cpp +733 -130
  141. data/ext/sources/ggml/src/ggml-openvino/ggml-decoder.h +76 -23
  142. data/ext/sources/ggml/src/ggml-openvino/ggml-openvino-extra.cpp +57 -3
  143. data/ext/sources/ggml/src/ggml-openvino/ggml-openvino-extra.h +29 -8
  144. data/ext/sources/ggml/src/ggml-openvino/ggml-openvino.cpp +307 -59
  145. data/ext/sources/ggml/src/ggml-openvino/ggml-quants.cpp +66 -0
  146. data/ext/sources/ggml/src/ggml-openvino/ggml-quants.h +10 -4
  147. data/ext/sources/ggml/src/ggml-openvino/openvino/decoder.h +56 -16
  148. data/ext/sources/ggml/src/ggml-openvino/openvino/frontend.h +1 -1
  149. data/ext/sources/ggml/src/ggml-openvino/openvino/input_model.h +4 -4
  150. data/ext/sources/ggml/src/ggml-openvino/openvino/node_context.h +94 -37
  151. data/ext/sources/ggml/src/ggml-openvino/openvino/op/add_id.cpp +76 -0
  152. data/ext/sources/ggml/src/ggml-openvino/openvino/op/argsort.cpp +47 -0
  153. data/ext/sources/ggml/src/ggml-openvino/openvino/op/clamp.cpp +33 -0
  154. data/ext/sources/ggml/src/ggml-openvino/openvino/op/concat.cpp +48 -0
  155. data/ext/sources/ggml/src/ggml-openvino/openvino/op/cont.cpp +8 -16
  156. data/ext/sources/ggml/src/ggml-openvino/openvino/op/cpy.cpp +14 -1
  157. data/ext/sources/ggml/src/ggml-openvino/openvino/op/div.cpp +146 -0
  158. data/ext/sources/ggml/src/ggml-openvino/openvino/op/flash_attn_ext.cpp +108 -21
  159. data/ext/sources/ggml/src/ggml-openvino/openvino/op/gated_delta_net.cpp +282 -0
  160. data/ext/sources/ggml/src/ggml-openvino/openvino/op/gated_delta_net.hpp +65 -0
  161. data/ext/sources/ggml/src/ggml-openvino/openvino/op/get_rows.cpp +2 -9
  162. data/ext/sources/ggml/src/ggml-openvino/openvino/op/glu_geglu.cpp +21 -7
  163. data/ext/sources/ggml/src/ggml-openvino/openvino/op/glu_swiglu.cpp +41 -8
  164. data/ext/sources/ggml/src/ggml-openvino/openvino/op/im2col.cpp +120 -0
  165. data/ext/sources/ggml/src/ggml-openvino/openvino/op/l2_norm.cpp +44 -0
  166. data/ext/sources/ggml/src/ggml-openvino/openvino/op/mul_mat_id.cpp +226 -0
  167. data/ext/sources/ggml/src/ggml-openvino/openvino/op/mulmat.cpp +19 -9
  168. data/ext/sources/ggml/src/ggml-openvino/openvino/op/norm.cpp +58 -0
  169. data/ext/sources/ggml/src/ggml-openvino/openvino/op/pad.cpp +95 -0
  170. data/ext/sources/ggml/src/ggml-openvino/openvino/op/permute.cpp +58 -13
  171. data/ext/sources/ggml/src/ggml-openvino/openvino/op/repeat.cpp +74 -0
  172. data/ext/sources/ggml/src/ggml-openvino/openvino/op/reshape.cpp +13 -6
  173. data/ext/sources/ggml/src/ggml-openvino/openvino/op/rms_norm.cpp +1 -1
  174. data/ext/sources/ggml/src/ggml-openvino/openvino/op/rope.cpp +134 -38
  175. data/ext/sources/ggml/src/ggml-openvino/openvino/op/set_rows.cpp +3 -3
  176. data/ext/sources/ggml/src/ggml-openvino/openvino/op/softmax.cpp +126 -49
  177. data/ext/sources/ggml/src/ggml-openvino/openvino/op/ssm_conv.cpp +59 -0
  178. data/ext/sources/ggml/src/ggml-openvino/openvino/op/sum_rows.cpp +27 -0
  179. data/ext/sources/ggml/src/ggml-openvino/openvino/op/transpose.cpp +32 -1
  180. data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_silu.cpp +1 -1
  181. data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_softplus.cpp +38 -0
  182. data/ext/sources/ggml/src/ggml-openvino/openvino/op/view.cpp +90 -25
  183. data/ext/sources/ggml/src/ggml-openvino/openvino/op_table.cpp +41 -23
  184. data/ext/sources/ggml/src/ggml-openvino/openvino/op_table.h +18 -5
  185. data/ext/sources/ggml/src/ggml-openvino/openvino/pass/mark_decompression_convert_constant_folding.h +1 -1
  186. data/ext/sources/ggml/src/ggml-openvino/openvino/translate_session.cpp +43 -40
  187. data/ext/sources/ggml/src/ggml-openvino/openvino/translate_session.h +5 -4
  188. data/ext/sources/ggml/src/ggml-openvino/openvino/utils.cpp +548 -3
  189. data/ext/sources/ggml/src/ggml-openvino/openvino/utils.h +28 -26
  190. data/ext/sources/ggml/src/ggml-openvino/utils.cpp +383 -94
  191. data/ext/sources/ggml/src/ggml-openvino/utils.h +11 -8
  192. data/ext/sources/ggml/src/ggml-quants.c +76 -0
  193. data/ext/sources/ggml/src/ggml-quants.h +3 -0
  194. data/ext/sources/ggml/src/ggml-sycl/CMakeLists.txt +5 -5
  195. data/ext/sources/ggml/src/ggml-sycl/backend.hpp +2 -0
  196. data/ext/sources/ggml/src/ggml-sycl/binbcast.cpp +12 -0
  197. data/ext/sources/ggml/src/ggml-sycl/col2im-1d.cpp +102 -0
  198. data/ext/sources/ggml/src/ggml-sycl/col2im-1d.hpp +8 -0
  199. data/ext/sources/ggml/src/ggml-sycl/common.cpp +6 -8
  200. data/ext/sources/ggml/src/ggml-sycl/common.hpp +19 -2
  201. data/ext/sources/ggml/src/ggml-sycl/concat.cpp +21 -1
  202. data/ext/sources/ggml/src/ggml-sycl/conv2d-dw.cpp +158 -0
  203. data/ext/sources/ggml/src/ggml-sycl/conv2d-dw.hpp +10 -0
  204. data/ext/sources/ggml/src/ggml-sycl/conv2d-transpose.cpp +125 -0
  205. data/ext/sources/ggml/src/ggml-sycl/conv2d-transpose.hpp +10 -0
  206. data/ext/sources/ggml/src/ggml-sycl/conv2d.cpp +150 -0
  207. data/ext/sources/ggml/src/ggml-sycl/conv2d.hpp +10 -0
  208. data/ext/sources/ggml/src/ggml-sycl/conv3d.cpp +224 -0
  209. data/ext/sources/ggml/src/ggml-sycl/conv3d.hpp +8 -0
  210. data/ext/sources/ggml/src/ggml-sycl/convert.cpp +6 -0
  211. data/ext/sources/ggml/src/ggml-sycl/cpy.cpp +706 -0
  212. data/ext/sources/ggml/src/ggml-sycl/cpy.hpp +281 -0
  213. data/ext/sources/ggml/src/ggml-sycl/cross_entropy_loss.cpp +255 -0
  214. data/ext/sources/ggml/src/ggml-sycl/cross_entropy_loss.hpp +7 -0
  215. data/ext/sources/ggml/src/ggml-sycl/dequantize.hpp +15 -0
  216. data/ext/sources/ggml/src/ggml-sycl/dmmv.cpp +492 -319
  217. data/ext/sources/ggml/src/ggml-sycl/dpct/helper.hpp +15 -7
  218. data/ext/sources/ggml/src/ggml-sycl/element_wise.cpp +215 -115
  219. data/ext/sources/ggml/src/ggml-sycl/element_wise.hpp +2 -0
  220. data/ext/sources/ggml/src/ggml-sycl/ggml-sycl.cpp +1006 -336
  221. data/ext/sources/ggml/src/ggml-sycl/mmvq.cpp +252 -67
  222. data/ext/sources/ggml/src/ggml-sycl/mmvq.hpp +17 -0
  223. data/ext/sources/ggml/src/ggml-sycl/norm.cpp +103 -49
  224. data/ext/sources/ggml/src/ggml-sycl/outprod.cpp +45 -9
  225. data/ext/sources/ggml/src/ggml-sycl/pool.cpp +185 -0
  226. data/ext/sources/ggml/src/ggml-sycl/pool.hpp +22 -0
  227. data/ext/sources/ggml/src/ggml-sycl/presets.hpp +3 -1
  228. data/ext/sources/ggml/src/ggml-sycl/set_rows.cpp +10 -2
  229. data/ext/sources/ggml/src/ggml-sycl/softmax.cpp +9 -10
  230. data/ext/sources/ggml/src/ggml-sycl/vecdotq.hpp +35 -0
  231. data/ext/sources/ggml/src/ggml-vulkan/CMakeLists.txt +5 -0
  232. data/ext/sources/ggml/src/ggml-vulkan/ggml-vulkan.cpp +833 -215
  233. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/col2im_1d.comp +61 -0
  234. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/conv2d_mm.comp +1 -1
  235. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/conv3d_mm.comp +431 -0
  236. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/diag.comp +3 -3
  237. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp +1 -0
  238. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp +1 -0
  239. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/generic_unary_head.glsl +21 -19
  240. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/get_rows_back.comp +25 -0
  241. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/glu_head.glsl +23 -4
  242. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/glu_main.glsl +14 -18
  243. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/l2_norm.comp +4 -7
  244. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq.comp +21 -24
  245. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp +31 -23
  246. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp +6 -5
  247. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl +84 -67
  248. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/norm.comp +10 -10
  249. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/repeat_back.comp +3 -3
  250. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/roll.comp +3 -3
  251. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/tri.comp +3 -3
  252. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/unary.comp +168 -0
  253. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +121 -74
  254. data/ext/sources/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp +26 -19
  255. data/ext/sources/ggml/src/ggml-webgpu/ggml-webgpu.cpp +31 -36
  256. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/common_decls.tmpl +16 -2
  257. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_split.wgsl +7 -7
  258. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/get_rows.wgsl +21 -0
  259. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_decls.tmpl +439 -320
  260. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_id_vec.wgsl +2 -2
  261. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec.wgsl +45 -39
  262. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_acc.tmpl +586 -465
  263. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_q_acc.tmpl +63 -69
  264. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/quantize_q8.wgsl +14 -9
  265. data/ext/sources/ggml/src/ggml.c +36 -14
  266. data/ext/sources/include/whisper.h +21 -0
  267. data/ext/sources/src/whisper.cpp +164 -14
  268. data/lib/whisper/log_settable.rb +5 -8
  269. data/lib/whisper/model/uri.rb +0 -7
  270. data/sig/whisper.rbs +6 -0
  271. data/test/test_vad.rb +9 -0
  272. data/test/test_vad_context.rb +2 -2
  273. data/whispercpp.gemspec +1 -1
  274. metadata +62 -37
  275. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-flash-attn-ops.c +0 -1878
  276. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-matmul-ops.c +0 -2066
  277. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-ops.c +0 -6
  278. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-ops.h +0 -88
  279. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-profile.h +0 -34
  280. data/ext/sources/ggml/src/ggml-hexagon/htp/vtcm-utils.h +0 -16
  281. data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_gelu.cpp +0 -25
  282. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/abs.comp +0 -21
  283. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/ceil.comp +0 -22
  284. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/clamp.comp +0 -17
  285. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/cos.comp +0 -17
  286. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/elu.comp +0 -27
  287. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/exp.comp +0 -20
  288. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/floor.comp +0 -22
  289. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu.comp +0 -25
  290. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu_erf.comp +0 -39
  291. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu_quick.comp +0 -23
  292. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/hardsigmoid.comp +0 -22
  293. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/hardswish.comp +0 -22
  294. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/leaky_relu.comp +0 -22
  295. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/neg.comp +0 -20
  296. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/relu.comp +0 -21
  297. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/round.comp +0 -29
  298. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sgn.comp +0 -21
  299. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sigmoid.comp +0 -20
  300. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/silu.comp +0 -22
  301. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sin.comp +0 -17
  302. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/softplus.comp +0 -23
  303. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sqrt.comp +0 -17
  304. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/square.comp +0 -17
  305. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/step.comp +0 -22
  306. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/tanh.comp +0 -20
  307. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/trunc.comp +0 -22
  308. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/xielu.comp +0 -35
@@ -32,8 +32,8 @@ fn inner_dot(src0_val: SRC0_TYPE, src1_val: SRC1_TYPE) -> f32 {
32
32
  #endif
33
33
 
34
34
  #ifdef MUL_ACC_FLOAT
35
- fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
36
- var acc: array<f32, OUTPUTS_PER_WG>;
35
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
36
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
37
37
 
38
38
  let k_vec = params.k / VEC_SIZE;
39
39
  let src1_idx_base_vec = src1_idx_base / VEC_SIZE;
@@ -41,12 +41,18 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
41
41
  // Each thread walks K, loads from the vector, and updates
42
42
  // a small block of output rows held in registers.
43
43
  for (var k = thread_id; k < k_vec; k += WG_SIZE) {
44
- let x = src1[src1_idx_base_vec + k];
44
+ var x_vals: array<SRC1_TYPE, NUM_COLS>;
45
+ for (var col = 0u;col < NUM_COLS;col += 1) {
46
+ x_vals[col] = src1[src1_idx_base_vec + col * (params.stride_11 / VEC_SIZE) + k];
47
+ }
45
48
  for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
46
49
  let output_row = row_base + row;
47
50
  if (output_row < params.m) {
48
51
  let src0_idx = (src0_batch_offset + output_row * params.stride_01) / VEC_SIZE + k;
49
- acc[row] += inner_dot(src0[src0_idx], x);
52
+ let w = src0[src0_idx];
53
+ for (var col = 0u;col < NUM_COLS;col += 1) {
54
+ acc[col][row] += inner_dot(w, x_vals[col]);
55
+ }
50
56
  }
51
57
  }
52
58
  }
@@ -60,30 +66,33 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
60
66
  #define BLOCK_SIZE_BYTES 18
61
67
  #define THREADS_PER_BLOCK 16
62
68
  #define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
63
- fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
64
- var acc: array<f32, OUTPUTS_PER_WG>;
69
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
70
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
65
71
 
66
72
  let num_blocks = params.k / BLOCK_SIZE;
67
73
  let thread_within_block = thread_id % THREADS_PER_BLOCK;
68
74
  for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
69
75
  let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * ELEMS_PER_THREAD;
70
- var x_block: array<f32, ELEMS_PER_THREAD>;
71
- for (var i = 0u; i < ELEMS_PER_THREAD; i++) {
72
- x_block[i] = f32(src1[x_base + i]);
76
+ var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
77
+ for (var col = 0u; col < NUM_COLS;col += 1) {
78
+ for (var i = 0u; i < ELEMS_PER_THREAD; i++) {
79
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
80
+ }
73
81
  }
74
-
75
82
  for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
76
83
  let output_row = row_base + row;
77
84
  if (output_row < params.m) {
78
85
  let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
79
86
  let d = f32(load_f16_at_src0(block_byte_base));
80
87
  let q_byte = load_u32_at_src0(block_byte_base + 2u + thread_within_block) & 0xFFu;
81
- var row_sum = 0.0;
82
- for (var bit = 0u; bit < 8u; bit++) {
83
- let w = select(-d, d, ((q_byte >> bit) & 1u) != 0u);
84
- row_sum += w * x_block[bit];
88
+ for (var col = 0u;col < NUM_COLS;col += 1) {
89
+ var row_sum = 0.0;
90
+ for (var bit = 0u; bit < 8u; bit++) {
91
+ let w = select(-d, d, ((q_byte >> bit) & 1u) != 0u);
92
+ row_sum += w * x_block[col][bit];
93
+ }
94
+ acc[col][row] += row_sum;
85
95
  }
86
- acc[row] += row_sum;
87
96
  }
88
97
  }
89
98
  }
@@ -97,35 +106,37 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
97
106
  #define BLOCK_SIZE_BYTES 18
98
107
  #define THREADS_PER_BLOCK 4
99
108
  #define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
100
- fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
101
- var acc: array<f32, OUTPUTS_PER_WG>;
109
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
110
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
102
111
 
103
112
  let num_blocks = params.k / BLOCK_SIZE;
104
113
  let thread_within_block = thread_id % 4;
105
114
  for (var block = thread_id/THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE/THREADS_PER_BLOCK) {
106
115
  let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * 4;
107
- var x_block: array<f32, ELEMS_PER_THREAD>;
108
- for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
109
- x_block[i] = f32(src1[x_base + i]);
110
- x_block[i + 4] = f32(src1[x_base + i + 16]);
116
+ var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
117
+ for (var col = 0u; col < NUM_COLS;col += 1) {
118
+ for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
119
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
120
+ x_block[col][i + 4] = f32(src1[x_base + col * params.stride_11 + i + 16]);
121
+ }
111
122
  }
112
-
113
123
  for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
114
124
  let output_row = row_base + row;
115
125
  if (output_row < params.m) {
116
126
  let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
117
127
  let d = f32(load_f16_at_src0(block_byte_base));
118
- var row_sum = 0.0;
119
-
120
128
  let q_packed = load_u32_at_src0(block_byte_base + 2u + 4u * thread_within_block);
121
- for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
122
- let q_byte = get_byte(q_packed, byte_idx);
123
- let q_lo = (f32(q_byte & 0xFu) - 8.0) * d;
124
- let q_hi = (f32((q_byte >> 4u) & 0xFu) - 8.0) * d;
125
- row_sum += q_lo * x_block[byte_idx];
126
- row_sum += q_hi * x_block[byte_idx + 4u];
129
+ for (var col = 0u;col < NUM_COLS;col += 1) {
130
+ var row_sum = 0.0;
131
+ for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
132
+ let q_byte = get_byte(q_packed, byte_idx);
133
+ let q_lo = (f32(q_byte & 0xFu) - 8.0) * d;
134
+ let q_hi = (f32((q_byte >> 4u) & 0xFu) - 8.0) * d;
135
+ row_sum += q_lo * x_block[col][byte_idx];
136
+ row_sum += q_hi * x_block[col][byte_idx + 4u];
137
+ }
138
+ acc[col][row] += row_sum;
127
139
  }
128
- acc[row] += row_sum;
129
140
  }
130
141
  }
131
142
  }
@@ -139,36 +150,38 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
139
150
  #define BLOCK_SIZE_BYTES 20
140
151
  #define THREADS_PER_BLOCK 4
141
152
  #define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
142
- fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
143
- var acc: array<f32, OUTPUTS_PER_WG>;
153
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
154
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
144
155
 
145
156
  let num_blocks = params.k / BLOCK_SIZE;
146
157
  let thread_within_block = thread_id % THREADS_PER_BLOCK;
147
158
  for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
148
159
  let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * 4;
149
- var x_block: array<f32, ELEMS_PER_THREAD>;
150
- for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
151
- x_block[i] = f32(src1[x_base + i]);
152
- x_block[i + 4] = f32(src1[x_base + i + 16]);
160
+ var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
161
+ for (var col = 0u; col < NUM_COLS;col += 1) {
162
+ for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
163
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
164
+ x_block[col][i + 4] = f32(src1[x_base + col * params.stride_11 + i + 16]);
165
+ }
153
166
  }
154
-
155
167
  for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
156
168
  let output_row = row_base + row;
157
169
  if (output_row < params.m) {
158
170
  let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
159
171
  let d = f32(load_f16_at_src0(block_byte_base));
160
172
  let m = f32(load_f16_at_src0(block_byte_base + 2u));
161
- var row_sum = 0.0;
162
-
163
173
  let q_packed = load_u32_at_src0(block_byte_base + 4u + 4u * thread_within_block);
164
- for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
165
- let q_byte = get_byte(q_packed, byte_idx);
166
- let q_lo = f32(q_byte & 0xFu) * d + m;
167
- let q_hi = f32((q_byte >> 4u) & 0xFu) * d + m;
168
- row_sum += q_lo * x_block[byte_idx];
169
- row_sum += q_hi * x_block[byte_idx + 4u];
174
+ for (var col = 0u;col < NUM_COLS;col += 1) {
175
+ var row_sum = 0.0;
176
+ for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
177
+ let q_byte = get_byte(q_packed, byte_idx);
178
+ let q_lo = f32(q_byte & 0xFu) * d + m;
179
+ let q_hi = f32((q_byte >> 4u) & 0xFu) * d + m;
180
+ row_sum += q_lo * x_block[col][byte_idx];
181
+ row_sum += q_hi * x_block[col][byte_idx + 4u];
182
+ }
183
+ acc[col][row] += row_sum;
170
184
  }
171
- acc[row] += row_sum;
172
185
  }
173
186
  }
174
187
  }
@@ -182,19 +195,20 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
182
195
  #define BLOCK_SIZE_BYTES 22
183
196
  #define THREADS_PER_BLOCK 4
184
197
  #define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
185
- fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
186
- var acc: array<f32, OUTPUTS_PER_WG>;
198
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
199
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
187
200
 
188
201
  let num_blocks = params.k / BLOCK_SIZE;
189
202
  let thread_within_block = thread_id % THREADS_PER_BLOCK;
190
203
  for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
191
204
  let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * 4;
192
- var x_block: array<f32, ELEMS_PER_THREAD>;
193
- for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
194
- x_block[i] = f32(src1[x_base + i]);
195
- x_block[i + 4] = f32(src1[x_base + i + 16]);
205
+ var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
206
+ for (var col = 0u; col < NUM_COLS;col += 1) {
207
+ for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
208
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
209
+ x_block[col][i + 4] = f32(src1[x_base + col * params.stride_11 + i + 16]);
210
+ }
196
211
  }
197
-
198
212
  for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
199
213
  let output_row = row_base + row;
200
214
  if (output_row < params.m) {
@@ -203,18 +217,19 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
203
217
  let qh_packed = load_u32_at_src0(block_byte_base + 2u);
204
218
  let q_packed = load_u32_at_src0(block_byte_base + 6u + 4u * thread_within_block);
205
219
  let qh_shift = thread_within_block * 4u;
206
- var row_sum = 0.0;
207
-
208
- for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
209
- let q_byte = get_byte(q_packed, byte_idx);
210
- let qh_lo = ((qh_packed >> (qh_shift + byte_idx)) << 4u) & 0x10u;
211
- let qh_hi = (qh_packed >> (qh_shift + byte_idx + 12u)) & 0x10u;
212
- let q_lo = (f32((q_byte & 0xFu) | qh_lo) - 16.0) * d;
213
- let q_hi = (f32(((q_byte >> 4u) & 0xFu) | qh_hi) - 16.0) * d;
214
- row_sum += q_lo * x_block[byte_idx];
215
- row_sum += q_hi * x_block[byte_idx + 4u];
220
+ for (var col = 0u;col < NUM_COLS;col += 1) {
221
+ var row_sum = 0.0;
222
+ for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
223
+ let q_byte = get_byte(q_packed, byte_idx);
224
+ let qh_lo = ((qh_packed >> (qh_shift + byte_idx)) << 4u) & 0x10u;
225
+ let qh_hi = (qh_packed >> (qh_shift + byte_idx + 12u)) & 0x10u;
226
+ let q_lo = (f32((q_byte & 0xFu) | qh_lo) - 16.0) * d;
227
+ let q_hi = (f32(((q_byte >> 4u) & 0xFu) | qh_hi) - 16.0) * d;
228
+ row_sum += q_lo * x_block[col][byte_idx];
229
+ row_sum += q_hi * x_block[col][byte_idx + 4u];
230
+ }
231
+ acc[col][row] += row_sum;
216
232
  }
217
- acc[row] += row_sum;
218
233
  }
219
234
  }
220
235
  }
@@ -228,19 +243,20 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
228
243
  #define BLOCK_SIZE_BYTES 24
229
244
  #define THREADS_PER_BLOCK 4
230
245
  #define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
231
- fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
232
- var acc: array<f32, OUTPUTS_PER_WG>;
246
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
247
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
233
248
 
234
249
  let num_blocks = params.k / BLOCK_SIZE;
235
250
  let thread_within_block = thread_id % THREADS_PER_BLOCK;
236
251
  for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
237
252
  let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * 4;
238
- var x_block: array<f32, ELEMS_PER_THREAD>;
239
- for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
240
- x_block[i] = f32(src1[x_base + i]);
241
- x_block[i + 4] = f32(src1[x_base + i + 16]);
253
+ var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
254
+ for (var col = 0u; col < NUM_COLS;col += 1) {
255
+ for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
256
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
257
+ x_block[col][i + 4] = f32(src1[x_base + col * params.stride_11 + i + 16]);
258
+ }
242
259
  }
243
-
244
260
  for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
245
261
  let output_row = row_base + row;
246
262
  if (output_row < params.m) {
@@ -250,18 +266,19 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
250
266
  let qh_packed = load_u32_at_src0(block_byte_base + 4u);
251
267
  let q_packed = load_u32_at_src0(block_byte_base + 8u + 4u * thread_within_block);
252
268
  let qh_shift = thread_within_block * 4u;
253
- var row_sum = 0.0;
254
-
255
- for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
256
- let q_byte = get_byte(q_packed, byte_idx);
257
- let qh_lo = ((qh_packed >> (qh_shift + byte_idx)) << 4u) & 0x10u;
258
- let qh_hi = (qh_packed >> (qh_shift + byte_idx + 12u)) & 0x10u;
259
- let q_lo = f32((q_byte & 0xFu) | qh_lo) * d + m;
260
- let q_hi = f32(((q_byte >> 4u) & 0xFu) | qh_hi) * d + m;
261
- row_sum += q_lo * x_block[byte_idx];
262
- row_sum += q_hi * x_block[byte_idx + 4u];
269
+ for (var col = 0u;col < NUM_COLS;col += 1) {
270
+ var row_sum = 0.0;
271
+ for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
272
+ let q_byte = get_byte(q_packed, byte_idx);
273
+ let qh_lo = ((qh_packed >> (qh_shift + byte_idx)) << 4u) & 0x10u;
274
+ let qh_hi = (qh_packed >> (qh_shift + byte_idx + 12u)) & 0x10u;
275
+ let q_lo = f32((q_byte & 0xFu) | qh_lo) * d + m;
276
+ let q_hi = f32(((q_byte >> 4u) & 0xFu) | qh_hi) * d + m;
277
+ row_sum += q_lo * x_block[col][byte_idx];
278
+ row_sum += q_hi * x_block[col][byte_idx + 4u];
279
+ }
280
+ acc[col][row] += row_sum;
263
281
  }
264
- acc[row] += row_sum;
265
282
  }
266
283
  }
267
284
  }
@@ -275,33 +292,38 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
275
292
  #define BLOCK_SIZE_BYTES 34
276
293
  #define THREADS_PER_BLOCK 4
277
294
  #define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
278
- fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
279
- var acc: array<f32, OUTPUTS_PER_WG>;
295
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
296
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
280
297
 
281
298
  let num_blocks = params.k / BLOCK_SIZE;
282
299
  let thread_within_block = thread_id % THREADS_PER_BLOCK;
283
300
  for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
284
301
  let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * ELEMS_PER_THREAD;
285
- var x_block: array<f32, ELEMS_PER_THREAD>;
286
- for (var i = 0u; i < ELEMS_PER_THREAD; i++) {
287
- x_block[i] = f32(src1[x_base + i]);
302
+ var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
303
+ for (var col = 0u; col < NUM_COLS;col += 1) {
304
+ for (var i = 0u; i < ELEMS_PER_THREAD; i++) {
305
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
306
+ }
288
307
  }
289
-
290
308
  for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
291
309
  let output_row = row_base + row;
292
310
  if (output_row < params.m) {
293
311
  let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
294
312
  let d = f32(load_f16_at_src0(block_byte_base));
295
- var row_sum = 0.0;
296
-
313
+ var q_packed: array<u32, ELEMS_PER_THREAD / 4u>;
297
314
  for (var packed_idx = 0u; packed_idx < ELEMS_PER_THREAD / 4u; packed_idx++) {
298
- let q_packed = load_u32_at_src0(block_byte_base + 2u + 4u * (thread_within_block * 2u + packed_idx));
299
- for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
300
- let q_val = f32(get_byte_i32(q_packed, byte_idx)) * d;
301
- row_sum += q_val * x_block[packed_idx * 4u + byte_idx];
315
+ q_packed[packed_idx] = load_u32_at_src0(block_byte_base + 2u + 4u * (thread_within_block * 2u + packed_idx));
316
+ }
317
+ for (var col = 0u;col < NUM_COLS;col += 1) {
318
+ var row_sum = 0.0;
319
+ for (var packed_idx = 0u; packed_idx < ELEMS_PER_THREAD / 4u; packed_idx++) {
320
+ for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
321
+ let q_val = f32(get_byte_i32(q_packed[packed_idx], byte_idx)) * d;
322
+ row_sum += q_val * x_block[col][packed_idx * 4u + byte_idx];
323
+ }
302
324
  }
325
+ acc[col][row] += row_sum;
303
326
  }
304
- acc[row] += row_sum;
305
327
  }
306
328
  }
307
329
  }
@@ -315,34 +337,39 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
315
337
  #define BLOCK_SIZE_BYTES 36
316
338
  #define THREADS_PER_BLOCK 4
317
339
  #define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
318
- fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
319
- var acc: array<f32, OUTPUTS_PER_WG>;
340
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
341
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
320
342
 
321
343
  let num_blocks = params.k / BLOCK_SIZE;
322
344
  let thread_within_block = thread_id % THREADS_PER_BLOCK;
323
345
  for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
324
346
  let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * ELEMS_PER_THREAD;
325
- var x_block: array<f32, ELEMS_PER_THREAD>;
326
- for (var i = 0u; i < ELEMS_PER_THREAD; i++) {
327
- x_block[i] = f32(src1[x_base + i]);
347
+ var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
348
+ for (var col = 0u; col < NUM_COLS;col += 1) {
349
+ for (var i = 0u; i < ELEMS_PER_THREAD; i++) {
350
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
351
+ }
328
352
  }
329
-
330
353
  for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
331
354
  let output_row = row_base + row;
332
355
  if (output_row < params.m) {
333
356
  let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
334
357
  let d = f32(load_f16_at_src0(block_byte_base));
335
358
  let m = f32(load_f16_at_src0(block_byte_base + 2u));
336
- var row_sum = 0.0;
337
-
359
+ var q_packed: array<u32, ELEMS_PER_THREAD / 4u>;
338
360
  for (var packed_idx = 0u; packed_idx < ELEMS_PER_THREAD / 4u; packed_idx++) {
339
- let q_packed = load_u32_at_src0(block_byte_base + 4u + 4u * (thread_within_block * 2u + packed_idx));
340
- for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
341
- let q_val = f32(get_byte_i32(q_packed, byte_idx)) * d + m;
342
- row_sum += q_val * x_block[packed_idx * 4u + byte_idx];
361
+ q_packed[packed_idx] = load_u32_at_src0(block_byte_base + 4u + 4u * (thread_within_block * 2u + packed_idx));
362
+ }
363
+ for (var col = 0u;col < NUM_COLS;col += 1) {
364
+ var row_sum = 0.0;
365
+ for (var packed_idx = 0u; packed_idx < ELEMS_PER_THREAD / 4u; packed_idx++) {
366
+ for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
367
+ let q_val = f32(get_byte_i32(q_packed[packed_idx], byte_idx)) * d + m;
368
+ row_sum += q_val * x_block[col][packed_idx * 4u + byte_idx];
369
+ }
343
370
  }
371
+ acc[col][row] += row_sum;
344
372
  }
345
- acc[row] += row_sum;
346
373
  }
347
374
  }
348
375
  }
@@ -355,8 +382,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
355
382
  #define BLOCK_SIZE 256
356
383
  #define BLOCK_SIZE_BYTES 84
357
384
  #define THREADS_PER_BLOCK 16
358
- fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
359
- var acc: array<f32, OUTPUTS_PER_WG>;
385
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
386
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
360
387
 
361
388
  let tid = thread_id % THREADS_PER_BLOCK;
362
389
  let block_group = thread_id / THREADS_PER_BLOCK;
@@ -379,14 +406,15 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
379
406
 
380
407
  for (var block = block_group; block < num_blocks; block += num_block_groups) {
381
408
  let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
382
- var x_block: array<f32, 16>;
383
- for (var i = 0u; i < 4u; i++) {
384
- x_block[i] = f32(src1[x_base + i]);
385
- x_block[i + 4u] = f32(src1[x_base + 32u + i]);
386
- x_block[i + 8u] = f32(src1[x_base + 64u + i]);
387
- x_block[i + 12u] = f32(src1[x_base + 96u + i]);
409
+ var x_block: array<array<f32, 16>, NUM_COLS>;
410
+ for (var col = 0u; col < NUM_COLS;col += 1) {
411
+ for (var i = 0u; i < 4u; i++) {
412
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
413
+ x_block[col][i + 4u] = f32(src1[x_base + col * params.stride_11 + 32u + i]);
414
+ x_block[col][i + 8u] = f32(src1[x_base + col * params.stride_11 + 64u + i]);
415
+ x_block[col][i + 12u] = f32(src1[x_base + col * params.stride_11 + 96u + i]);
416
+ }
388
417
  }
389
-
390
418
  for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
391
419
  let output_row = row_base + row;
392
420
  if (output_row < params.m) {
@@ -404,30 +432,32 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
404
432
  let qs0 = q_u32 & 0xFFFFu;
405
433
  let qs1 = q_u32 >> 16u;
406
434
 
407
- var sumy = vec4<f32>(0.0, 0.0, 0.0, 0.0);
408
- var acc1 = vec4<f32>(0.0, 0.0, 0.0, 0.0);
409
- var acc2 = vec4<f32>(0.0, 0.0, 0.0, 0.0);
410
-
411
- sumy[0] = x_block[0] + x_block[1] + x_block[2] + x_block[3];
412
- sumy[1] = x_block[4] + x_block[5] + x_block[6] + x_block[7];
413
- sumy[2] = x_block[8] + x_block[9] + x_block[10] + x_block[11];
414
- sumy[3] = x_block[12] + x_block[13] + x_block[14] + x_block[15];
415
-
416
- acc1[0] = x_block[0] * f32(qs0 & 0x0003u) + x_block[2] * f32(qs1 & 0x0003u);
417
- acc2[0] = x_block[1] * f32(qs0 & 0x0300u) + x_block[3] * f32(qs1 & 0x0300u);
418
- acc1[1] = x_block[4] * f32(qs0 & 0x000Cu) + x_block[6] * f32(qs1 & 0x000Cu);
419
- acc2[1] = x_block[5] * f32(qs0 & 0x0C00u) + x_block[7] * f32(qs1 & 0x0C00u);
420
- acc1[2] = x_block[8] * f32(qs0 & 0x0030u) + x_block[10] * f32(qs1 & 0x0030u);
421
- acc2[2] = x_block[9] * f32(qs0 & 0x3000u) + x_block[11] * f32(qs1 & 0x3000u);
422
- acc1[3] = x_block[12] * f32(qs0 & 0x00C0u) + x_block[14] * f32(qs1 & 0x00C0u);
423
- acc2[3] = x_block[13] * f32(qs0 & 0xC000u) + x_block[15] * f32(qs1 & 0xC000u);
424
-
425
- acc[row] += dall * ((acc1[0] + (1.0/256.0) * acc2[0]) * f32(sc0 & 0xFu) +
426
- (acc1[1] + (1.0/256.0) * acc2[1]) * f32(sc2 & 0xFu) / 4.0 +
427
- (acc1[2] + (1.0/256.0) * acc2[2]) * f32(sc4 & 0xFu) / 16.0 +
428
- (acc1[3] + (1.0/256.0) * acc2[3]) * f32(sc6 & 0xFu) / 64.0)
429
- - dmin * (sumy[0] * f32(sc0 & 0xF0u) + sumy[1] * f32(sc2 & 0xF0u) +
430
- sumy[2] * f32(sc4 & 0xF0u) + sumy[3] * f32(sc6 & 0xF0u));
435
+ for (var col = 0u;col < NUM_COLS;col += 1) {
436
+ var sumy = vec4<f32>(0.0, 0.0, 0.0, 0.0);
437
+ var acc1 = vec4<f32>(0.0, 0.0, 0.0, 0.0);
438
+ var acc2 = vec4<f32>(0.0, 0.0, 0.0, 0.0);
439
+
440
+ sumy[0] = x_block[col][0] + x_block[col][1] + x_block[col][2] + x_block[col][3];
441
+ sumy[1] = x_block[col][4] + x_block[col][5] + x_block[col][6] + x_block[col][7];
442
+ sumy[2] = x_block[col][8] + x_block[col][9] + x_block[col][10] + x_block[col][11];
443
+ sumy[3] = x_block[col][12] + x_block[col][13] + x_block[col][14] + x_block[col][15];
444
+
445
+ acc1[0] = x_block[col][0] * f32(qs0 & 0x0003u) + x_block[col][2] * f32(qs1 & 0x0003u);
446
+ acc2[0] = x_block[col][1] * f32(qs0 & 0x0300u) + x_block[col][3] * f32(qs1 & 0x0300u);
447
+ acc1[1] = x_block[col][4] * f32(qs0 & 0x000Cu) + x_block[col][6] * f32(qs1 & 0x000Cu);
448
+ acc2[1] = x_block[col][5] * f32(qs0 & 0x0C00u) + x_block[col][7] * f32(qs1 & 0x0C00u);
449
+ acc1[2] = x_block[col][8] * f32(qs0 & 0x0030u) + x_block[col][10] * f32(qs1 & 0x0030u);
450
+ acc2[2] = x_block[col][9] * f32(qs0 & 0x3000u) + x_block[col][11] * f32(qs1 & 0x3000u);
451
+ acc1[3] = x_block[col][12] * f32(qs0 & 0x00C0u) + x_block[col][14] * f32(qs1 & 0x00C0u);
452
+ acc2[3] = x_block[col][13] * f32(qs0 & 0xC000u) + x_block[col][15] * f32(qs1 & 0xC000u);
453
+
454
+ acc[col][row] += dall * ((acc1[0] + (1.0/256.0) * acc2[0]) * f32(sc0 & 0xFu) +
455
+ (acc1[1] + (1.0/256.0) * acc2[1]) * f32(sc2 & 0xFu) / 4.0 +
456
+ (acc1[2] + (1.0/256.0) * acc2[2]) * f32(sc4 & 0xFu) / 16.0 +
457
+ (acc1[3] + (1.0/256.0) * acc2[3]) * f32(sc6 & 0xFu) / 64.0)
458
+ - dmin * (sumy[0] * f32(sc0 & 0xF0u) + sumy[1] * f32(sc2 & 0xF0u) +
459
+ sumy[2] * f32(sc4 & 0xF0u) + sumy[3] * f32(sc6 & 0xF0u));
460
+ }
431
461
  }
432
462
  }
433
463
  }
@@ -440,8 +470,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
440
470
  #define BLOCK_SIZE 256
441
471
  #define BLOCK_SIZE_BYTES 110
442
472
  #define THREADS_PER_BLOCK 16
443
- fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
444
- var acc: array<f32, OUTPUTS_PER_WG>;
473
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
474
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
445
475
 
446
476
  let tid = thread_id % THREADS_PER_BLOCK;
447
477
  let block_group = thread_id / THREADS_PER_BLOCK;
@@ -485,12 +515,13 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
485
515
 
486
516
  for (var block = block_group; block < num_blocks; block += num_block_groups) {
487
517
  let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
488
- var x_block: array<f32, 16>;
489
- for (var i = 0u; i < 8u; i++) {
490
- x_block[i] = f32(src1[x_base + i]);
491
- x_block[i + 8u] = f32(src1[x_base + 32u + i]);
518
+ var x_block: array<array<f32, 16>, NUM_COLS>;
519
+ for (var col = 0u; col < NUM_COLS;col += 1) {
520
+ for (var i = 0u; i < 8u; i++) {
521
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
522
+ x_block[col][i + 8u] = f32(src1[x_base + col * params.stride_11 + 32u + i]);
523
+ }
492
524
  }
493
-
494
525
  for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
495
526
  let output_row = row_base + row;
496
527
  if (output_row < params.m) {
@@ -516,28 +547,30 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
516
547
  let h_u32_0 = load_u32_at_src0(block_byte_base + h_byte + 0u);
517
548
  let h_u32_1 = load_u32_at_src0(block_byte_base + h_byte + 4u);
518
549
 
519
- var s1 = 0.0; var s2 = 0.0; var s3 = 0.0;
520
- var s4 = 0.0; var s5 = 0.0; var s6 = 0.0;
521
-
522
- for (var l = 0u; l < 8u; l += 2u) {
523
- let q_u32 = select(q_u32_0, q_u32_1, l >= 4u);
524
- let qs = select(q_u32 & 0xFFFFu, q_u32 >> 16u, (l & 2u) != 0u);
525
- let h_u32 = select(h_u32_0, h_u32_1, l >= 4u);
526
- let hv = select(h_u32 & 0xFFFFu, h_u32 >> 16u, (l & 2u) != 0u);
527
-
528
- s1 += x_block[l + 0u] * f32(qs & qm0);
529
- s2 += x_block[l + 1u] * f32(qs & qm1);
530
- s3 += select(0.0, x_block[l + 0u], (hv & hm0) == 0u) +
531
- select(0.0, x_block[l + 1u], (hv & hm1) == 0u);
532
- s4 += x_block[l + 8u] * f32(qs & qm2);
533
- s5 += x_block[l + 9u] * f32(qs & qm3);
534
- s6 += select(0.0, x_block[l + 8u], (hv & hm2) == 0u) +
535
- select(0.0, x_block[l + 9u], (hv & hm3) == 0u);
536
- }
550
+ for (var col = 0u;col < NUM_COLS;col += 1) {
551
+ var s1 = 0.0; var s2 = 0.0; var s3 = 0.0;
552
+ var s4 = 0.0; var s5 = 0.0; var s6 = 0.0;
553
+
554
+ for (var l = 0u; l < 8u; l += 2u) {
555
+ let q_u32 = select(q_u32_0, q_u32_1, l >= 4u);
556
+ let qs = select(q_u32 & 0xFFFFu, q_u32 >> 16u, (l & 2u) != 0u);
557
+ let h_u32 = select(h_u32_0, h_u32_1, l >= 4u);
558
+ let hv = select(h_u32 & 0xFFFFu, h_u32 >> 16u, (l & 2u) != 0u);
559
+
560
+ s1 += x_block[col][l + 0u] * f32(qs & qm0);
561
+ s2 += x_block[col][l + 1u] * f32(qs & qm1);
562
+ s3 += select(0.0, x_block[col][l + 0u], (hv & hm0) == 0u) +
563
+ select(0.0, x_block[col][l + 1u], (hv & hm1) == 0u);
564
+ s4 += x_block[col][l + 8u] * f32(qs & qm2);
565
+ s5 += x_block[col][l + 9u] * f32(qs & qm3);
566
+ s6 += select(0.0, x_block[col][l + 8u], (hv & hm2) == 0u) +
567
+ select(0.0, x_block[col][l + 9u], (hv & hm3) == 0u);
568
+ }
537
569
 
538
- let d1 = d * (s1 + (1.0/256.0) * s2 - s3 * v1);
539
- let d2 = d * (s4 + (1.0/256.0) * s5 - s6 * v2);
540
- acc[row] += (d1 * scale0 + 0.25 * d2 * scale1) / f32(1u << shift);
570
+ let d1 = d * (s1 + (1.0/256.0) * s2 - s3 * v1);
571
+ let d2 = d * (s4 + (1.0/256.0) * s5 - s6 * v2);
572
+ acc[col][row] += (d1 * scale0 + 0.25 * d2 * scale1) / f32(1u << shift);
573
+ }
541
574
  }
542
575
  }
543
576
  }
@@ -550,8 +583,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
550
583
  #define BLOCK_SIZE 256
551
584
  #define BLOCK_SIZE_BYTES 144
552
585
  #define THREADS_PER_BLOCK 16
553
- fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
554
- var acc: array<f32, OUTPUTS_PER_WG>;
586
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
587
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
555
588
 
556
589
  let tid = thread_id % THREADS_PER_BLOCK;
557
590
  let block_group = thread_id / THREADS_PER_BLOCK;
@@ -573,12 +606,15 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
573
606
 
574
607
  for (var block = block_group; block < num_blocks; block += num_block_groups) {
575
608
  let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
576
- var x_block: array<f32, 16>;
577
- for (var i = 0u; i < 4u; i++) {
578
- x_block[i] = f32(src1[x_base + i]);
579
- x_block[i + 4u] = f32(src1[x_base + 32u + i]);
580
- x_block[i + 8u] = f32(src1[x_base + 128u + i]);
581
- x_block[i + 12u] = f32(src1[x_base + 160u + i]);
609
+ var x_block: array<array<f32, 16>, NUM_COLS>;
610
+ for (var col = 0u; col < NUM_COLS;col += 1) {
611
+ let col_base = x_base + col * params.stride_11;
612
+ for (var i = 0u; i < 4u; i++) {
613
+ x_block[col][i] = f32(src1[col_base + i]);
614
+ x_block[col][i + 4u] = f32(src1[col_base + 32u + i]);
615
+ x_block[col][i + 8u] = f32(src1[col_base + 128u + i]);
616
+ x_block[col][i + 12u] = f32(src1[col_base + 160u + i]);
617
+ }
582
618
  }
583
619
 
584
620
  for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
@@ -613,23 +649,25 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
613
649
  let q1_u32 = load_u32_at_src0_aligned(block_byte_base + 16u + q_offset);
614
650
  let q2_u32 = load_u32_at_src0_aligned(block_byte_base + 80u + q_offset);
615
651
 
616
- var dot = vec4<f32>(0.0, 0.0, 0.0, 0.0);
617
- var sumx = vec4<f32>(0.0, 0.0, 0.0, 0.0);
618
- for (var i = 0u; i < 4u; i++) {
619
- let q1b = byte_of(q1_u32, i);
620
- let q2b = byte_of(q2_u32, i);
621
- dot[0] += x_block[i] * f32(q1b & 0x0Fu);
622
- dot[1] += x_block[i + 4u] * f32(q1b >> 4u);
623
- dot[2] += x_block[i + 8u] * f32(q2b & 0x0Fu);
624
- dot[3] += x_block[i + 12u] * f32(q2b >> 4u);
625
- sumx[0] += x_block[i];
626
- sumx[1] += x_block[i + 4u];
627
- sumx[2] += x_block[i + 8u];
628
- sumx[3] += x_block[i + 12u];
629
- }
652
+ for (var col = 0u;col < NUM_COLS;col += 1) {
653
+ var dot = vec4<f32>(0.0, 0.0, 0.0, 0.0);
654
+ var sumx = vec4<f32>(0.0, 0.0, 0.0, 0.0);
655
+ for (var i = 0u; i < 4u; i++) {
656
+ let q1b = byte_of(q1_u32, i);
657
+ let q2b = byte_of(q2_u32, i);
658
+ dot[0] += x_block[col][i] * f32(q1b & 0x0Fu);
659
+ dot[1] += x_block[col][i + 4u] * f32(q1b >> 4u);
660
+ dot[2] += x_block[col][i + 8u] * f32(q2b & 0x0Fu);
661
+ dot[3] += x_block[col][i + 12u] * f32(q2b >> 4u);
662
+ sumx[0] += x_block[col][i];
663
+ sumx[1] += x_block[col][i + 4u];
664
+ sumx[2] += x_block[col][i + 8u];
665
+ sumx[3] += x_block[col][i + 12u];
666
+ }
630
667
 
631
- acc[row] += d * (dot[0] * scale0 + dot[1] * scale1 + dot[2] * scale2 + dot[3] * scale3)
632
- - dmin * (sumx[0] * min0 + sumx[1] * min1 + sumx[2] * min2 + sumx[3] * min3);
668
+ acc[col][row] += d * (dot[0] * scale0 + dot[1] * scale1 + dot[2] * scale2 + dot[3] * scale3)
669
+ - dmin * (sumx[0] * min0 + sumx[1] * min1 + sumx[2] * min2 + sumx[3] * min3);
670
+ }
633
671
  }
634
672
  }
635
673
  }
@@ -642,8 +680,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
642
680
  #define BLOCK_SIZE 256
643
681
  #define BLOCK_SIZE_BYTES 176
644
682
  #define THREADS_PER_BLOCK 16
645
- fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
646
- var acc: array<f32, OUTPUTS_PER_WG>;
683
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
684
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
647
685
 
648
686
  let tid = thread_id % THREADS_PER_BLOCK;
649
687
  let block_group = thread_id / THREADS_PER_BLOCK;
@@ -671,14 +709,16 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
671
709
 
672
710
  for (var block = block_group; block < num_blocks; block += num_block_groups) {
673
711
  let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
674
- var x_block: array<f32, 16>;
675
- for (var i = 0u; i < 4u; i++) {
676
- x_block[i] = f32(src1[x_base + i]);
677
- x_block[i + 4u] = f32(src1[x_base + 32u + i]);
678
- x_block[i + 8u] = f32(src1[x_base + 128u + i]);
679
- x_block[i + 12u] = f32(src1[x_base + 160u + i]);
712
+ var x_block: array<array<f32, 16>, NUM_COLS>;
713
+ for (var col = 0u; col < NUM_COLS;col += 1) {
714
+ let col_base = x_base + col * params.stride_11;
715
+ for (var i = 0u; i < 4u; i++) {
716
+ x_block[col][i] = f32(src1[col_base + i]);
717
+ x_block[col][i + 4u] = f32(src1[col_base + 32u + i]);
718
+ x_block[col][i + 8u] = f32(src1[col_base + 128u + i]);
719
+ x_block[col][i + 12u] = f32(src1[col_base + 160u + i]);
720
+ }
680
721
  }
681
-
682
722
  for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
683
723
  let output_row = row_base + row;
684
724
  if (output_row < params.m) {
@@ -712,37 +752,39 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
712
752
  let q2_u32 = load_u32_at_src0_aligned(block_byte_base + q_offset + 64u);
713
753
  let qh_u32 = load_u32_at_src0_aligned(block_byte_base + qh_offset);
714
754
 
715
- var vals = vec4<f32>(0.0, 0.0, 0.0, 0.0);
716
- var sumy = vec4<f32>(0.0, 0.0, 0.0, 0.0);
717
- for (var i = 0u; i < 4u; i++) {
718
- let q1b = byte_of(q1_u32, i);
719
- let q2b = byte_of(q2_u32, i);
720
- let qhb = byte_of(qh_u32, i);
721
-
722
- let yl0 = x_block[i];
723
- let yl8 = x_block[i + 4u];
724
- let yh0 = x_block[i + 8u];
725
- let yh8 = x_block[i + 12u];
726
-
727
- sumy[0] += yl0;
728
- sumy[1] += yl8;
729
- sumy[2] += yh0;
730
- sumy[3] += yh8;
731
-
732
- let q0 = f32((q1b & 0x0Fu) | select(0u, 0x10u, (qhb & hm1) != 0u));
733
- let q1 = f32((q1b >> 4u) | select(0u, 0x10u, (qhb & hm2) != 0u));
734
- let q2 = f32((q2b & 0x0Fu) | select(0u, 0x10u, (qhb & hm3) != 0u));
735
- let q3 = f32((q2b >> 4u) | select(0u, 0x10u, (qhb & hm4) != 0u));
736
-
737
- vals[0] += yl0 * q0;
738
- vals[1] += yl8 * q1;
739
- vals[2] += yh0 * q2;
740
- vals[3] += yh8 * q3;
741
- }
755
+ for (var col = 0u;col < NUM_COLS;col += 1) {
756
+ var vals = vec4<f32>(0.0, 0.0, 0.0, 0.0);
757
+ var sumy = vec4<f32>(0.0, 0.0, 0.0, 0.0);
758
+ for (var i = 0u; i < 4u; i++) {
759
+ let q1b = byte_of(q1_u32, i);
760
+ let q2b = byte_of(q2_u32, i);
761
+ let qhb = byte_of(qh_u32, i);
762
+
763
+ let yl0 = x_block[col][i];
764
+ let yl8 = x_block[col][i + 4u];
765
+ let yh0 = x_block[col][i + 8u];
766
+ let yh8 = x_block[col][i + 12u];
767
+
768
+ sumy[0] += yl0;
769
+ sumy[1] += yl8;
770
+ sumy[2] += yh0;
771
+ sumy[3] += yh8;
772
+
773
+ let q0 = f32((q1b & 0x0Fu) | select(0u, 0x10u, (qhb & hm1) != 0u));
774
+ let q1 = f32((q1b >> 4u) | select(0u, 0x10u, (qhb & hm2) != 0u));
775
+ let q2 = f32((q2b & 0x0Fu) | select(0u, 0x10u, (qhb & hm3) != 0u));
776
+ let q3 = f32((q2b >> 4u) | select(0u, 0x10u, (qhb & hm4) != 0u));
777
+
778
+ vals[0] += yl0 * q0;
779
+ vals[1] += yl8 * q1;
780
+ vals[2] += yh0 * q2;
781
+ vals[3] += yh8 * q3;
782
+ }
742
783
 
743
- acc[row] += d * (f0 * vals[0] + f1 * vals[1] + f4 * vals[2] + f5 * vals[3])
744
- - dmin * (sumy[0] * m0 + sumy[1] * m1 +
745
- sumy[2] * m4 + sumy[3] * m5);
784
+ acc[col][row] += d * (f0 * vals[0] + f1 * vals[1] + f4 * vals[2] + f5 * vals[3])
785
+ - dmin * (sumy[0] * m0 + sumy[1] * m1 +
786
+ sumy[2] * m4 + sumy[3] * m5);
787
+ }
746
788
  }
747
789
  }
748
790
  }
@@ -755,8 +797,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
755
797
  #define BLOCK_SIZE 256
756
798
  #define BLOCK_SIZE_BYTES 210
757
799
  #define THREADS_PER_BLOCK 16
758
- fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
759
- var acc: array<f32, OUTPUTS_PER_WG>;
800
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
801
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
760
802
 
761
803
  let tid = thread_id % THREADS_PER_BLOCK;
762
804
  let block_group = thread_id / THREADS_PER_BLOCK;
@@ -777,14 +819,16 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
777
819
 
778
820
  for (var block = block_group; block < num_blocks; block += num_block_groups) {
779
821
  let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
780
- var x_block: array<f32, 16>;
781
- for (var l = 0u; l < 4u; l++) {
782
- x_block[l] = f32(src1[x_base + l]);
783
- x_block[l + 4u] = f32(src1[x_base + 32u + l]);
784
- x_block[l + 8u] = f32(src1[x_base + 64u + l]);
785
- x_block[l + 12u] = f32(src1[x_base + 96u + l]);
822
+ var x_block: array<array<f32, 16>, NUM_COLS>;
823
+ for (var col = 0u; col < NUM_COLS;col += 1) {
824
+ let col_base = x_base + col * params.stride_11;
825
+ for (var l = 0u; l < 4u; l++) {
826
+ x_block[col][l] = f32(src1[col_base + l]);
827
+ x_block[col][l + 4u] = f32(src1[col_base + 32u + l]);
828
+ x_block[col][l + 8u] = f32(src1[col_base + 64u + l]);
829
+ x_block[col][l + 12u] = f32(src1[col_base + 96u + l]);
830
+ }
786
831
  }
787
-
788
832
  for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
789
833
  let output_row = row_base + row;
790
834
  if (output_row < params.m) {
@@ -802,26 +846,28 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
802
846
  let sc4 = sbyte_of(sc_u32_1, sc_byte_pos);
803
847
  let sc6 = sbyte_of(sc_u32_1, sc_byte_pos + 2u);
804
848
 
805
- var sums = vec4<f32>(0.0, 0.0, 0.0, 0.0);
849
+ for (var col = 0u;col < NUM_COLS;col += 1) {
850
+ var sums = vec4<f32>(0.0, 0.0, 0.0, 0.0);
806
851
 
807
- for (var l = 0u; l < 4u; l++) {
808
- let q1b = byte_of(ql1_u32, l);
809
- let q2b = byte_of(ql2_u32, l);
810
- let qhb = byte_of(qh_u32, l);
852
+ for (var l = 0u; l < 4u; l++) {
853
+ let q1b = byte_of(ql1_u32, l);
854
+ let q2b = byte_of(ql2_u32, l);
855
+ let qhb = byte_of(qh_u32, l);
811
856
 
812
- let dq0 = f32(i32((q1b & 0x0Fu) | ((qhb & 0x03u) << 4u)) - 32);
813
- let dq1 = f32(i32((q2b & 0x0Fu) | ((qhb & 0x0Cu) << 2u)) - 32);
814
- let dq2 = f32(i32((q1b >> 4u) | (qhb & 0x30u)) - 32);
815
- let dq3 = f32(i32((q2b >> 4u) | ((qhb & 0xC0u) >> 2u)) - 32);
857
+ let dq0 = f32(i32((q1b & 0x0Fu) | ((qhb & 0x03u) << 4u)) - 32);
858
+ let dq1 = f32(i32((q2b & 0x0Fu) | ((qhb & 0x0Cu) << 2u)) - 32);
859
+ let dq2 = f32(i32((q1b >> 4u) | (qhb & 0x30u)) - 32);
860
+ let dq3 = f32(i32((q2b >> 4u) | ((qhb & 0xC0u) >> 2u)) - 32);
816
861
 
817
- sums[0] += x_block[l] * dq0;
818
- sums[1] += x_block[l + 4u] * dq1;
819
- sums[2] += x_block[l + 8u] * dq2;
820
- sums[3] += x_block[l + 12u] * dq3;
821
- }
862
+ sums[0] += x_block[col][l] * dq0;
863
+ sums[1] += x_block[col][l + 4u] * dq1;
864
+ sums[2] += x_block[col][l + 8u] * dq2;
865
+ sums[3] += x_block[col][l + 12u] * dq3;
866
+ }
822
867
 
823
- acc[row] += d * (sums[0] * f32(sc0) + sums[1] * f32(sc2) +
824
- sums[2] * f32(sc4) + sums[3] * f32(sc6));
868
+ acc[col][row] += d * (sums[0] * f32(sc0) + sums[1] * f32(sc2) +
869
+ sums[2] * f32(sc4) + sums[3] * f32(sc6));
870
+ }
825
871
  }
826
872
  }
827
873
  }
@@ -834,8 +880,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
834
880
  #define BLOCK_SIZE 256
835
881
  #define BLOCK_SIZE_BYTES 50
836
882
  #define THREADS_PER_BLOCK 16
837
- fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
838
- var acc: array<f32, OUTPUTS_PER_WG>;
883
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
884
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
839
885
 
840
886
  let tid = thread_id % THREADS_PER_BLOCK;
841
887
  let block_group = thread_id / THREADS_PER_BLOCK;
@@ -850,11 +896,12 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
850
896
 
851
897
  for (var block = block_group; block < num_blocks; block += num_block_groups) {
852
898
  let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
853
- var x_block: array<f32, 16>;
854
- for (var i = 0u; i < 16u; i++) {
855
- x_block[i] = f32(src1[x_base + i]);
899
+ var x_block: array<array<f32, 16>, NUM_COLS>;
900
+ for (var col = 0u; col < NUM_COLS;col += 1) {
901
+ for (var i = 0u; i < 16u; i++) {
902
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
903
+ }
856
904
  }
857
-
858
905
  for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
859
906
  let output_row = row_base + row;
860
907
  if (output_row < params.m) {
@@ -866,20 +913,22 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
866
913
  let delta = select(IQ1_DELTA, -IQ1_DELTA, (qh & 0x8000u) != 0u);
867
914
  let qs_w = load_u32_at_src0(block_byte_base + 2u + sub_blk * 4u);
868
915
 
869
- var row_sum = 0.0;
870
- for (var ll = 0u; ll < 2u; ll++) {
871
- let l = slot0 + ll;
872
- let qs_byte = get_byte(qs_w, l);
873
- let ig = (qs_byte | (((qh >> (3u * l)) & 7u) << 8u)) * 8u;
874
- let gw = iq1_grid[ig / 16u];
875
- let bit_base = (ig % 16u) * 2u;
876
- for (var j = 0u; j < 8u; j++) {
877
- let g = (gw >> (bit_base + j * 2u)) & 3u;
878
- let gs = select(f32(g), f32(g) - 4.0, (g & 2u) != 0u);
879
- row_sum += dl * (gs + delta) * x_block[ll * 8u + j];
916
+ for (var col = 0u;col < NUM_COLS;col += 1) {
917
+ var row_sum = 0.0;
918
+ for (var ll = 0u; ll < 2u; ll++) {
919
+ let l = slot0 + ll;
920
+ let qs_byte = get_byte(qs_w, l);
921
+ let ig = (qs_byte | (((qh >> (3u * l)) & 7u) << 8u)) * 8u;
922
+ let gw = iq1_grid[ig / 16u];
923
+ let bit_base = (ig % 16u) * 2u;
924
+ for (var j = 0u; j < 8u; j++) {
925
+ let g = (gw >> (bit_base + j * 2u)) & 3u;
926
+ let gs = select(f32(g), f32(g) - 4.0, (g & 2u) != 0u);
927
+ row_sum += dl * (gs + delta) * x_block[col][ll * 8u + j];
928
+ }
880
929
  }
930
+ acc[col][row] += row_sum;
881
931
  }
882
- acc[row] += row_sum;
883
932
  }
884
933
  }
885
934
  }
@@ -892,8 +941,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
892
941
  #define BLOCK_SIZE 256
893
942
  #define BLOCK_SIZE_BYTES 56
894
943
  #define THREADS_PER_BLOCK 16
895
- fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
896
- var acc: array<f32, OUTPUTS_PER_WG>;
944
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
945
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
897
946
 
898
947
  let tid = thread_id % THREADS_PER_BLOCK;
899
948
  let block_group = thread_id / THREADS_PER_BLOCK;
@@ -908,11 +957,12 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
908
957
 
909
958
  for (var block = block_group; block < num_blocks; block += num_block_groups) {
910
959
  let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
911
- var x_block: array<f32, 16>;
912
- for (var i = 0u; i < 16u; i++) {
913
- x_block[i] = f32(src1[x_base + i]);
960
+ var x_block: array<array<f32, 16>, NUM_COLS>;
961
+ for (var col = 0u; col < NUM_COLS;col += 1) {
962
+ for (var i = 0u; i < 16u; i++) {
963
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
964
+ }
914
965
  }
915
-
916
966
  for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
917
967
  let output_row = row_base + row;
918
968
  if (output_row < params.m) {
@@ -936,26 +986,28 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
936
986
  let qh_lo = qh & 0xFFu;
937
987
  let qh_hi = (qh >> 8u) & 0xFFu;
938
988
 
939
- var row_sum = 0.0;
940
- for (var ll = 0u; ll < 2u; ll++) {
941
- let l = slot0 + ll;
942
- let bit_off = 6u * (sub_blk % 2u) + 3u * (l / 2u);
943
- let sub_scale = (sc_u16 >> bit_off) & 0x7u;
944
- let dl = d * f32(2u * sub_scale + 1u);
945
- let qh_byte = select(qh_lo, qh_hi, l >= 2u);
946
- let ll2 = l % 2u;
947
- let grid_idx = get_byte(qs_w, l) | (((qh_byte >> (4u * ll2)) & 7u) << 8u);
948
- let delta = select(IQ1_DELTA, -IQ1_DELTA, ((qh_byte >> (3u + 4u * ll2)) & 1u) != 0u);
949
- let ig = grid_idx * 8u;
950
- let gw = iq1_grid[ig / 16u];
951
- let bit_base = (ig % 16u) * 2u;
952
- for (var j = 0u; j < 8u; j++) {
953
- let g = (gw >> (bit_base + j * 2u)) & 3u;
954
- let gs = select(f32(g), f32(g) - 4.0, (g & 2u) != 0u);
955
- row_sum += dl * (gs + delta) * x_block[ll * 8u + j];
989
+ for (var col = 0u;col < NUM_COLS;col += 1) {
990
+ var row_sum = 0.0;
991
+ for (var ll = 0u; ll < 2u; ll++) {
992
+ let l = slot0 + ll;
993
+ let bit_off = 6u * (sub_blk % 2u) + 3u * (l / 2u);
994
+ let sub_scale = (sc_u16 >> bit_off) & 0x7u;
995
+ let dl = d * f32(2u * sub_scale + 1u);
996
+ let qh_byte = select(qh_lo, qh_hi, l >= 2u);
997
+ let ll2 = l % 2u;
998
+ let grid_idx = get_byte(qs_w, l) | (((qh_byte >> (4u * ll2)) & 7u) << 8u);
999
+ let delta = select(IQ1_DELTA, -IQ1_DELTA, ((qh_byte >> (3u + 4u * ll2)) & 1u) != 0u);
1000
+ let ig = grid_idx * 8u;
1001
+ let gw = iq1_grid[ig / 16u];
1002
+ let bit_base = (ig % 16u) * 2u;
1003
+ for (var j = 0u; j < 8u; j++) {
1004
+ let g = (gw >> (bit_base + j * 2u)) & 3u;
1005
+ let gs = select(f32(g), f32(g) - 4.0, (g & 2u) != 0u);
1006
+ row_sum += dl * (gs + delta) * x_block[col][ll * 8u + j];
1007
+ }
956
1008
  }
1009
+ acc[col][row] += row_sum;
957
1010
  }
958
- acc[row] += row_sum;
959
1011
  }
960
1012
  }
961
1013
  }
@@ -968,8 +1020,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
968
1020
  #define BLOCK_SIZE 256
969
1021
  #define BLOCK_SIZE_BYTES 66
970
1022
  #define THREADS_PER_BLOCK 16
971
- fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
972
- var acc: array<f32, OUTPUTS_PER_WG>;
1023
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
1024
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
973
1025
 
974
1026
  let tid = thread_id % THREADS_PER_BLOCK;
975
1027
  let block_group = thread_id / THREADS_PER_BLOCK;
@@ -984,11 +1036,12 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
984
1036
 
985
1037
  for (var block = block_group; block < num_blocks; block += num_block_groups) {
986
1038
  let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
987
- var x_block: array<f32, 16>;
988
- for (var i = 0u; i < 16u; i++) {
989
- x_block[i] = f32(src1[x_base + i]);
1039
+ var x_block: array<array<f32, 16>, NUM_COLS>;
1040
+ for (var col = 0u; col < NUM_COLS;col += 1) {
1041
+ for (var i = 0u; i < 16u; i++) {
1042
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
1043
+ }
990
1044
  }
991
-
992
1045
  for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
993
1046
  let output_row = row_base + row;
994
1047
  if (output_row < params.m) {
@@ -999,22 +1052,24 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
999
1052
  let ls = aux_hi >> 28u;
1000
1053
  let db = d * (0.5 + f32(ls)) * 0.25;
1001
1054
 
1002
- var row_sum = 0.0;
1003
- for (var ll = 0u; ll < 2u; ll++) {
1004
- let l = slot0 + ll;
1005
- let grid_idx = (aux_lo >> (8u * l)) & 0xFFu;
1006
- let signs_idx = (aux_hi >> (7u * l)) & 0x7Fu;
1007
- let signs = (ksigns_iq2xs[signs_idx / 4u] >> ((signs_idx % 4u) * 8u)) & 0xFFu;
1008
- let gw_lo = iq2xxs_grid[grid_idx * 2u];
1009
- let gw_hi = iq2xxs_grid[grid_idx * 2u + 1u];
1010
- for (var j = 0u; j < 8u; j++) {
1011
- let gw = select(gw_hi, gw_lo, j < 4u);
1012
- let b = f32((gw >> ((j & 3u) * 8u)) & 0xFFu);
1013
- let s = select(1.0, -1.0, ((signs >> j) & 1u) != 0u);
1014
- row_sum += db * b * s * x_block[ll * 8u + j];
1055
+ for (var col = 0u;col < NUM_COLS;col += 1) {
1056
+ var row_sum = 0.0;
1057
+ for (var ll = 0u; ll < 2u; ll++) {
1058
+ let l = slot0 + ll;
1059
+ let grid_idx = (aux_lo >> (8u * l)) & 0xFFu;
1060
+ let signs_idx = (aux_hi >> (7u * l)) & 0x7Fu;
1061
+ let signs = (ksigns_iq2xs[signs_idx / 4u] >> ((signs_idx % 4u) * 8u)) & 0xFFu;
1062
+ let gw_lo = iq2xxs_grid[grid_idx * 2u];
1063
+ let gw_hi = iq2xxs_grid[grid_idx * 2u + 1u];
1064
+ for (var j = 0u; j < 8u; j++) {
1065
+ let gw = select(gw_hi, gw_lo, j < 4u);
1066
+ let b = f32((gw >> ((j & 3u) * 8u)) & 0xFFu);
1067
+ let s = select(1.0, -1.0, ((signs >> j) & 1u) != 0u);
1068
+ row_sum += db * b * s * x_block[col][ll * 8u + j];
1069
+ }
1015
1070
  }
1071
+ acc[col][row] += row_sum;
1016
1072
  }
1017
- acc[row] += row_sum;
1018
1073
  }
1019
1074
  }
1020
1075
  }
@@ -1027,8 +1082,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
1027
1082
  #define BLOCK_SIZE 256
1028
1083
  #define BLOCK_SIZE_BYTES 74
1029
1084
  #define THREADS_PER_BLOCK 16
1030
- fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
1031
- var acc: array<f32, OUTPUTS_PER_WG>;
1085
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
1086
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
1032
1087
 
1033
1088
  let tid = thread_id % THREADS_PER_BLOCK;
1034
1089
  let block_group = thread_id / THREADS_PER_BLOCK;
@@ -1043,11 +1098,12 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
1043
1098
 
1044
1099
  for (var block = block_group; block < num_blocks; block += num_block_groups) {
1045
1100
  let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
1046
- var x_block: array<f32, 16>;
1047
- for (var i = 0u; i < 16u; i++) {
1048
- x_block[i] = f32(src1[x_base + i]);
1101
+ var x_block: array<array<f32, 16>, NUM_COLS>;
1102
+ for (var col = 0u; col < NUM_COLS;col += 1) {
1103
+ for (var i = 0u; i < 16u; i++) {
1104
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
1105
+ }
1049
1106
  }
1050
-
1051
1107
  for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
1052
1108
  let output_row = row_base + row;
1053
1109
  if (output_row < params.m) {
@@ -1058,27 +1114,29 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
1058
1114
  let scales_word = load_u32_at_src0(block_byte_base + 66u + (sub_blk / 4u) * 4u);
1059
1115
  let scales_byte = get_byte(scales_word, sub_blk % 4u);
1060
1116
 
1061
- var row_sum = 0.0;
1062
- for (var ll = 0u; ll < 2u; ll++) {
1063
- let l = slot0 + ll;
1064
- let qs_word = select(qs_hi, qs_lo, l < 2u);
1065
- let half2 = (l % 2u) * 16u;
1066
- let qs_val = (qs_word >> half2) & 0xFFFFu;
1067
- let grid_idx = qs_val & 0x1FFu;
1068
- let signs_idx = (qs_val >> 9u) & 0x7Fu;
1069
- let sub_scale = (scales_byte >> (4u * (l / 2u))) & 0xFu;
1070
- let db = d * (0.5 + f32(sub_scale)) * 0.25;
1071
- let signs = (ksigns_iq2xs[signs_idx / 4u] >> ((signs_idx % 4u) * 8u)) & 0xFFu;
1072
- let gw_lo = iq2xs_grid[grid_idx * 2u];
1073
- let gw_hi = iq2xs_grid[grid_idx * 2u + 1u];
1074
- for (var j = 0u; j < 8u; j++) {
1075
- let gw = select(gw_hi, gw_lo, j < 4u);
1076
- let b = f32((gw >> ((j & 3u) * 8u)) & 0xFFu);
1077
- let s = select(1.0, -1.0, ((signs >> j) & 1u) != 0u);
1078
- row_sum += db * b * s * x_block[ll * 8u + j];
1117
+ for (var col = 0u;col < NUM_COLS;col += 1) {
1118
+ var row_sum = 0.0;
1119
+ for (var ll = 0u; ll < 2u; ll++) {
1120
+ let l = slot0 + ll;
1121
+ let qs_word = select(qs_hi, qs_lo, l < 2u);
1122
+ let half2 = (l % 2u) * 16u;
1123
+ let qs_val = (qs_word >> half2) & 0xFFFFu;
1124
+ let grid_idx = qs_val & 0x1FFu;
1125
+ let signs_idx = (qs_val >> 9u) & 0x7Fu;
1126
+ let sub_scale = (scales_byte >> (4u * (l / 2u))) & 0xFu;
1127
+ let db = d * (0.5 + f32(sub_scale)) * 0.25;
1128
+ let signs = (ksigns_iq2xs[signs_idx / 4u] >> ((signs_idx % 4u) * 8u)) & 0xFFu;
1129
+ let gw_lo = iq2xs_grid[grid_idx * 2u];
1130
+ let gw_hi = iq2xs_grid[grid_idx * 2u + 1u];
1131
+ for (var j = 0u; j < 8u; j++) {
1132
+ let gw = select(gw_hi, gw_lo, j < 4u);
1133
+ let b = f32((gw >> ((j & 3u) * 8u)) & 0xFFu);
1134
+ let s = select(1.0, -1.0, ((signs >> j) & 1u) != 0u);
1135
+ row_sum += db * b * s * x_block[col][ll * 8u + j];
1136
+ }
1079
1137
  }
1138
+ acc[col][row] += row_sum;
1080
1139
  }
1081
- acc[row] += row_sum;
1082
1140
  }
1083
1141
  }
1084
1142
  }
@@ -1091,8 +1149,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
1091
1149
  #define BLOCK_SIZE 256
1092
1150
  #define BLOCK_SIZE_BYTES 82
1093
1151
  #define THREADS_PER_BLOCK 16
1094
- fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
1095
- var acc: array<f32, OUTPUTS_PER_WG>;
1152
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
1153
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
1096
1154
 
1097
1155
  let tid = thread_id % THREADS_PER_BLOCK;
1098
1156
  let block_group = thread_id / THREADS_PER_BLOCK;
@@ -1107,11 +1165,12 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
1107
1165
 
1108
1166
  for (var block = block_group; block < num_blocks; block += num_block_groups) {
1109
1167
  let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
1110
- var x_block: array<f32, 16>;
1111
- for (var i = 0u; i < 16u; i++) {
1112
- x_block[i] = f32(src1[x_base + i]);
1168
+ var x_block: array<array<f32, 16>, NUM_COLS>;
1169
+ for (var col = 0u; col < NUM_COLS;col += 1) {
1170
+ for (var i = 0u; i < 16u; i++) {
1171
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
1172
+ }
1113
1173
  }
1114
-
1115
1174
  for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
1116
1175
  let output_row = row_base + row;
1117
1176
  if (output_row < params.m) {
@@ -1124,24 +1183,26 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
1124
1183
  let sc_word = load_u32_at_src0(block_byte_base + 74u + (sub_blk / 4u) * 4u);
1125
1184
  let scales_byte = get_byte(sc_word, sub_blk % 4u);
1126
1185
 
1127
- var row_sum = 0.0;
1128
- for (var ll = 0u; ll < 2u; ll++) {
1129
- let l = slot0 + ll;
1130
- let qs_byte = get_byte(qs_w, l);
1131
- let sign_byte = get_byte(sg_w, l);
1132
- let grid_idx = qs_byte | (((qh_byte >> (2u * l)) & 3u) << 8u);
1133
- let sub_scale = (scales_byte >> (4u * (l / 2u))) & 0xFu;
1134
- let db = d * (0.5 + f32(sub_scale)) * 0.25;
1135
- let gw_lo = iq2s_grid[grid_idx * 2u];
1136
- let gw_hi = iq2s_grid[grid_idx * 2u + 1u];
1137
- for (var j = 0u; j < 8u; j++) {
1138
- let gw = select(gw_hi, gw_lo, j < 4u);
1139
- let b = f32((gw >> ((j & 3u) * 8u)) & 0xFFu);
1140
- let s = select(1.0, -1.0, ((sign_byte >> j) & 1u) != 0u);
1141
- row_sum += db * b * s * x_block[ll * 8u + j];
1186
+ for (var col = 0u;col < NUM_COLS;col += 1) {
1187
+ var row_sum = 0.0;
1188
+ for (var ll = 0u; ll < 2u; ll++) {
1189
+ let l = slot0 + ll;
1190
+ let qs_byte = get_byte(qs_w, l);
1191
+ let sign_byte = get_byte(sg_w, l);
1192
+ let grid_idx = qs_byte | (((qh_byte >> (2u * l)) & 3u) << 8u);
1193
+ let sub_scale = (scales_byte >> (4u * (l / 2u))) & 0xFu;
1194
+ let db = d * (0.5 + f32(sub_scale)) * 0.25;
1195
+ let gw_lo = iq2s_grid[grid_idx * 2u];
1196
+ let gw_hi = iq2s_grid[grid_idx * 2u + 1u];
1197
+ for (var j = 0u; j < 8u; j++) {
1198
+ let gw = select(gw_hi, gw_lo, j < 4u);
1199
+ let b = f32((gw >> ((j & 3u) * 8u)) & 0xFFu);
1200
+ let s = select(1.0, -1.0, ((sign_byte >> j) & 1u) != 0u);
1201
+ row_sum += db * b * s * x_block[col][ll * 8u + j];
1202
+ }
1142
1203
  }
1204
+ acc[col][row] += row_sum;
1143
1205
  }
1144
- acc[row] += row_sum;
1145
1206
  }
1146
1207
  }
1147
1208
  }
@@ -1154,8 +1215,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
1154
1215
  #define BLOCK_SIZE 256
1155
1216
  #define BLOCK_SIZE_BYTES 98
1156
1217
  #define THREADS_PER_BLOCK 16
1157
- fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
1158
- var acc: array<f32, OUTPUTS_PER_WG>;
1218
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
1219
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
1159
1220
 
1160
1221
  let tid = thread_id % THREADS_PER_BLOCK;
1161
1222
  let block_group = thread_id / THREADS_PER_BLOCK;
@@ -1170,11 +1231,12 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
1170
1231
 
1171
1232
  for (var block = block_group; block < num_blocks; block += num_block_groups) {
1172
1233
  let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
1173
- var x_block: array<f32, 16>;
1174
- for (var i = 0u; i < 16u; i++) {
1175
- x_block[i] = f32(src1[x_base + i]);
1234
+ var x_block: array<array<f32, 16>, NUM_COLS>;
1235
+ for (var col = 0u; col < NUM_COLS;col += 1) {
1236
+ for (var i = 0u; i < 16u; i++) {
1237
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
1238
+ }
1176
1239
  }
1177
-
1178
1240
  for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
1179
1241
  let output_row = row_base + row;
1180
1242
  if (output_row < params.m) {
@@ -1186,27 +1248,29 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
1186
1248
  let ls = aux >> 28u;
1187
1249
  let db = d * (0.5 + f32(ls)) * 0.5;
1188
1250
 
1189
- var row_sum = 0.0;
1190
- for (var ll = 0u; ll < 2u; ll++) {
1191
- let l = slot0 + ll;
1192
- let qs_word = select(qs_hi, qs_lo, l < 2u);
1193
- let byte_pos = (l % 2u) * 2u;
1194
- let grid_idx_0 = (qs_word >> (byte_pos * 8u)) & 0xFFu;
1195
- let grid_idx_1 = (qs_word >> ((byte_pos + 1u) * 8u)) & 0xFFu;
1196
- let signs_idx = (aux >> (7u * l)) & 0x7Fu;
1197
- let signs = (ksigns_iq2xs[signs_idx / 4u] >> ((signs_idx % 4u) * 8u)) & 0xFFu;
1198
- let grid1 = iq3xxs_grid[grid_idx_0];
1199
- let grid2 = iq3xxs_grid[grid_idx_1];
1200
- for (var j = 0u; j < 4u; j++) {
1201
- let b1 = f32((grid1 >> (j * 8u)) & 0xFFu);
1202
- let b2 = f32((grid2 >> (j * 8u)) & 0xFFu);
1203
- let s1 = select(1.0, -1.0, ((signs >> j) & 1u) != 0u);
1204
- let s2 = select(1.0, -1.0, ((signs >> (j + 4u)) & 1u) != 0u);
1205
- row_sum += db * b1 * s1 * x_block[ll * 8u + j];
1206
- row_sum += db * b2 * s2 * x_block[ll * 8u + j + 4u];
1251
+ for (var col = 0u;col < NUM_COLS;col += 1) {
1252
+ var row_sum = 0.0;
1253
+ for (var ll = 0u; ll < 2u; ll++) {
1254
+ let l = slot0 + ll;
1255
+ let qs_word = select(qs_hi, qs_lo, l < 2u);
1256
+ let byte_pos = (l % 2u) * 2u;
1257
+ let grid_idx_0 = (qs_word >> (byte_pos * 8u)) & 0xFFu;
1258
+ let grid_idx_1 = (qs_word >> ((byte_pos + 1u) * 8u)) & 0xFFu;
1259
+ let signs_idx = (aux >> (7u * l)) & 0x7Fu;
1260
+ let signs = (ksigns_iq2xs[signs_idx / 4u] >> ((signs_idx % 4u) * 8u)) & 0xFFu;
1261
+ let grid1 = iq3xxs_grid[grid_idx_0];
1262
+ let grid2 = iq3xxs_grid[grid_idx_1];
1263
+ for (var j = 0u; j < 4u; j++) {
1264
+ let b1 = f32((grid1 >> (j * 8u)) & 0xFFu);
1265
+ let b2 = f32((grid2 >> (j * 8u)) & 0xFFu);
1266
+ let s1 = select(1.0, -1.0, ((signs >> j) & 1u) != 0u);
1267
+ let s2 = select(1.0, -1.0, ((signs >> (j + 4u)) & 1u) != 0u);
1268
+ row_sum += db * b1 * s1 * x_block[col][ll * 8u + j];
1269
+ row_sum += db * b2 * s2 * x_block[col][ll * 8u + j + 4u];
1270
+ }
1207
1271
  }
1272
+ acc[col][row] += row_sum;
1208
1273
  }
1209
- acc[row] += row_sum;
1210
1274
  }
1211
1275
  }
1212
1276
  }
@@ -1219,8 +1283,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
1219
1283
  #define BLOCK_SIZE 256
1220
1284
  #define BLOCK_SIZE_BYTES 110
1221
1285
  #define THREADS_PER_BLOCK 16
1222
- fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
1223
- var acc: array<f32, OUTPUTS_PER_WG>;
1286
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
1287
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
1224
1288
 
1225
1289
  let tid = thread_id % THREADS_PER_BLOCK;
1226
1290
  let block_group = thread_id / THREADS_PER_BLOCK;
@@ -1235,11 +1299,12 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
1235
1299
 
1236
1300
  for (var block = block_group; block < num_blocks; block += num_block_groups) {
1237
1301
  let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
1238
- var x_block: array<f32, 16>;
1239
- for (var i = 0u; i < 16u; i++) {
1240
- x_block[i] = f32(src1[x_base + i]);
1302
+ var x_block: array<array<f32, 16>, NUM_COLS>;
1303
+ for (var col = 0u; col < NUM_COLS;col += 1) {
1304
+ for (var i = 0u; i < 16u; i++) {
1305
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
1306
+ }
1241
1307
  }
1242
-
1243
1308
  for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
1244
1309
  let output_row = row_base + row;
1245
1310
  if (output_row < params.m) {
@@ -1255,28 +1320,30 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
1255
1320
  let sub_scale = (scales_byte >> (4u * (sub_blk % 2u))) & 0xFu;
1256
1321
  let db = d * (1.0 + 2.0 * f32(sub_scale));
1257
1322
 
1258
- var row_sum = 0.0;
1259
- for (var ll = 0u; ll < 2u; ll++) {
1260
- let l = slot0 + ll;
1261
- let qs_word = select(qs_hi, qs_lo, l < 2u);
1262
- let byte_pos = (l % 2u) * 2u;
1263
- let qs0 = (qs_word >> (byte_pos * 8u)) & 0xFFu;
1264
- let qs1 = (qs_word >> ((byte_pos + 1u) * 8u)) & 0xFFu;
1265
- let grid_idx_1 = qs0 | (((qh_byte >> (2u * l)) & 1u) << 8u);
1266
- let grid_idx_2 = qs1 | (((qh_byte >> (2u * l + 1u)) & 1u) << 8u);
1267
- let sign_byte = get_byte(sg_w, l);
1268
- let grid1 = iq3s_grid[grid_idx_1];
1269
- let grid2 = iq3s_grid[grid_idx_2];
1270
- for (var j = 0u; j < 4u; j++) {
1271
- let b1 = f32((grid1 >> (j * 8u)) & 0xFFu);
1272
- let b2 = f32((grid2 >> (j * 8u)) & 0xFFu);
1273
- let s1 = select(1.0, -1.0, ((sign_byte >> j) & 1u) != 0u);
1274
- let s2 = select(1.0, -1.0, ((sign_byte >> (j + 4u)) & 1u) != 0u);
1275
- row_sum += db * b1 * s1 * x_block[ll * 8u + j];
1276
- row_sum += db * b2 * s2 * x_block[ll * 8u + j + 4u];
1323
+ for (var col = 0u;col < NUM_COLS;col += 1) {
1324
+ var row_sum = 0.0;
1325
+ for (var ll = 0u; ll < 2u; ll++) {
1326
+ let l = slot0 + ll;
1327
+ let qs_word = select(qs_hi, qs_lo, l < 2u);
1328
+ let byte_pos = (l % 2u) * 2u;
1329
+ let qs0 = (qs_word >> (byte_pos * 8u)) & 0xFFu;
1330
+ let qs1 = (qs_word >> ((byte_pos + 1u) * 8u)) & 0xFFu;
1331
+ let grid_idx_1 = qs0 | (((qh_byte >> (2u * l)) & 1u) << 8u);
1332
+ let grid_idx_2 = qs1 | (((qh_byte >> (2u * l + 1u)) & 1u) << 8u);
1333
+ let sign_byte = get_byte(sg_w, l);
1334
+ let grid1 = iq3s_grid[grid_idx_1];
1335
+ let grid2 = iq3s_grid[grid_idx_2];
1336
+ for (var j = 0u; j < 4u; j++) {
1337
+ let b1 = f32((grid1 >> (j * 8u)) & 0xFFu);
1338
+ let b2 = f32((grid2 >> (j * 8u)) & 0xFFu);
1339
+ let s1 = select(1.0, -1.0, ((sign_byte >> j) & 1u) != 0u);
1340
+ let s2 = select(1.0, -1.0, ((sign_byte >> (j + 4u)) & 1u) != 0u);
1341
+ row_sum += db * b1 * s1 * x_block[col][ll * 8u + j];
1342
+ row_sum += db * b2 * s2 * x_block[col][ll * 8u + j + 4u];
1343
+ }
1277
1344
  }
1345
+ acc[col][row] += row_sum;
1278
1346
  }
1279
- acc[row] += row_sum;
1280
1347
  }
1281
1348
  }
1282
1349
  }
@@ -1290,35 +1357,37 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
1290
1357
  #define BLOCK_SIZE_BYTES 18
1291
1358
  #define THREADS_PER_BLOCK 4
1292
1359
  #define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
1293
- fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
1294
- var acc: array<f32, OUTPUTS_PER_WG>;
1360
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
1361
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
1295
1362
 
1296
1363
  let num_blocks = params.k / BLOCK_SIZE;
1297
1364
  let thread_within_block = thread_id % THREADS_PER_BLOCK;
1298
1365
  for (var block = thread_id / THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE / THREADS_PER_BLOCK) {
1299
1366
  let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * 4u;
1300
- var x_block: array<f32, ELEMS_PER_THREAD>;
1301
- for (var i = 0u; i < ELEMS_PER_THREAD / 2u; i++) {
1302
- x_block[i] = f32(src1[x_base + i]);
1303
- x_block[i + 4u] = f32(src1[x_base + i + 16u]);
1367
+ var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
1368
+ for (var col = 0u; col < NUM_COLS;col += 1) {
1369
+ for (var i = 0u; i < ELEMS_PER_THREAD / 2u; i++) {
1370
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
1371
+ x_block[col][i + 4u] = f32(src1[x_base + col * params.stride_11 + i + 16u]);
1372
+ }
1304
1373
  }
1305
-
1306
1374
  for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
1307
1375
  let output_row = row_base + row;
1308
1376
  if (output_row < params.m) {
1309
1377
  let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
1310
1378
  let d = f32(load_f16_at_src0(block_byte_base));
1311
- var row_sum = 0.0;
1312
-
1313
1379
  let q_packed = load_u32_at_src0(block_byte_base + 2u + 4u * thread_within_block);
1314
- for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
1315
- let q_byte = get_byte(q_packed, byte_idx);
1316
- let q_lo = f32(kvalues_iq4nl[q_byte & 0xFu]) * d;
1317
- let q_hi = f32(kvalues_iq4nl[(q_byte >> 4u) & 0xFu]) * d;
1318
- row_sum += q_lo * x_block[byte_idx];
1319
- row_sum += q_hi * x_block[byte_idx + 4u];
1380
+ for (var col = 0u;col < NUM_COLS;col += 1) {
1381
+ var row_sum = 0.0;
1382
+ for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
1383
+ let q_byte = get_byte(q_packed, byte_idx);
1384
+ let q_lo = f32(kvalues_iq4nl[q_byte & 0xFu]) * d;
1385
+ let q_hi = f32(kvalues_iq4nl[(q_byte >> 4u) & 0xFu]) * d;
1386
+ row_sum += q_lo * x_block[col][byte_idx];
1387
+ row_sum += q_hi * x_block[col][byte_idx + 4u];
1388
+ }
1389
+ acc[col][row] += row_sum;
1320
1390
  }
1321
- acc[row] += row_sum;
1322
1391
  }
1323
1392
  }
1324
1393
  }
@@ -1331,8 +1400,8 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
1331
1400
  #define BLOCK_SIZE 256
1332
1401
  #define BLOCK_SIZE_BYTES 136
1333
1402
  #define THREADS_PER_BLOCK 16
1334
- fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
1335
- var acc: array<f32, OUTPUTS_PER_WG>;
1403
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
1404
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
1336
1405
 
1337
1406
  let tid = thread_id % THREADS_PER_BLOCK;
1338
1407
  let block_group = thread_id / THREADS_PER_BLOCK;
@@ -1346,11 +1415,12 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
1346
1415
 
1347
1416
  for (var block = block_group; block < num_blocks; block += num_block_groups) {
1348
1417
  let x_base = src1_idx_base + block * BLOCK_SIZE + y_offset;
1349
- var x_block: array<f32, 16>;
1350
- for (var i = 0u; i < 16u; i++) {
1351
- x_block[i] = f32(src1[x_base + i]);
1418
+ var x_block: array<array<f32, 16>, NUM_COLS>;
1419
+ for (var col = 0u; col < NUM_COLS;col += 1) {
1420
+ for (var i = 0u; i < 16u; i++) {
1421
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
1422
+ }
1352
1423
  }
1353
-
1354
1424
  for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
1355
1425
  let output_row = row_base + row;
1356
1426
  if (output_row < params.m) {
@@ -1370,17 +1440,19 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
1370
1440
  let q_w2 = load_u32_at_src0(block_byte_base + qs_byte_off + 8u);
1371
1441
  let q_w3 = load_u32_at_src0(block_byte_base + qs_byte_off + 12u);
1372
1442
 
1373
- var row_sum = 0.0;
1374
- for (var i = 0u; i < 16u; i++) {
1375
- let q_word = select(
1376
- select(q_w0, q_w1, i >= 4u),
1377
- select(q_w2, q_w3, i >= 12u),
1378
- i >= 8u);
1379
- let q_byte = get_byte(q_word, i % 4u);
1380
- let nib = select(q_byte & 0xFu, (q_byte >> 4u) & 0xFu, half == 1u);
1381
- row_sum += f32(kvalues_iq4nl[nib]) * dl * x_block[i];
1443
+ for (var col = 0u;col < NUM_COLS;col += 1) {
1444
+ var row_sum = 0.0;
1445
+ for (var i = 0u; i < 16u; i++) {
1446
+ let q_word = select(
1447
+ select(q_w0, q_w1, i >= 4u),
1448
+ select(q_w2, q_w3, i >= 12u),
1449
+ i >= 8u);
1450
+ let q_byte = get_byte(q_word, i % 4u);
1451
+ let nib = select(q_byte & 0xFu, (q_byte >> 4u) & 0xFu, half == 1u);
1452
+ row_sum += f32(kvalues_iq4nl[nib]) * dl * x_block[col][i];
1453
+ }
1454
+ acc[col][row] += row_sum;
1382
1455
  }
1383
- acc[row] += row_sum;
1384
1456
  }
1385
1457
  }
1386
1458
  }
@@ -1394,35 +1466,84 @@ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src
1394
1466
  #define BLOCK_SIZE_BYTES 17
1395
1467
  #define THREADS_PER_BLOCK 4
1396
1468
  #define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
1397
- fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<f32, OUTPUTS_PER_WG> {
1398
- var acc: array<f32, OUTPUTS_PER_WG>;
1469
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
1470
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
1399
1471
 
1400
1472
  let num_blocks = params.k / BLOCK_SIZE;
1401
1473
  let thread_within_block = thread_id % 4;
1402
1474
  for (var block = thread_id/THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE/THREADS_PER_BLOCK) {
1403
1475
  let x_base = src1_idx_base + block * BLOCK_SIZE + thread_within_block * 4;
1404
- var x_block: array<f32, ELEMS_PER_THREAD>;
1405
- for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
1406
- x_block[i] = f32(src1[x_base + i]);
1407
- x_block[i + 4] = f32(src1[x_base + i + 16]);
1476
+ var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
1477
+ for (var col = 0u; col < NUM_COLS;col += 1) {
1478
+ for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
1479
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
1480
+ x_block[col][i + 4] = f32(src1[x_base + col * params.stride_11 + i + 16]);
1481
+ }
1408
1482
  }
1409
-
1410
1483
  for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
1411
1484
  let output_row = row_base + row;
1412
1485
  if (output_row < params.m) {
1413
1486
  let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
1414
1487
  let eu8 = get_byte(load_u32_at_src0(block_byte_base), 0);
1415
1488
  let e = ldexp(1.0, i32(eu8) - 128);
1416
- var row_sum = 0.0;
1417
1489
  let q_packed = load_u32_at_src0(block_byte_base + 1u + 4u * thread_within_block);
1418
- for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
1419
- let q_byte = get_byte(q_packed, byte_idx);
1420
- let q_lo = f32(kvalues_mxfp4[q_byte & 0xFu]) * e;
1421
- let q_hi = f32(kvalues_mxfp4[(q_byte >> 4u) & 0xFu]) * e;
1422
- row_sum += q_lo * x_block[byte_idx];
1423
- row_sum += q_hi * x_block[byte_idx + 4u];
1490
+ for (var col = 0u;col < NUM_COLS;col += 1) {
1491
+ var row_sum = 0.0;
1492
+ for (var byte_idx = 0u; byte_idx < 4u; byte_idx++) {
1493
+ let q_byte = get_byte(q_packed, byte_idx);
1494
+ let q_lo = f32(kvalues_mxfp4[q_byte & 0xFu]) * e;
1495
+ let q_hi = f32(kvalues_mxfp4[(q_byte >> 4u) & 0xFu]) * e;
1496
+ row_sum += q_lo * x_block[col][byte_idx];
1497
+ row_sum += q_hi * x_block[col][byte_idx + 4u];
1498
+ }
1499
+ acc[col][row] += row_sum;
1500
+ }
1501
+ }
1502
+ }
1503
+ }
1504
+
1505
+ return acc;
1506
+ }
1507
+ #endif
1508
+
1509
+ #ifdef MUL_ACC_NVFP4
1510
+ #define BLOCK_SIZE 64
1511
+ #define BLOCK_SIZE_BYTES 36
1512
+ #define THREADS_PER_BLOCK 4
1513
+ #define ELEMS_PER_THREAD (BLOCK_SIZE/THREADS_PER_BLOCK)
1514
+ fn accumulate_vec_dot(thread_id: u32, row_base: u32, src0_batch_offset: u32, src1_idx_base: u32) -> array<array<f32, OUTPUTS_PER_WG>, NUM_COLS> {
1515
+ var acc: array<array<f32, OUTPUTS_PER_WG>, NUM_COLS>;
1516
+
1517
+ let num_blocks = params.k / BLOCK_SIZE;
1518
+ let sub = thread_id % THREADS_PER_BLOCK;
1519
+ for (var block = thread_id/THREADS_PER_BLOCK; block < num_blocks; block += WG_SIZE/THREADS_PER_BLOCK) {
1520
+ let x_base = src1_idx_base + block * BLOCK_SIZE + sub * ELEMS_PER_THREAD;
1521
+ var x_block: array<array<f32, ELEMS_PER_THREAD>, NUM_COLS>;
1522
+ for (var col = 0u; col < NUM_COLS;col += 1) {
1523
+ for (var i = 0u; i < ELEMS_PER_THREAD / 2; i++) {
1524
+ x_block[col][i] = f32(src1[x_base + col * params.stride_11 + i]);
1525
+ x_block[col][i + 8] = f32(src1[x_base + col * params.stride_11 + i + 8]);
1526
+ }
1527
+ }
1528
+ for (var row = 0u; row < OUTPUTS_PER_WG; row++) {
1529
+ let output_row = row_base + row;
1530
+ if (output_row < params.m) {
1531
+ let block_byte_base = (src0_batch_offset + output_row * params.stride_01 + block) * BLOCK_SIZE_BYTES;
1532
+ let d = ue4m3_to_fp32(get_byte(load_u32_at_src0_aligned(block_byte_base), sub)) * 0.5;
1533
+ let q_w0 = load_u32_at_src0_aligned(block_byte_base + 4u + 8u * sub);
1534
+ let q_w1 = load_u32_at_src0_aligned(block_byte_base + 8u + 8u * sub);
1535
+ for (var col = 0u;col < NUM_COLS;col += 1) {
1536
+ var row_sum = 0.0;
1537
+ for (var l = 0u; l < 8u; l++) {
1538
+ let q_word = select(q_w0, q_w1, l >= 4u);
1539
+ let q_byte = get_byte(q_word, l % 4u);
1540
+ let q_lo = f32(kvalues_mxfp4[q_byte & 0xFu]) * d;
1541
+ let q_hi = f32(kvalues_mxfp4[(q_byte >> 4u) & 0xFu]) * d;
1542
+ row_sum += q_lo * x_block[col][l];
1543
+ row_sum += q_hi * x_block[col][l + 8u];
1544
+ }
1545
+ acc[col][row] += row_sum;
1424
1546
  }
1425
- acc[row] += row_sum;
1426
1547
  }
1427
1548
  }
1428
1549
  }