whispercpp 1.3.7 → 1.3.8

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (308) hide show
  1. checksums.yaml +4 -4
  2. data/README.md +5 -4
  3. data/ext/options.rb +1 -1
  4. data/ext/ruby_whisper.c +0 -1
  5. data/ext/ruby_whisper.h +7 -1
  6. data/ext/ruby_whisper_context.c +50 -1
  7. data/ext/ruby_whisper_log_settable.h +1 -2
  8. data/ext/ruby_whisper_params.c +9 -8
  9. data/ext/ruby_whisper_transcribe.cpp +0 -19
  10. data/ext/ruby_whisper_vad_context.c +30 -10
  11. data/ext/ruby_whisper_vad_context_detect.cpp +8 -9
  12. data/ext/ruby_whisper_vad_params.c +4 -4
  13. data/ext/ruby_whisper_vad_segment.c +2 -2
  14. data/ext/sources/CMakeLists.txt +2 -1
  15. data/ext/sources/cmake/parakeet.pc.in +2 -2
  16. data/ext/sources/cmake/whisper.pc.in +2 -2
  17. data/ext/sources/examples/cli/cli.cpp +9 -1
  18. data/ext/sources/examples/common-ggml.cpp +2 -0
  19. data/ext/sources/examples/vad-speech-segments/speech.cpp +3 -2
  20. data/ext/sources/ggml/CMakeLists.txt +3 -4
  21. data/ext/sources/ggml/include/ggml-cuda.h +0 -3
  22. data/ext/sources/ggml/include/ggml-sycl.h +8 -0
  23. data/ext/sources/ggml/include/ggml.h +3 -1
  24. data/ext/sources/ggml/src/CMakeLists.txt +8 -1
  25. data/ext/sources/ggml/src/ggml-backend-meta.cpp +7 -4
  26. data/ext/sources/ggml/src/ggml-common.h +13 -2
  27. data/ext/sources/ggml/src/ggml-cpu/CMakeLists.txt +1 -1
  28. data/ext/sources/ggml/src/ggml-cpu/amx/mmq.cpp +5 -6
  29. data/ext/sources/ggml/src/ggml-cpu/arch/arm/quants.c +78 -4
  30. data/ext/sources/ggml/src/ggml-cpu/arch/x86/quants.c +142 -4
  31. data/ext/sources/ggml/src/ggml-cpu/arch-fallback.h +7 -2
  32. data/ext/sources/ggml/src/ggml-cpu/ggml-cpu.c +14 -0
  33. data/ext/sources/ggml/src/ggml-cpu/llamafile/sgemm.cpp +26 -19
  34. data/ext/sources/ggml/src/ggml-cpu/ops.cpp +129 -46
  35. data/ext/sources/ggml/src/ggml-cpu/quants.c +51 -0
  36. data/ext/sources/ggml/src/ggml-cpu/quants.h +3 -0
  37. data/ext/sources/ggml/src/ggml-cpu/simd-gemm.h +1 -1
  38. data/ext/sources/ggml/src/ggml-cpu/simd-mappings.h +11 -0
  39. data/ext/sources/ggml/src/ggml-cpu/vec.cpp +2 -2
  40. data/ext/sources/ggml/src/ggml-cuda/binbcast.cu +90 -46
  41. data/ext/sources/ggml/src/ggml-cuda/col2im-1d.cu +81 -0
  42. data/ext/sources/ggml/src/ggml-cuda/col2im-1d.cuh +3 -0
  43. data/ext/sources/ggml/src/ggml-cuda/common.cuh +4 -0
  44. data/ext/sources/ggml/src/ggml-cuda/concat.cu +33 -21
  45. data/ext/sources/ggml/src/ggml-cuda/conv-transpose-1d.cu +14 -12
  46. data/ext/sources/ggml/src/ggml-cuda/convert.cu +86 -34
  47. data/ext/sources/ggml/src/ggml-cuda/cpy.cu +80 -29
  48. data/ext/sources/ggml/src/ggml-cuda/fattn-common.cuh +9 -5
  49. data/ext/sources/ggml/src/ggml-cuda/fattn-mma-f16.cuh +4 -0
  50. data/ext/sources/ggml/src/ggml-cuda/fattn-tile.cuh +9 -5
  51. data/ext/sources/ggml/src/ggml-cuda/fattn.cu +27 -21
  52. data/ext/sources/ggml/src/ggml-cuda/gated_delta_net.cu +40 -25
  53. data/ext/sources/ggml/src/ggml-cuda/gated_delta_net.cuh +10 -0
  54. data/ext/sources/ggml/src/ggml-cuda/getrows.cu +15 -12
  55. data/ext/sources/ggml/src/ggml-cuda/ggml-cuda.cu +718 -1248
  56. data/ext/sources/ggml/src/ggml-cuda/mmq.cu +7 -0
  57. data/ext/sources/ggml/src/ggml-cuda/mmvq.cu +77 -40
  58. data/ext/sources/ggml/src/ggml-cuda/out-prod.cu +55 -12
  59. data/ext/sources/ggml/src/ggml-cuda/set-rows.cu +64 -4
  60. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_16-ncols2_2.cu +1 -0
  61. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_32-ncols2_2.cu +1 -0
  62. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_2.cu +1 -0
  63. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_8-ncols2_2.cu +1 -0
  64. data/ext/sources/ggml/src/ggml-cuda/topk-moe.cu +7 -1
  65. data/ext/sources/ggml/src/ggml-cuda/vendors/hip.h +1 -0
  66. data/ext/sources/ggml/src/ggml-cuda/vendors/musa.h +1 -0
  67. data/ext/sources/ggml/src/ggml-hexagon/CMakeLists.txt +0 -5
  68. data/ext/sources/ggml/src/ggml-hexagon/ggml-hexagon.cpp +1634 -1293
  69. data/ext/sources/ggml/src/ggml-hexagon/htp/CMakeLists.txt +11 -40
  70. data/ext/sources/ggml/src/ggml-hexagon/htp/cmake-toolchain.cmake +13 -15
  71. data/ext/sources/ggml/src/ggml-hexagon/htp/concat-ops.c +1 -1
  72. data/ext/sources/ggml/src/ggml-hexagon/htp/flash-attn-ops.c +1749 -399
  73. data/ext/sources/ggml/src/ggml-hexagon/htp/flash-attn-ops.h +303 -0
  74. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-common.h +80 -0
  75. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-dma.h +26 -23
  76. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-profile.h +64 -0
  77. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-utils.h +1 -83
  78. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-fa-kernels.h +555 -0
  79. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-mm-kernels-tiled.h +1303 -0
  80. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-queue.c +9 -0
  81. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-queue.h +27 -4
  82. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-utils.h +59 -37
  83. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-ctx.h +11 -3
  84. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-ops.h +52 -12
  85. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-vtcm.h +19 -0
  86. data/ext/sources/ggml/src/ggml-hexagon/htp/htp_iface.idl +2 -1
  87. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-base.h +14 -30
  88. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-exp.h +39 -0
  89. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-fa-kernels.h +232 -0
  90. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-flat.h +1511 -0
  91. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-tiled.h +1200 -0
  92. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-sigmoid.h +39 -0
  93. data/ext/sources/ggml/src/ggml-hexagon/htp/main.c +127 -32
  94. data/ext/sources/ggml/src/ggml-hexagon/htp/matmul-ops.c +3023 -4425
  95. data/ext/sources/ggml/src/ggml-hexagon/htp/matmul-ops.h +650 -0
  96. data/ext/sources/ggml/src/ggml-hexagon/htp/rope-ops.c +48 -13
  97. data/ext/sources/ggml/src/ggml-hexagon/htp/ssm-conv.c +10 -9
  98. data/ext/sources/ggml/src/ggml-hexagon/htp/worker-pool.c +15 -3
  99. data/ext/sources/ggml/src/ggml-hexagon/htp/worker-pool.h +8 -0
  100. data/ext/sources/ggml/src/ggml-hexagon/htp-opnode.h +168 -50
  101. data/ext/sources/ggml/src/ggml-hexagon/libggml-htp.inf +0 -4
  102. data/ext/sources/ggml/src/ggml-hip/CMakeLists.txt +5 -0
  103. data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.cpp +69 -5
  104. data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.h +4 -1
  105. data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.m +27 -6
  106. data/ext/sources/ggml/src/ggml-metal/ggml-metal-impl.h +38 -0
  107. data/ext/sources/ggml/src/ggml-metal/ggml-metal-ops.cpp +132 -2
  108. data/ext/sources/ggml/src/ggml-metal/ggml-metal-ops.h +2 -0
  109. data/ext/sources/ggml/src/ggml-metal/ggml-metal.metal +345 -87
  110. data/ext/sources/ggml/src/ggml-opencl/CMakeLists.txt +13 -0
  111. data/ext/sources/ggml/src/ggml-opencl/fa_tune.h +92 -0
  112. data/ext/sources/ggml/src/ggml-opencl/ggml-opencl.cpp +4060 -357
  113. data/ext/sources/ggml/src/ggml-opencl/kernels/cvt.cl +198 -0
  114. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f16.cl +81 -41
  115. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32.cl +88 -39
  116. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_f16.cl +1995 -96
  117. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_q4_0.cl +1615 -0
  118. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_q8_0.cl +1486 -0
  119. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_pre_f16.cl +156 -0
  120. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_mxfp4_f32_ns.cl +74 -6
  121. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_0_f32_ns.cl +74 -6
  122. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_1_f32_ns.cl +74 -6
  123. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_k_f32_ns.cl +71 -6
  124. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_0_f32_ns.cl +74 -6
  125. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_1_f32_ns.cl +74 -6
  126. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_k_f32_ns.cl +74 -6
  127. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q6_k_f32_ns.cl +74 -6
  128. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q1_0_f32.cl +94 -0
  129. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q1_0_f32.cl +121 -0
  130. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q8_0_f32.cl +1 -1
  131. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q1_0_f32_l4_lm.cl +156 -0
  132. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_f16_f32_l4.cl +1149 -0
  133. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q1_0_f32.cl +141 -0
  134. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q1_0_f32_flat.cl +190 -0
  135. data/ext/sources/ggml/src/ggml-opencl/kernels/norm.cl +5 -2
  136. data/ext/sources/ggml/src/ggml-opencl/kernels/set_rows.cl +500 -0
  137. data/ext/sources/ggml/src/ggml-opencl/libdl.h +79 -0
  138. data/ext/sources/ggml/src/ggml-openvino/.clang-format +0 -5
  139. data/ext/sources/ggml/src/ggml-openvino/CMakeLists.txt +2 -4
  140. data/ext/sources/ggml/src/ggml-openvino/ggml-decoder.cpp +733 -130
  141. data/ext/sources/ggml/src/ggml-openvino/ggml-decoder.h +76 -23
  142. data/ext/sources/ggml/src/ggml-openvino/ggml-openvino-extra.cpp +57 -3
  143. data/ext/sources/ggml/src/ggml-openvino/ggml-openvino-extra.h +29 -8
  144. data/ext/sources/ggml/src/ggml-openvino/ggml-openvino.cpp +307 -59
  145. data/ext/sources/ggml/src/ggml-openvino/ggml-quants.cpp +66 -0
  146. data/ext/sources/ggml/src/ggml-openvino/ggml-quants.h +10 -4
  147. data/ext/sources/ggml/src/ggml-openvino/openvino/decoder.h +56 -16
  148. data/ext/sources/ggml/src/ggml-openvino/openvino/frontend.h +1 -1
  149. data/ext/sources/ggml/src/ggml-openvino/openvino/input_model.h +4 -4
  150. data/ext/sources/ggml/src/ggml-openvino/openvino/node_context.h +94 -37
  151. data/ext/sources/ggml/src/ggml-openvino/openvino/op/add_id.cpp +76 -0
  152. data/ext/sources/ggml/src/ggml-openvino/openvino/op/argsort.cpp +47 -0
  153. data/ext/sources/ggml/src/ggml-openvino/openvino/op/clamp.cpp +33 -0
  154. data/ext/sources/ggml/src/ggml-openvino/openvino/op/concat.cpp +48 -0
  155. data/ext/sources/ggml/src/ggml-openvino/openvino/op/cont.cpp +8 -16
  156. data/ext/sources/ggml/src/ggml-openvino/openvino/op/cpy.cpp +14 -1
  157. data/ext/sources/ggml/src/ggml-openvino/openvino/op/div.cpp +146 -0
  158. data/ext/sources/ggml/src/ggml-openvino/openvino/op/flash_attn_ext.cpp +108 -21
  159. data/ext/sources/ggml/src/ggml-openvino/openvino/op/gated_delta_net.cpp +282 -0
  160. data/ext/sources/ggml/src/ggml-openvino/openvino/op/gated_delta_net.hpp +65 -0
  161. data/ext/sources/ggml/src/ggml-openvino/openvino/op/get_rows.cpp +2 -9
  162. data/ext/sources/ggml/src/ggml-openvino/openvino/op/glu_geglu.cpp +21 -7
  163. data/ext/sources/ggml/src/ggml-openvino/openvino/op/glu_swiglu.cpp +41 -8
  164. data/ext/sources/ggml/src/ggml-openvino/openvino/op/im2col.cpp +120 -0
  165. data/ext/sources/ggml/src/ggml-openvino/openvino/op/l2_norm.cpp +44 -0
  166. data/ext/sources/ggml/src/ggml-openvino/openvino/op/mul_mat_id.cpp +226 -0
  167. data/ext/sources/ggml/src/ggml-openvino/openvino/op/mulmat.cpp +19 -9
  168. data/ext/sources/ggml/src/ggml-openvino/openvino/op/norm.cpp +58 -0
  169. data/ext/sources/ggml/src/ggml-openvino/openvino/op/pad.cpp +95 -0
  170. data/ext/sources/ggml/src/ggml-openvino/openvino/op/permute.cpp +58 -13
  171. data/ext/sources/ggml/src/ggml-openvino/openvino/op/repeat.cpp +74 -0
  172. data/ext/sources/ggml/src/ggml-openvino/openvino/op/reshape.cpp +13 -6
  173. data/ext/sources/ggml/src/ggml-openvino/openvino/op/rms_norm.cpp +1 -1
  174. data/ext/sources/ggml/src/ggml-openvino/openvino/op/rope.cpp +134 -38
  175. data/ext/sources/ggml/src/ggml-openvino/openvino/op/set_rows.cpp +3 -3
  176. data/ext/sources/ggml/src/ggml-openvino/openvino/op/softmax.cpp +126 -49
  177. data/ext/sources/ggml/src/ggml-openvino/openvino/op/ssm_conv.cpp +59 -0
  178. data/ext/sources/ggml/src/ggml-openvino/openvino/op/sum_rows.cpp +27 -0
  179. data/ext/sources/ggml/src/ggml-openvino/openvino/op/transpose.cpp +32 -1
  180. data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_silu.cpp +1 -1
  181. data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_softplus.cpp +38 -0
  182. data/ext/sources/ggml/src/ggml-openvino/openvino/op/view.cpp +90 -25
  183. data/ext/sources/ggml/src/ggml-openvino/openvino/op_table.cpp +41 -23
  184. data/ext/sources/ggml/src/ggml-openvino/openvino/op_table.h +18 -5
  185. data/ext/sources/ggml/src/ggml-openvino/openvino/pass/mark_decompression_convert_constant_folding.h +1 -1
  186. data/ext/sources/ggml/src/ggml-openvino/openvino/translate_session.cpp +43 -40
  187. data/ext/sources/ggml/src/ggml-openvino/openvino/translate_session.h +5 -4
  188. data/ext/sources/ggml/src/ggml-openvino/openvino/utils.cpp +548 -3
  189. data/ext/sources/ggml/src/ggml-openvino/openvino/utils.h +28 -26
  190. data/ext/sources/ggml/src/ggml-openvino/utils.cpp +383 -94
  191. data/ext/sources/ggml/src/ggml-openvino/utils.h +11 -8
  192. data/ext/sources/ggml/src/ggml-quants.c +76 -0
  193. data/ext/sources/ggml/src/ggml-quants.h +3 -0
  194. data/ext/sources/ggml/src/ggml-sycl/CMakeLists.txt +5 -5
  195. data/ext/sources/ggml/src/ggml-sycl/backend.hpp +2 -0
  196. data/ext/sources/ggml/src/ggml-sycl/binbcast.cpp +12 -0
  197. data/ext/sources/ggml/src/ggml-sycl/col2im-1d.cpp +102 -0
  198. data/ext/sources/ggml/src/ggml-sycl/col2im-1d.hpp +8 -0
  199. data/ext/sources/ggml/src/ggml-sycl/common.cpp +6 -8
  200. data/ext/sources/ggml/src/ggml-sycl/common.hpp +19 -2
  201. data/ext/sources/ggml/src/ggml-sycl/concat.cpp +21 -1
  202. data/ext/sources/ggml/src/ggml-sycl/conv2d-dw.cpp +158 -0
  203. data/ext/sources/ggml/src/ggml-sycl/conv2d-dw.hpp +10 -0
  204. data/ext/sources/ggml/src/ggml-sycl/conv2d-transpose.cpp +125 -0
  205. data/ext/sources/ggml/src/ggml-sycl/conv2d-transpose.hpp +10 -0
  206. data/ext/sources/ggml/src/ggml-sycl/conv2d.cpp +150 -0
  207. data/ext/sources/ggml/src/ggml-sycl/conv2d.hpp +10 -0
  208. data/ext/sources/ggml/src/ggml-sycl/conv3d.cpp +224 -0
  209. data/ext/sources/ggml/src/ggml-sycl/conv3d.hpp +8 -0
  210. data/ext/sources/ggml/src/ggml-sycl/convert.cpp +6 -0
  211. data/ext/sources/ggml/src/ggml-sycl/cpy.cpp +706 -0
  212. data/ext/sources/ggml/src/ggml-sycl/cpy.hpp +281 -0
  213. data/ext/sources/ggml/src/ggml-sycl/cross_entropy_loss.cpp +255 -0
  214. data/ext/sources/ggml/src/ggml-sycl/cross_entropy_loss.hpp +7 -0
  215. data/ext/sources/ggml/src/ggml-sycl/dequantize.hpp +15 -0
  216. data/ext/sources/ggml/src/ggml-sycl/dmmv.cpp +492 -319
  217. data/ext/sources/ggml/src/ggml-sycl/dpct/helper.hpp +15 -7
  218. data/ext/sources/ggml/src/ggml-sycl/element_wise.cpp +215 -115
  219. data/ext/sources/ggml/src/ggml-sycl/element_wise.hpp +2 -0
  220. data/ext/sources/ggml/src/ggml-sycl/ggml-sycl.cpp +1006 -336
  221. data/ext/sources/ggml/src/ggml-sycl/mmvq.cpp +252 -67
  222. data/ext/sources/ggml/src/ggml-sycl/mmvq.hpp +17 -0
  223. data/ext/sources/ggml/src/ggml-sycl/norm.cpp +103 -49
  224. data/ext/sources/ggml/src/ggml-sycl/outprod.cpp +45 -9
  225. data/ext/sources/ggml/src/ggml-sycl/pool.cpp +185 -0
  226. data/ext/sources/ggml/src/ggml-sycl/pool.hpp +22 -0
  227. data/ext/sources/ggml/src/ggml-sycl/presets.hpp +3 -1
  228. data/ext/sources/ggml/src/ggml-sycl/set_rows.cpp +10 -2
  229. data/ext/sources/ggml/src/ggml-sycl/softmax.cpp +9 -10
  230. data/ext/sources/ggml/src/ggml-sycl/vecdotq.hpp +35 -0
  231. data/ext/sources/ggml/src/ggml-vulkan/CMakeLists.txt +5 -0
  232. data/ext/sources/ggml/src/ggml-vulkan/ggml-vulkan.cpp +833 -215
  233. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/col2im_1d.comp +61 -0
  234. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/conv2d_mm.comp +1 -1
  235. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/conv3d_mm.comp +431 -0
  236. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/diag.comp +3 -3
  237. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp +1 -0
  238. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp +1 -0
  239. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/generic_unary_head.glsl +21 -19
  240. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/get_rows_back.comp +25 -0
  241. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/glu_head.glsl +23 -4
  242. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/glu_main.glsl +14 -18
  243. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/l2_norm.comp +4 -7
  244. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq.comp +21 -24
  245. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp +31 -23
  246. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp +6 -5
  247. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl +84 -67
  248. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/norm.comp +10 -10
  249. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/repeat_back.comp +3 -3
  250. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/roll.comp +3 -3
  251. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/tri.comp +3 -3
  252. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/unary.comp +168 -0
  253. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +121 -74
  254. data/ext/sources/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp +26 -19
  255. data/ext/sources/ggml/src/ggml-webgpu/ggml-webgpu.cpp +31 -36
  256. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/common_decls.tmpl +16 -2
  257. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_split.wgsl +7 -7
  258. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/get_rows.wgsl +21 -0
  259. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_decls.tmpl +439 -320
  260. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_id_vec.wgsl +2 -2
  261. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec.wgsl +45 -39
  262. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_acc.tmpl +586 -465
  263. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_q_acc.tmpl +63 -69
  264. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/quantize_q8.wgsl +14 -9
  265. data/ext/sources/ggml/src/ggml.c +36 -14
  266. data/ext/sources/include/whisper.h +21 -0
  267. data/ext/sources/src/whisper.cpp +164 -14
  268. data/lib/whisper/log_settable.rb +5 -8
  269. data/lib/whisper/model/uri.rb +0 -7
  270. data/sig/whisper.rbs +6 -0
  271. data/test/test_vad.rb +9 -0
  272. data/test/test_vad_context.rb +2 -2
  273. data/whispercpp.gemspec +1 -1
  274. metadata +62 -37
  275. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-flash-attn-ops.c +0 -1878
  276. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-matmul-ops.c +0 -2066
  277. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-ops.c +0 -6
  278. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-ops.h +0 -88
  279. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-profile.h +0 -34
  280. data/ext/sources/ggml/src/ggml-hexagon/htp/vtcm-utils.h +0 -16
  281. data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_gelu.cpp +0 -25
  282. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/abs.comp +0 -21
  283. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/ceil.comp +0 -22
  284. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/clamp.comp +0 -17
  285. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/cos.comp +0 -17
  286. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/elu.comp +0 -27
  287. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/exp.comp +0 -20
  288. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/floor.comp +0 -22
  289. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu.comp +0 -25
  290. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu_erf.comp +0 -39
  291. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu_quick.comp +0 -23
  292. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/hardsigmoid.comp +0 -22
  293. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/hardswish.comp +0 -22
  294. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/leaky_relu.comp +0 -22
  295. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/neg.comp +0 -20
  296. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/relu.comp +0 -21
  297. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/round.comp +0 -29
  298. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sgn.comp +0 -21
  299. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sigmoid.comp +0 -20
  300. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/silu.comp +0 -22
  301. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sin.comp +0 -17
  302. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/softplus.comp +0 -23
  303. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sqrt.comp +0 -17
  304. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/square.comp +0 -17
  305. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/step.comp +0 -22
  306. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/tanh.comp +0 -20
  307. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/trunc.comp +0 -22
  308. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/xielu.comp +0 -35
@@ -98,6 +98,7 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
98
98
  }
99
99
  #endif // INIT_SRC0_SHMEM_Q1_0
100
100
 
101
+ // legacy-quants
101
102
  #if defined(INIT_SRC0_SHMEM_Q4_0) || defined(INIT_SRC0_SHMEM_Q4_1) || defined(INIT_SRC0_SHMEM_Q5_0) || defined(INIT_SRC0_SHMEM_Q5_1) || defined(INIT_SRC0_SHMEM_Q8_0) || defined(INIT_SRC0_SHMEM_Q8_1) || defined(INIT_SRC0_SHMEM_MXFP4)
102
103
  const BLOCK_SIZE = 32u;
103
104
  // the number of blocks per k-tile. Note that this currently only works if TILE_K is a multiple of BLOCK_SIZE, which may need to be rethought for larger quantized types.
@@ -124,7 +125,7 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
124
125
  if (global_m < params.m && global_block_k < params.k / BLOCK_SIZE) {
125
126
  let src0_idx = batch_offset + global_m * params.stride_01 + global_block_k;
126
127
 
127
- #ifdef INIT_SRC0_SHMEM_Q4_0
128
+ #if defined(INIT_SRC0_SHMEM_Q4_0)
128
129
  let block_byte_base = src0_idx * 18u; // BLOCK_SIZE_BYTES = 18u;
129
130
  let d = load_f16_at_src0(block_byte_base);
130
131
 
@@ -134,7 +135,9 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
134
135
  let q_packed = load_u32_at_src0(q_byte_offset);
135
136
  dequant_q4_0_packed_to_shmem(q_packed, d, shmem_idx + j * BYTES_PER_INNER_LOOP);
136
137
  }
137
- #elif INIT_SRC0_SHMEM_Q4_1
138
+ #endif // INIT_SRC0_SHMEM_Q4_0
139
+
140
+ #if defined(INIT_SRC0_SHMEM_Q4_1)
138
141
  let block_byte_base = src0_idx * 20u; // BLOCK_SIZE_BYTES = 20u;
139
142
  let dm = unpack2x16float(load_u32_at_src0_aligned(block_byte_base));
140
143
  let d = f16(dm[0]);
@@ -153,7 +156,9 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
153
156
  shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k + 16u] = q_hi;
154
157
  }
155
158
  }
156
- #elif INIT_SRC0_SHMEM_Q5_0
159
+ #endif // INIT_SRC0_SHMEM_Q4_1
160
+
161
+ #if defined(INIT_SRC0_SHMEM_Q5_0)
157
162
  let block_byte_base = src0_idx * 22u; // BLOCK_SIZE_BYTES = 22u;
158
163
 
159
164
  let d = load_f16_at_src0(block_byte_base);
@@ -176,7 +181,9 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
176
181
  shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k + 16u] = q_hi;
177
182
  }
178
183
  }
179
- #elif INIT_SRC0_SHMEM_Q5_1
184
+ #endif // INIT_SRC0_SHMEM_Q5_0
185
+
186
+ #if defined(INIT_SRC0_SHMEM_Q5_1)
180
187
  let block_byte_base = src0_idx * 24u; // BLOCK_SIZE_BYTES = 24u;
181
188
 
182
189
  let dm = unpack2x16float(load_u32_at_src0_aligned(block_byte_base));
@@ -201,7 +208,9 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
201
208
  shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k + 16u] = q_hi;
202
209
  }
203
210
  }
204
- #elif INIT_SRC0_SHMEM_Q8_0
211
+ #endif // INIT_SRC0_SHMEM_Q5_1
212
+
213
+ #if defined(INIT_SRC0_SHMEM_Q8_0)
205
214
  let block_byte_base = src0_idx * 34u; // BLOCK_SIZE_BYTES = 34u;
206
215
  let d = load_f16_at_src0(block_byte_base);
207
216
 
@@ -211,7 +220,9 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
211
220
  let q_packed = load_u32_at_src0(q_byte_offset);
212
221
  dequant_q8_0_packed_to_shmem(q_packed, d, shmem_idx + j * BYTES_PER_INNER_LOOP);
213
222
  }
214
- #elif INIT_SRC0_SHMEM_Q8_1
223
+ #endif // INIT_SRC0_SHMEM_Q8_0
224
+
225
+ #if defined(INIT_SRC0_SHMEM_Q8_1)
215
226
  let block_byte_base = src0_idx * 36u; // BLOCK_SIZE_BYTES = 36u;
216
227
  let dm = unpack2x16float(load_u32_at_src0_aligned(block_byte_base));
217
228
  let d = f16(dm[0]);
@@ -227,8 +238,10 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
227
238
  shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k] = q_val;
228
239
  }
229
240
  }
230
- #elif INIT_SRC0_SHMEM_MXFP4
231
- let block_byte_base = src0_idx * 17u;
241
+ #endif // INIT_SRC0_SHMEM_Q8_1
242
+
243
+ #if defined(INIT_SRC0_SHMEM_MXFP4)
244
+ let block_byte_base = src0_idx * 17u; // BLOCK_SIZE_BYTES = 17u;
232
245
  let eu8 = get_byte(load_u32_at_src0_aligned(block_byte_base), block_byte_base & 3u);
233
246
  let e = ldexp(1.0, i32(eu8) - 128);
234
247
 
@@ -244,11 +257,52 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
244
257
  shmem[shmem_idx + j * BYTES_PER_INNER_LOOP + k + 16u] = f16(q_hi);
245
258
  }
246
259
  }
247
- #endif
260
+ #endif // INIT_SRC0_SHMEM_MXFP4
248
261
  }
249
262
  }
250
263
  }
251
- #endif
264
+ #endif // legacy-quants
265
+
266
+ #if defined(INIT_SRC0_SHMEM_NVFP4)
267
+ const BLOCK_SIZE = 64u;
268
+ const BLOCK_SIZE_BYTES = 36u;
269
+ const SUB_BLOCK_SIZE = 16u; // elements sharing one UE4M3 scale
270
+ const NQ = 16u;
271
+ const BYTES_PER_THREAD = 8u;
272
+ const BYTES_PER_INNER_LOOP = 4u;
273
+
274
+ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
275
+ for (var i = thread_id * NQ; i < TILE_SRC0_SHMEM; i += TOTAL_WORKGROUP_SIZE * NQ) {
276
+ let tile_m = i / TILE_K;
277
+ let tile_k_start = i % TILE_K;
278
+ let global_m = offset_m + tile_m;
279
+ let global_k_start = k_outer + tile_k_start;
280
+
281
+ if (global_m >= params.m) {
282
+ break;
283
+ }
284
+
285
+ let block_k = global_k_start / BLOCK_SIZE;
286
+ let sub_block = (global_k_start % BLOCK_SIZE) / SUB_BLOCK_SIZE;
287
+ let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
288
+
289
+ let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
290
+ let d_byte_base = block_byte_base;
291
+ let qs_byte_base = block_byte_base + 4u;
292
+
293
+ let d = ue4m3_to_fp32(get_byte(load_u32_at_src0_aligned(d_byte_base), sub_block)) * 0.5;
294
+
295
+ for (var j = 0u; j < BYTES_PER_THREAD / BYTES_PER_INNER_LOOP; j++) {
296
+ let q_packed = load_u32_at_src0_aligned(qs_byte_base + sub_block * 8u + j * 4u);
297
+ for (var k = 0u; k < BYTES_PER_INNER_LOOP; k++) {
298
+ let q_byte = get_byte(q_packed, k);
299
+ shmem[i + j * BYTES_PER_INNER_LOOP + k] = f16(f32(kvalues_mxfp4[q_byte & 0xF]) * d);
300
+ shmem[i + j * BYTES_PER_INNER_LOOP + k + 8u] = f16(f32(kvalues_mxfp4[(q_byte >> 4) & 0xF]) * d);
301
+ }
302
+ }
303
+ }
304
+ }
305
+ #endif // INIT_SRC0_SHMEM_NVFP4
252
306
 
253
307
  // k-quants
254
308
  #if defined(INIT_SRC0_SHMEM_Q2_K) || defined(INIT_SRC0_SHMEM_Q3_K) || defined(INIT_SRC0_SHMEM_Q4_K) || defined(INIT_SRC0_SHMEM_Q5_K) || defined(INIT_SRC0_SHMEM_Q6_K)
@@ -284,7 +338,7 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
284
338
 
285
339
  let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
286
340
 
287
- #ifdef INIT_SRC0_SHMEM_Q2_K
341
+ #if defined(INIT_SRC0_SHMEM_Q2_K)
288
342
  let block_byte_base = src0_idx * 84u; // BLOCK_SIZE_BYTES = 84u;
289
343
  let scales_byte_base = block_byte_base;
290
344
  let qs_byte_base = block_byte_base + 16u;
@@ -314,7 +368,9 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
314
368
  let ml = dmin * f16(scale >> 4u);
315
369
 
316
370
  store_shmem_kquants(qs_vec4 * dl - ml, elem_idx);
317
- #elif INIT_SRC0_SHMEM_Q3_K
371
+ #endif // INIT_SRC0_SHMEM_Q2_K
372
+
373
+ #if defined(INIT_SRC0_SHMEM_Q3_K)
318
374
  let block_byte_base = src0_idx * 110u; // BLOCK_SIZE_BYTES = 110u;
319
375
  let hmask_byte_base = block_byte_base + 0u;
320
376
  let qs_byte_base = block_byte_base + 32u;
@@ -355,7 +411,9 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
355
411
  let dl = d_all * (f16((scale_hi2 << 4u) | scale_low4) - 32.0);
356
412
 
357
413
  store_shmem_kquants(dl * q_vec4, elem_idx);
358
- #elif INIT_SRC0_SHMEM_Q4_K
414
+ #endif // INIT_SRC0_SHMEM_Q3_K
415
+
416
+ #if defined(INIT_SRC0_SHMEM_Q4_K)
359
417
  let block_byte_base = src0_idx * 144u; // BLOCK_SIZE_BYTES = 144u;
360
418
  let dm_byte_base = block_byte_base + 0u;
361
419
  let scale_byte_base = block_byte_base + 4u;
@@ -399,7 +457,9 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
399
457
  let ml = dmin * f16(mn);
400
458
 
401
459
  store_shmem_kquants(dl * qs_vec4 - vec4(ml, ml, ml, ml), elem_idx);
402
- #elif INIT_SRC0_SHMEM_Q5_K
460
+ #endif // INIT_SRC0_SHMEM_Q4_K
461
+
462
+ #if defined(INIT_SRC0_SHMEM_Q5_K)
403
463
  let block_byte_base = src0_idx * 176u; // BLOCK_SIZE_BYTES = 176u;
404
464
  let dm_byte_base = block_byte_base + 0u;
405
465
  let scale_byte_base = block_byte_base + 4u;
@@ -456,7 +516,9 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
456
516
  let ml = dmin * f16(mn);
457
517
 
458
518
  store_shmem_kquants((qh_vec4 + qs_lo4_vec4) * dl - vec4<f16>(ml, ml, ml, ml), elem_idx);
459
- #elif INIT_SRC0_SHMEM_Q6_K
519
+ #endif // INIT_SRC0_SHMEM_Q5_K
520
+
521
+ #if defined(INIT_SRC0_SHMEM_Q6_K)
460
522
  let block_byte_base = src0_idx * 210u; // BLOCK_SIZE_BYTES = 210u;
461
523
  let ql_byte_base = block_byte_base;
462
524
  let qh_byte_base = block_byte_base + 128u;
@@ -497,17 +559,18 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
497
559
  let scale = get_byte_i32(scale_word, scale_byte & 3u);
498
560
 
499
561
  store_shmem_kquants(d * q_vec4 * f16(scale), elem_idx);
500
- #endif
562
+ #endif // INIT_SRC0_SHMEM_Q6_K
501
563
  }
502
564
  }
503
565
  #endif // k-quants
504
566
 
505
- #ifdef INIT_SRC0_SHMEM_IQ4_NL
567
+ #if defined(INIT_SRC0_SHMEM_IQ4_NL)
506
568
  const BLOCK_SIZE = 32u;
507
569
  const BLOCK_SIZE_BYTES = 18u;
570
+ const NQ = 4u;
508
571
 
509
572
  fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
510
- for (var elem_idx = thread_id; elem_idx < TILE_SRC0_SHMEM; elem_idx += TOTAL_WORKGROUP_SIZE) {
573
+ for (var elem_idx = thread_id * NQ; elem_idx < TILE_SRC0_SHMEM; elem_idx += NQ * TOTAL_WORKGROUP_SIZE) {
511
574
  let tile_m = elem_idx / TILE_K;
512
575
  let tile_k = elem_idx % TILE_K;
513
576
  let global_m = offset_m + tile_m;
@@ -519,408 +582,464 @@ fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u3
519
582
  }
520
583
 
521
584
  let block_k = global_k / BLOCK_SIZE;
522
- let k_in_block = global_k % BLOCK_SIZE;
585
+ let k_in_block = global_k % BLOCK_SIZE; // k_in_block % 4 == 0;
586
+
587
+ let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
523
588
 
524
- let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
525
589
  let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
526
- let d = load_f16_at_src0(block_byte_base);
590
+ let d_byte_base = block_byte_base + 0u;
591
+ let qs_byte_base = block_byte_base + 2u;
592
+
593
+ let d = load_f16_at_src0(d_byte_base);
594
+
595
+ let id_qtr = (k_in_block % 16u) / 4u;
596
+ let shift_phase = k_in_block / 16u;
527
597
 
528
- let pos = k_in_block % 16u;
529
- let nib_shift = (k_in_block / 16u) * 4u;
530
- let q_packed = load_u32_at_src0(block_byte_base + 2u + (pos / 4u) * 4u);
531
- let nib = (get_byte(q_packed, pos % 4u) >> nib_shift) & 0xFu;
598
+ let qs_u32 = load_u32_at_src0(qs_byte_base + 4u * id_qtr);
532
599
 
533
- shmem[elem_idx] = d * f16(kvalues_iq4nl[nib]);
600
+ shmem[elem_idx + 0u] = d * f16(kvalues_iq4nl[(qs_u32 >> ( 0u + 4u * shift_phase)) & 0xFu]);
601
+ shmem[elem_idx + 1u] = d * f16(kvalues_iq4nl[(qs_u32 >> ( 8u + 4u * shift_phase)) & 0xFu]);
602
+ shmem[elem_idx + 2u] = d * f16(kvalues_iq4nl[(qs_u32 >> (16u + 4u * shift_phase)) & 0xFu]);
603
+ shmem[elem_idx + 3u] = d * f16(kvalues_iq4nl[(qs_u32 >> (24u + 4u * shift_phase)) & 0xFu]);
534
604
  }
535
605
  }
536
606
  #endif // INIT_SRC0_SHMEM_IQ4_NL
537
607
 
538
- #ifdef INIT_SRC0_SHMEM_IQ4_XS
608
+ // i-quants (super block size: 256)
609
+ #if defined(INIT_SRC0_SHMEM_IQ4_XS) || defined(INIT_SRC0_SHMEM_IQ1_S) || defined(INIT_SRC0_SHMEM_IQ1_M) || defined(INIT_SRC0_SHMEM_IQ2_XXS) \
610
+ || defined(INIT_SRC0_SHMEM_IQ2_XS) || defined(INIT_SRC0_SHMEM_IQ2_S) || defined(INIT_SRC0_SHMEM_IQ3_XXS) || defined(INIT_SRC0_SHMEM_IQ3_S)
539
611
  const BLOCK_SIZE = 256u;
540
- const BLOCK_SIZE_BYTES = 136u;
612
+ const NQ = 16u;
541
613
 
542
- fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
543
- for (var elem_idx = thread_id; elem_idx < TILE_SRC0_SHMEM; elem_idx += TOTAL_WORKGROUP_SIZE) {
544
- let tile_m = elem_idx / TILE_K;
545
- let tile_k = elem_idx % TILE_K;
546
- let global_m = offset_m + tile_m;
547
- let global_k = k_outer + tile_k;
614
+ fn store_shmem_iquants(val: vec4<f16>, idx: u32) {
615
+ shmem[idx] = val.x;
616
+ shmem[idx + 1] = val.y;
617
+ shmem[idx + 2] = val.z;
618
+ shmem[idx + 3] = val.w;
619
+ }
548
620
 
549
- if (global_m >= params.m || global_k >= params.k) {
550
- shmem[elem_idx] = f16(0.0);
551
- continue;
552
- }
621
+ fn load_byte_at_src0_aligned(byte_offset: u32) -> u32 {
622
+ return get_byte(load_u32_at_src0_aligned(byte_offset), byte_offset % 4u);
623
+ }
553
624
 
554
- let block_k = global_k / BLOCK_SIZE;
555
- let k_in_block = global_k % BLOCK_SIZE;
625
+ #if defined(INIT_SRC0_SHMEM_IQ1_M) || defined(INIT_SRC0_SHMEM_IQ1_S)
626
+ fn create_iq_gw4(dl: f32, gw: u32, shift_base: u32, delta: f32) -> vec4<f16> {
627
+ return vec4<f16>(
628
+ f16(dl * (f32((bitcast<i32>(((gw >> (shift_base + 0u)) & 3u) << 30u) >> 30u)) + delta)),
629
+ f16(dl * (f32((bitcast<i32>(((gw >> (shift_base + 2u)) & 3u) << 30u) >> 30u)) + delta)),
630
+ f16(dl * (f32((bitcast<i32>(((gw >> (shift_base + 4u)) & 3u) << 30u) >> 30u)) + delta)),
631
+ f16(dl * (f32((bitcast<i32>(((gw >> (shift_base + 6u)) & 3u) << 30u) >> 30u)) + delta)),
632
+ );
633
+ }
634
+ #endif
556
635
 
557
- let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
558
- let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
636
+ #if defined(INIT_SRC0_SHMEM_IQ4_XS)
637
+ fn create_iq_gw4(dl: f16, qs_u32: u32, shift_phase: u32) -> vec4<f16> {
638
+ return vec4<f16>(
639
+ dl * f16(kvalues_iq4nl[(qs_u32 >> (4 * shift_phase + 0u)) & 0xFu]),
640
+ dl * f16(kvalues_iq4nl[(qs_u32 >> (4 * shift_phase + 8u)) & 0xFu]),
641
+ dl * f16(kvalues_iq4nl[(qs_u32 >> (4 * shift_phase + 16u)) & 0xFu]),
642
+ dl * f16(kvalues_iq4nl[(qs_u32 >> (4 * shift_phase + 24u)) & 0xFu]),
643
+ );
644
+ }
645
+ #endif
559
646
 
560
- let d_scales_h = load_u32_at_src0(block_byte_base);
561
- let d = bitcast<vec2<f16>>(d_scales_h).x;
562
- let scales_h = d_scales_h >> 16u;
647
+ #if defined(INIT_SRC0_SHMEM_IQ2_XXS)
648
+ fn create_iq_gw4(ig: u32, grid_phase: u32) -> vec4<f32> {
649
+ return vec4<f32>(
650
+ f32(get_byte(iq2xxs_grid[(ig + grid_phase + 0u) / 4u], (ig + grid_phase + 0u) % 4u)),
651
+ f32(get_byte(iq2xxs_grid[(ig + grid_phase + 1u) / 4u], (ig + grid_phase + 1u) % 4u)),
652
+ f32(get_byte(iq2xxs_grid[(ig + grid_phase + 2u) / 4u], (ig + grid_phase + 2u) % 4u)),
653
+ f32(get_byte(iq2xxs_grid[(ig + grid_phase + 3u) / 4u], (ig + grid_phase + 3u) % 4u)),
654
+ );
655
+ }
656
+ #endif
563
657
 
564
- let ib = k_in_block / 32u;
565
- let pos = k_in_block % 32u;
658
+ #if defined(INIT_SRC0_SHMEM_IQ2_XS)
659
+ fn create_iq_gw4(ig: u32, grid_phase: u32) -> vec4<f32> {
660
+ return vec4<f32>(
661
+ f32(get_byte(iq2xs_grid[(ig + grid_phase + 0u) / 4u], (ig + grid_phase + 0u) % 4u)),
662
+ f32(get_byte(iq2xs_grid[(ig + grid_phase + 1u) / 4u], (ig + grid_phase + 1u) % 4u)),
663
+ f32(get_byte(iq2xs_grid[(ig + grid_phase + 2u) / 4u], (ig + grid_phase + 2u) % 4u)),
664
+ f32(get_byte(iq2xs_grid[(ig + grid_phase + 3u) / 4u], (ig + grid_phase + 3u) % 4u)),
665
+ );
666
+ }
667
+ #endif
566
668
 
567
- let scales_l_word = load_u32_at_src0(block_byte_base + 4u);
568
- let ls_lo = (get_byte(scales_l_word, ib / 2u) >> ((ib & 1u) * 4u)) & 0xFu;
569
- let ls_hi = ((scales_h >> (2u * ib)) & 3u) << 4u;
570
- let dl = d * f16(i32(ls_lo | ls_hi) - 32);
669
+ #if defined(INIT_SRC0_SHMEM_IQ2_S)
670
+ fn create_iq_gw4(ig: u32, grid_phase: u32) -> vec4<f32> {
671
+ return vec4<f32>(
672
+ f32(get_byte(iq2s_grid[(ig + grid_phase + 0u) / 4u], (ig + grid_phase + 0u) % 4u)),
673
+ f32(get_byte(iq2s_grid[(ig + grid_phase + 1u) / 4u], (ig + grid_phase + 1u) % 4u)),
674
+ f32(get_byte(iq2s_grid[(ig + grid_phase + 2u) / 4u], (ig + grid_phase + 2u) % 4u)),
675
+ f32(get_byte(iq2s_grid[(ig + grid_phase + 3u) / 4u], (ig + grid_phase + 3u) % 4u)),
676
+ );
677
+ }
678
+ #endif
571
679
 
572
- let iqs = ib * 16u + (pos % 16u);
573
- let nib_shift = (pos / 16u) * 4u;
574
- let q_packed = load_u32_at_src0(block_byte_base + 8u + (iqs / 4u) * 4u);
575
- let nib = (get_byte(q_packed, iqs % 4u) >> nib_shift) & 0xFu;
680
+ #if defined(INIT_SRC0_SHMEM_IQ3_XXS)
681
+ fn create_iq_gw4(ig: u32) -> vec4<f32> {
682
+ return vec4<f32>(
683
+ f32(get_byte(iq3xxs_grid[ig], 0)),
684
+ f32(get_byte(iq3xxs_grid[ig], 1)),
685
+ f32(get_byte(iq3xxs_grid[ig], 2)),
686
+ f32(get_byte(iq3xxs_grid[ig], 3)),
687
+ );
688
+ }
689
+ #endif
576
690
 
577
- shmem[elem_idx] = dl * f16(kvalues_iq4nl[nib]);
578
- }
691
+ #if defined(INIT_SRC0_SHMEM_IQ3_S)
692
+ fn create_iq_gw4(ig: u32) -> vec4<f32> {
693
+ return vec4<f32>(
694
+ f32(get_byte(iq3s_grid[ig], 0)),
695
+ f32(get_byte(iq3s_grid[ig], 1)),
696
+ f32(get_byte(iq3s_grid[ig], 2)),
697
+ f32(get_byte(iq3s_grid[ig], 3)),
698
+ );
579
699
  }
580
- #endif // INIT_SRC0_SHMEM_IQ4_XS
700
+ #endif
581
701
 
582
- #ifdef INIT_SRC0_SHMEM_IQ1_S
583
- const BLOCK_SIZE = 256u;
584
- const BLOCK_SIZE_BYTES = 50u;
702
+ #if defined(INIT_SRC0_SHMEM_IQ2_XXS) || defined(INIT_SRC0_SHMEM_IQ2_XS) || defined(INIT_SRC0_SHMEM_IQ2_S) \
703
+ || defined(INIT_SRC0_SHMEM_IQ3_XXS) || defined(INIT_SRC0_SHMEM_IQ3_S)
704
+ fn create_iq2_m4(signs: u32, mask_phase: u32) -> vec4<f32> {
705
+ return vec4<f32>(
706
+ select(1.0, -1.0, (get_byte(kmask_iq2xs[mask_phase], 0) & signs) != 0u),
707
+ select(1.0, -1.0, (get_byte(kmask_iq2xs[mask_phase], 1) & signs) != 0u),
708
+ select(1.0, -1.0, (get_byte(kmask_iq2xs[mask_phase], 2) & signs) != 0u),
709
+ select(1.0, -1.0, (get_byte(kmask_iq2xs[mask_phase], 3) & signs) != 0u),
710
+ );
711
+ }
712
+ #endif
585
713
 
586
714
  fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
587
- for (var elem_idx = thread_id; elem_idx < TILE_SRC0_SHMEM; elem_idx += TOTAL_WORKGROUP_SIZE) {
715
+ for (var elem_idx = thread_id * NQ; elem_idx < TILE_SRC0_SHMEM; elem_idx += NQ * TOTAL_WORKGROUP_SIZE) {
588
716
  let tile_m = elem_idx / TILE_K;
589
717
  let tile_k = elem_idx % TILE_K;
590
718
  let global_m = offset_m + tile_m;
591
719
  let global_k = k_outer + tile_k;
592
720
 
593
721
  if (global_m >= params.m || global_k >= params.k) {
594
- shmem[elem_idx] = f16(0.0);
722
+ let zero_vec4 = vec4<f16>(f16(0.0), f16(0.0), f16(0.0), f16(0.0));
723
+ store_shmem_iquants(zero_vec4, elem_idx + 0u);
724
+ store_shmem_iquants(zero_vec4, elem_idx + 4u);
725
+ store_shmem_iquants(zero_vec4, elem_idx + 8u);
726
+ store_shmem_iquants(zero_vec4, elem_idx + 12u);
595
727
  continue;
596
728
  }
597
729
 
598
730
  let block_k = global_k / BLOCK_SIZE;
599
- let k_in_block = global_k % BLOCK_SIZE;
731
+ let k_in_block = global_k % BLOCK_SIZE; // k_in_block % 16 == 0;
600
732
 
601
- let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
602
- let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
603
- let d = load_f16_as_f32_at_src0(block_byte_base);
733
+ let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
604
734
 
605
- let ib = k_in_block / 32u;
606
- let pos = k_in_block % 32u;
607
- let l = pos / 8u;
608
- let j = pos % 8u;
735
+ #if defined(INIT_SRC0_SHMEM_IQ4_XS)
736
+ let block_byte_base = src0_idx * 136u; // BLOCK_SIZE_BYTES = 136u;
737
+ let d_byte_base = block_byte_base + 0u;
738
+ let scales_l_byte_base = block_byte_base + 4u;
739
+ let qs_byte_base = block_byte_base + 8u;
609
740
 
610
- let qh = load_u32_at_src0(block_byte_base + 34u + ib * 2u) & 0xFFFFu;
611
- let dl = d * (2.0 * f32((qh >> 12u) & 7u) + 1.0);
612
- let delta = select(IQ1_DELTA, -IQ1_DELTA, (qh & 0x8000u) != 0u);
741
+ let d_scales_h = load_u32_at_src0_aligned(d_byte_base);
742
+ let d = bitcast<vec2<f16>>(d_scales_h).x;
743
+ let scales_h = d_scales_h >> 16u;
613
744
 
614
- let qs_w = load_u32_at_src0(block_byte_base + 2u + ib * 4u);
615
- let ig = (get_byte(qs_w, l) | (((qh >> (3u * l)) & 7u) << 8u)) * 8u;
745
+ let sub_block = k_in_block / 32u;
746
+ let phase = (k_in_block / NQ) % 2u;
616
747
 
617
- let gw = iq1_grid[(ig + j) / 16u];
618
- let g = (gw >> (((ig + j) % 16u) * 2u)) & 3u;
619
- let gs = bitcast<i32>(g << 30u) >> 30u;
748
+ let scales_l_u32 = load_u32_at_src0_aligned(scales_l_byte_base);
749
+ let ls_lo = (get_byte(scales_l_u32, sub_block / 2u) >> (4u * (sub_block % 2u))) & 0xFu;
750
+ let ls_hi = ((scales_h >> (2u * sub_block)) & 3u) << 4u;
751
+ let dl = d * f16(i32(ls_lo | ls_hi) - 32);
620
752
 
621
- shmem[elem_idx] = f16(dl * (f32(gs) + delta));
622
- }
623
- }
624
- #endif // INIT_SRC0_SHMEM_IQ1_S
753
+ let qs_0_3_u32 = load_u32_at_src0_aligned(qs_byte_base + 16u * sub_block + 0u);
754
+ let qs_4_7_u32 = load_u32_at_src0_aligned(qs_byte_base + 16u * sub_block + 4u);
755
+ let qs_8_11_u32 = load_u32_at_src0_aligned(qs_byte_base + 16u * sub_block + 8u);
756
+ let qs_12_15_u32 = load_u32_at_src0_aligned(qs_byte_base + 16u * sub_block + 12u);
625
757
 
626
- #ifdef INIT_SRC0_SHMEM_IQ1_M
627
- const BLOCK_SIZE = 256u;
628
- const BLOCK_SIZE_BYTES = 56u;
758
+ store_shmem_iquants(create_iq_gw4(dl, qs_0_3_u32, phase), elem_idx + 0u);
759
+ store_shmem_iquants(create_iq_gw4(dl, qs_4_7_u32, phase), elem_idx + 4u);
760
+ store_shmem_iquants(create_iq_gw4(dl, qs_8_11_u32, phase), elem_idx + 8u);
761
+ store_shmem_iquants(create_iq_gw4(dl, qs_12_15_u32, phase), elem_idx + 12u);
762
+ #endif // INIT_SRC0_SHMEM_IQ4_XS
629
763
 
630
- fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
631
- for (var elem_idx = thread_id; elem_idx < TILE_SRC0_SHMEM; elem_idx += TOTAL_WORKGROUP_SIZE) {
632
- let tile_m = elem_idx / TILE_K;
633
- let tile_k = elem_idx % TILE_K;
634
- let global_m = offset_m + tile_m;
635
- let global_k = k_outer + tile_k;
764
+ #if defined(INIT_SRC0_SHMEM_IQ1_S)
765
+ let block_byte_base = src0_idx * 50u; // BLOCK_SIZE_BYTES = 50u;
766
+ let d_byte_base = block_byte_base + 0u;
767
+ let qs_byte_base = block_byte_base + 2u;
768
+ let qh_byte_base = block_byte_base + 34u;
636
769
 
637
- if (global_m >= params.m || global_k >= params.k) {
638
- shmem[elem_idx] = f16(0.0);
639
- continue;
640
- }
770
+ let d = load_f16_as_f32_at_src0(d_byte_base);
641
771
 
642
- let block_k = global_k / BLOCK_SIZE;
643
- let k_in_block = global_k % BLOCK_SIZE;
772
+ let sub_block = k_in_block / 32u;
773
+ let phase = (k_in_block / NQ) % 2u;
644
774
 
645
- let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
646
- let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
775
+ let qh_u16 = load_u32_at_src0(qh_byte_base + sub_block * 2u) & 0xFFFFu;
776
+ let qs_u16 = load_u32_at_src0(qs_byte_base + sub_block * 4u + phase * 2u) & 0xFFFFu;
777
+
778
+ let dl = d * (2.0 * f32((qh_u16 >> 12u) & 7u) + 1.0);
779
+ let delta = select(IQ1_DELTA, -IQ1_DELTA, (qh_u16 & 0x8000u) != 0u);
780
+
781
+ let gp0_grid_id = ((qs_u16 & 0xFFu) | (((qh_u16 >> (phase * 6u)) & 7u) << 8u)) * 8u;
782
+ let gp1_grid_id = (((qs_u16 >> 8) & 0xFFu) | (((qh_u16 >> (phase * 6u + 3u)) & 7u) << 8u)) * 8u;
783
+
784
+ let gp0_gw = iq1_grid[(gp0_grid_id) / 16u];
785
+ let gp1_gw = iq1_grid[(gp1_grid_id) / 16u];
786
+
787
+ let gp0_shift_base = (gp0_grid_id % 16u) * 2u;
788
+ let gp1_shift_base = (gp1_grid_id % 16u) * 2u;
647
789
 
648
- let scales0 = load_u32_at_src0(block_byte_base + 48u);
649
- let scales1 = load_u32_at_src0(block_byte_base + 52u);
790
+ store_shmem_iquants(create_iq_gw4(dl, gp0_gw, gp0_shift_base + 0u, delta), elem_idx + 0u);
791
+ store_shmem_iquants(create_iq_gw4(dl, gp0_gw, gp0_shift_base + 8u, delta), elem_idx + 4u);
792
+ store_shmem_iquants(create_iq_gw4(dl, gp1_gw, gp1_shift_base + 0u, delta), elem_idx + 8u);
793
+ store_shmem_iquants(create_iq_gw4(dl, gp1_gw, gp1_shift_base + 8u, delta), elem_idx + 12u);
794
+ #endif // INIT_SRC0_SHMEM_IQ1_S
795
+
796
+ #if defined(INIT_SRC0_SHMEM_IQ1_M)
797
+ let block_byte_base = src0_idx * 56u; // BLOCK_SIZE_BYTES = 56u;
798
+ let qs_byte_base = block_byte_base + 0u;
799
+ let qh_byte_base = block_byte_base + 32u;
800
+ let scales_byte_base = block_byte_base + 48u;
801
+
802
+ let scales0 = load_u32_at_src0_aligned(scales_byte_base);
803
+ let scales1 = load_u32_at_src0_aligned(scales_byte_base + 4u);
650
804
  let scale_packed = ((scales0 >> 12u) & 0xFu) |
651
805
  ((scales0 >> 24u) & 0x00F0u) |
652
806
  ((scales1 >> 4u) & 0x0F00u) |
653
807
  ((scales1 >> 16u) & 0xF000u);
654
808
  let d = f32(bitcast<vec2<f16>>(scale_packed).x);
655
809
 
656
- let ib = k_in_block / 32u;
657
- let pos = k_in_block % 32u;
658
- let l = pos / 8u;
659
- let j = pos % 8u;
810
+ let sub_block = k_in_block / 32u;
811
+ let phase = (k_in_block / NQ) % 2u;
660
812
 
661
- let scales = select(scales0, scales1, ib >= 4u);
662
- let sw = (scales >> (16u * ((ib / 2u) % 2u))) & 0xFFFFu;
663
- let s_pair = (sw >> (6u * (ib % 2u) + 3u * (l / 2u))) & 0x7u;
664
- let dl = d * f32(2u * s_pair + 1u);
813
+ let scale_u32 = select(scales0, scales1, sub_block >= 4u);
814
+ let scale_u3 = (scale_u32 >> (16u * ((sub_block / 2u) % 2u) + 6u * (sub_block % 2u) + 3u * phase)) & 0x7u;
815
+ let dl = d * f32(2u * scale_u3 + 1u);
665
816
 
666
- let qh_word = load_u32_at_src0(block_byte_base + 32u + (ib / 2u) * 4u);
667
- let qh = qh_word >> (16u * (ib % 2u));
668
- let qh_nib = (qh >> (4u * l)) & 0xFu;
817
+ let qh_u8 = (load_u32_at_src0_aligned(qh_byte_base + 4u * (sub_block / 2u)) >> (16u * (sub_block % 2u) + 8u * phase)) & 0xFFu;
818
+ let qs_u16 = (load_u32_at_src0_aligned(qs_byte_base + 4u * sub_block) >> (16u * phase)) & 0xFFFFu;
669
819
 
670
- let qs_w = load_u32_at_src0(block_byte_base + ib * 4u);
671
- let idx = get_byte(qs_w, l) | ((qh_nib & 7u) << 8u);
672
- let delta = select(IQ1_DELTA, -IQ1_DELTA, (qh_nib & 0x8u) != 0u);
820
+ let gp0_grid_id = ((qs_u16 & 0xFFu) | ((qh_u8 & 7u) << 8u)) * 8u;
821
+ let gp0_delta = select(IQ1_DELTA, -IQ1_DELTA, (qh_u8 & 0x8u) != 0u);
673
822
 
674
- let ig = idx * 8u;
675
- let gw = iq1_grid[(ig + j) / 16u];
676
- let g = (gw >> (((ig + j) % 16u) * 2u)) & 3u;
677
- let gs = bitcast<i32>(g << 30u) >> 30u;
823
+ let gp1_grid_id = (((qs_u16 >> 8u) & 0xFFu) | (((qh_u8 >> 4u) & 7u) << 8u)) * 8u;
824
+ let gp1_delta = select(IQ1_DELTA, -IQ1_DELTA, (qh_u8 & 0x80u) != 0u);
678
825
 
679
- shmem[elem_idx] = f16(dl * (f32(gs) + delta));
680
- }
681
- }
682
- #endif // INIT_SRC0_SHMEM_IQ1_M
826
+ let gp0_gw = iq1_grid[(gp0_grid_id) / 16u];
827
+ let gp1_gw = iq1_grid[(gp1_grid_id) / 16u];
683
828
 
684
- #ifdef INIT_SRC0_SHMEM_IQ2_XXS
685
- const BLOCK_SIZE = 256u;
686
- const BLOCK_SIZE_BYTES = 66u;
829
+ let gp0_shift_base = (gp0_grid_id % 16u) * 2u;
830
+ let gp1_shift_base = (gp1_grid_id % 16u) * 2u;
687
831
 
688
- fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
689
- for (var elem_idx = thread_id; elem_idx < TILE_SRC0_SHMEM; elem_idx += TOTAL_WORKGROUP_SIZE) {
690
- let tile_m = elem_idx / TILE_K;
691
- let tile_k = elem_idx % TILE_K;
692
- let global_m = offset_m + tile_m;
693
- let global_k = k_outer + tile_k;
832
+ store_shmem_iquants(create_iq_gw4(dl, gp0_gw, gp0_shift_base + 0u, gp0_delta), elem_idx + 0u);
833
+ store_shmem_iquants(create_iq_gw4(dl, gp0_gw, gp0_shift_base + 8u, gp0_delta), elem_idx + 4u);
834
+ store_shmem_iquants(create_iq_gw4(dl, gp1_gw, gp1_shift_base + 0u, gp1_delta), elem_idx + 8u);
835
+ store_shmem_iquants(create_iq_gw4(dl, gp1_gw, gp1_shift_base + 8u, gp1_delta), elem_idx + 12u);
836
+ #endif // INIT_SRC0_SHMEM_IQ1_M
694
837
 
695
- if (global_m >= params.m || global_k >= params.k) {
696
- shmem[elem_idx] = f16(0.0);
697
- continue;
698
- }
838
+ #if defined(INIT_SRC0_SHMEM_IQ2_XXS)
839
+ let block_byte_base = src0_idx * 66u; // BLOCK_SIZE_BYTES = 66u;
840
+ let d_byte_base = block_byte_base + 0u;
841
+ let qs_byte_base = block_byte_base + 2u;
699
842
 
700
- let block_k = global_k / BLOCK_SIZE;
701
- let k_in_block = global_k % BLOCK_SIZE;
843
+ let d = load_f16_as_f32_at_src0(d_byte_base);
702
844
 
703
- let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
704
- let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
705
- let d = load_f16_as_f32_at_src0(block_byte_base);
845
+ let sub_block = k_in_block / 32u;
846
+ let phase = (k_in_block / NQ) % 2u;
706
847
 
707
- let entry_idx = k_in_block / 8u;
708
- let j = k_in_block % 8u;
848
+ let aux0 = load_u32_at_src0(qs_byte_base + 8u * sub_block + 0u);
849
+ let aux1 = load_u32_at_src0(qs_byte_base + 8u * sub_block + 4u);
850
+ let db = d * (0.5 + f32(aux1 >> 28u)) * 0.25;
709
851
 
710
- let ib = entry_idx & ~3u;
711
- let l = entry_idx & 3u;
852
+ let gp0_ig = get_byte(aux0, 2u * phase + 0u) * 8u;
853
+ let gp1_ig = get_byte(aux0, 2u * phase + 1u) * 8u;
712
854
 
713
- let aux0 = load_u32_at_src0(block_byte_base + 2u + ib * 2u);
714
- let aux1 = load_u32_at_src0(block_byte_base + 2u + (ib + 2u) * 2u);
715
- let db = d * (0.5 + f32(aux1 >> 28u)) * 0.25;
855
+ let gp0_is = (aux1 >> (14u * phase + 0u)) & 127u;
856
+ let gp1_is = (aux1 >> (14u * phase + 7u)) & 127u;
716
857
 
717
- let ig = get_byte(aux0, l) * 8u;
718
- let is = (aux1 >> (7u * l)) & 127u;
719
- let signs = get_byte(ksigns_iq2xs[is / 4u], is % 4u);
858
+ let gp0_signs = get_byte(ksigns_iq2xs[gp0_is / 4u], gp0_is % 4u);
859
+ let gp1_signs = get_byte(ksigns_iq2xs[gp1_is / 4u], gp1_is % 4u);
720
860
 
721
- let g = get_byte(iq2xxs_grid[(ig + j) / 4u], (ig + j) % 4u);
722
- let m = select(1.0, -1.0, (get_byte(kmask_iq2xs[j / 4u], j % 4u) & signs) != 0u);
861
+ let m_0_3_val4 = create_iq2_m4(gp0_signs, 0);
862
+ let m_4_7_val4 = create_iq2_m4(gp0_signs, 1);
863
+ let m_8_11_val4 = create_iq2_m4(gp1_signs, 0);
864
+ let m_12_15_val4 = create_iq2_m4(gp1_signs, 1);
723
865
 
724
- shmem[elem_idx] = f16(db * f32(g) * m);
725
- }
726
- }
866
+ let gw_0_3_val4 = create_iq_gw4(gp0_ig, 0);
867
+ let gw_4_7_val4 = create_iq_gw4(gp0_ig, 4);
868
+ let gw_8_11_val4 = create_iq_gw4(gp1_ig, 0);
869
+ let gw_12_15_val4 = create_iq_gw4(gp1_ig, 4);
870
+
871
+ store_shmem_iquants(vec4<f16>(db * m_0_3_val4 * gw_0_3_val4), elem_idx + 0u);
872
+ store_shmem_iquants(vec4<f16>(db * m_4_7_val4 * gw_4_7_val4), elem_idx + 4u);
873
+ store_shmem_iquants(vec4<f16>(db * m_8_11_val4 * gw_8_11_val4), elem_idx + 8u);
874
+ store_shmem_iquants(vec4<f16>(db * m_12_15_val4 * gw_12_15_val4), elem_idx + 12u);
727
875
  #endif // INIT_SRC0_SHMEM_IQ2_XXS
728
876
 
729
- #ifdef INIT_SRC0_SHMEM_IQ2_XS
730
- const BLOCK_SIZE = 256u;
731
- const BLOCK_SIZE_BYTES = 74u;
877
+ #if defined(INIT_SRC0_SHMEM_IQ2_XS)
878
+ let block_byte_base = src0_idx * 74u; // BLOCK_SIZE_BYTES = 74u;
879
+ let d_byte_base = block_byte_base + 0u;
880
+ let qs_byte_base = block_byte_base + 2u;
881
+ let scales_byte_base = block_byte_base + 66u;
732
882
 
733
- fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
734
- for (var elem_idx = thread_id; elem_idx < TILE_SRC0_SHMEM; elem_idx += TOTAL_WORKGROUP_SIZE) {
735
- let tile_m = elem_idx / TILE_K;
736
- let tile_k = elem_idx % TILE_K;
737
- let global_m = offset_m + tile_m;
738
- let global_k = k_outer + tile_k;
883
+ let d = load_f16_as_f32_at_src0(d_byte_base);
739
884
 
740
- if (global_m >= params.m || global_k >= params.k) {
741
- shmem[elem_idx] = f16(0.0);
742
- continue;
743
- }
885
+ let sub_block = k_in_block / 32u;
886
+ let phase = (k_in_block / NQ) % 2u;
744
887
 
745
- let block_k = global_k / BLOCK_SIZE;
746
- let k_in_block = global_k % BLOCK_SIZE;
888
+ let scale = (load_byte_at_src0_aligned(scales_byte_base + 1u * sub_block) >> (4u * phase)) & 0xFu;
889
+ let db = d * (0.5 + f32(scale)) * 0.25;
747
890
 
748
- let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
749
- let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
750
- let d = load_f16_as_f32_at_src0(block_byte_base);
891
+ let qs_u32 = load_u32_at_src0(qs_byte_base + 8u * sub_block + 4u * phase);
751
892
 
752
- let entry_idx = k_in_block / 8u;
753
- let j = k_in_block % 8u;
893
+ let gp0_ig = (qs_u32 & 0x1FFu) * 8u;
894
+ let gp1_ig = ((qs_u32 >> 16u) & 0x1FFu) * 8u;
754
895
 
755
- let ib = entry_idx & ~3u;
756
- let l = entry_idx & 3u;
896
+ let gp0_is = (qs_u32 >> 9u) & 0x7Fu;
897
+ let gp1_is = (qs_u32 >> 25u) & 0x7Fu;
757
898
 
758
- let scales_word = load_u32_at_src0(block_byte_base + 66u + (ib / 16u) * 4u);
759
- let s = get_byte(scales_word, (ib % 16u) / 4u);
760
- let s_nib = select(s & 0xFu, (s >> 4u) & 0xFu, (l / 2u) != 0u);
761
- let dl = d * (0.5 + f32(s_nib)) * 0.25;
899
+ let gp0_signs = get_byte(ksigns_iq2xs[gp0_is / 4u], gp0_is % 4u);
900
+ let gp1_signs = get_byte(ksigns_iq2xs[gp1_is / 4u], gp1_is % 4u);
762
901
 
763
- let qs_word = load_u32_at_src0(block_byte_base + 2u + (ib + l) * 2u);
764
- let qs_val = qs_word & 0xFFFFu;
765
- let ig = (qs_val & 511u) * 8u;
766
- let is = qs_val >> 9u;
767
- let signs = get_byte(ksigns_iq2xs[is / 4u], is % 4u);
902
+ let m_0_3_val4 = create_iq2_m4(gp0_signs, 0);
903
+ let m_4_7_val4 = create_iq2_m4(gp0_signs, 1);
904
+ let m_8_11_val4 = create_iq2_m4(gp1_signs, 0);
905
+ let m_12_15_val4 = create_iq2_m4(gp1_signs, 1);
768
906
 
769
- let g = get_byte(iq2xs_grid[(ig + j) / 4u], (ig + j) % 4u);
770
- let m = select(1.0, -1.0, (get_byte(kmask_iq2xs[j / 4u], j % 4u) & signs) != 0u);
907
+ let gw_0_3_val4 = create_iq_gw4(gp0_ig, 0);
908
+ let gw_4_7_val4 = create_iq_gw4(gp0_ig, 4);
909
+ let gw_8_11_val4 = create_iq_gw4(gp1_ig, 0);
910
+ let gw_12_15_val4 = create_iq_gw4(gp1_ig, 4);
771
911
 
772
- shmem[elem_idx] = f16(dl * f32(g) * m);
773
- }
774
- }
912
+ store_shmem_iquants(vec4<f16>(db * m_0_3_val4 * gw_0_3_val4), elem_idx + 0u);
913
+ store_shmem_iquants(vec4<f16>(db * m_4_7_val4 * gw_4_7_val4), elem_idx + 4u);
914
+ store_shmem_iquants(vec4<f16>(db * m_8_11_val4 * gw_8_11_val4), elem_idx + 8u);
915
+ store_shmem_iquants(vec4<f16>(db * m_12_15_val4 * gw_12_15_val4), elem_idx + 12u);
775
916
  #endif // INIT_SRC0_SHMEM_IQ2_XS
776
917
 
777
- #ifdef INIT_SRC0_SHMEM_IQ2_S
778
- const BLOCK_SIZE = 256u;
779
- const BLOCK_SIZE_BYTES = 82u;
918
+ #if defined(INIT_SRC0_SHMEM_IQ2_S)
919
+ let block_byte_base = src0_idx * 82u; // BLOCK_SIZE_BYTES = 82u;
920
+ let d_byte_base = block_byte_base + 0u;
921
+ let qs_byte_base = block_byte_base + 2u;
922
+ let qh_byte_base = block_byte_base + 66u;
923
+ let scales_byte_base = block_byte_base + 74u;
780
924
 
781
- fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
782
- for (var elem_idx = thread_id; elem_idx < TILE_SRC0_SHMEM; elem_idx += TOTAL_WORKGROUP_SIZE) {
783
- let tile_m = elem_idx / TILE_K;
784
- let tile_k = elem_idx % TILE_K;
785
- let global_m = offset_m + tile_m;
786
- let global_k = k_outer + tile_k;
925
+ let d = load_f16_as_f32_at_src0(d_byte_base);
787
926
 
788
- if (global_m >= params.m || global_k >= params.k) {
789
- shmem[elem_idx] = f16(0.0);
790
- continue;
791
- }
792
-
793
- let block_k = global_k / BLOCK_SIZE;
794
- let k_in_block = global_k % BLOCK_SIZE;
927
+ let sub_block = k_in_block / 32u;
928
+ let phase = (k_in_block / NQ) % 2u;
795
929
 
796
- let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
797
- let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
798
- let d = load_f16_as_f32_at_src0(block_byte_base);
930
+ let scale = (load_byte_at_src0_aligned(scales_byte_base + 1u * sub_block) >> (4u * phase)) & 0xFu;
931
+ let db = d * (0.5 + f32(scale)) * 0.25;
799
932
 
800
- let ib = k_in_block / 32u;
801
- let l = (k_in_block % 32u) / 8u;
802
- let j = k_in_block % 8u;
933
+ let qs_u16 = load_u32_at_src0(qs_byte_base + 4u * sub_block + 2u * phase) & 0xFFFFu;
934
+ let signs_u16 = load_u32_at_src0(qs_byte_base + 32u + 4u * sub_block + 2u * phase) & 0xFFFFu;
935
+ let qh_u4 = (load_byte_at_src0_aligned(qh_byte_base + 1u * sub_block) >> (4u * phase)) & 0xFu;
803
936
 
804
- let scales_word = load_u32_at_src0(block_byte_base + 74u + (ib / 4u) * 4u);
805
- let s = get_byte(scales_word, ib % 4u);
806
- let s_nib = select(s & 0xFu, (s >> 4u) & 0xFu, (l / 2u) != 0u);
807
- let dl = d * (0.5 + f32(s_nib)) * 0.25;
937
+ let gp0_ig = ((qs_u16 & 0xFFu) | ((qh_u4 & 0x3u) << 8u)) * 8u;
938
+ let gp1_ig = (((qs_u16 >> 8u) & 0xFFu) | ((qh_u4 & 0xCu) << 6u)) * 8u;
808
939
 
809
- let qs_word = load_u32_at_src0(block_byte_base + 2u + ib * 4u);
810
- let qh_word = load_u32_at_src0(block_byte_base + 66u + (ib / 4u) * 4u);
811
- let qh_b = (get_byte(qh_word, ib % 4u) << (8u - 2u * l)) & 0x300u;
812
- let ig = (get_byte(qs_word, l) | qh_b) * 8u;
940
+ let gp0_signs = get_byte(signs_u16, 0);
941
+ let gp1_signs = get_byte(signs_u16, 1);
813
942
 
814
- let signs_word = load_u32_at_src0(block_byte_base + 34u + ib * 4u);
815
- let signs = get_byte(signs_word, l);
943
+ let m_0_3_val4 = create_iq2_m4(gp0_signs, 0);
944
+ let m_4_7_val4 = create_iq2_m4(gp0_signs, 1);
945
+ let m_8_11_val4 = create_iq2_m4(gp1_signs, 0);
946
+ let m_12_15_val4 = create_iq2_m4(gp1_signs, 1);
816
947
 
817
- let g = get_byte(iq2s_grid[(ig + j) / 4u], (ig + j) % 4u);
818
- let m = select(1.0, -1.0, (get_byte(kmask_iq2xs[j / 4u], j % 4u) & signs) != 0u);
948
+ let gw_0_3_val4 = create_iq_gw4(gp0_ig, 0);
949
+ let gw_4_7_val4 = create_iq_gw4(gp0_ig, 4);
950
+ let gw_8_11_val4 = create_iq_gw4(gp1_ig, 0);
951
+ let gw_12_15_val4 = create_iq_gw4(gp1_ig, 4);
819
952
 
820
- shmem[elem_idx] = f16(dl * f32(g) * m);
821
- }
822
- }
953
+ store_shmem_iquants(vec4<f16>(db * m_0_3_val4 * gw_0_3_val4), elem_idx + 0u);
954
+ store_shmem_iquants(vec4<f16>(db * m_4_7_val4 * gw_4_7_val4), elem_idx + 4u);
955
+ store_shmem_iquants(vec4<f16>(db * m_8_11_val4 * gw_8_11_val4), elem_idx + 8u);
956
+ store_shmem_iquants(vec4<f16>(db * m_12_15_val4 * gw_12_15_val4), elem_idx + 12u);
823
957
  #endif // INIT_SRC0_SHMEM_IQ2_S
824
958
 
825
- #ifdef INIT_SRC0_SHMEM_IQ3_XXS
826
- const BLOCK_SIZE = 256u;
827
- const BLOCK_SIZE_BYTES = 98u;
959
+ #if defined(INIT_SRC0_SHMEM_IQ3_XXS)
960
+ let block_byte_base = src0_idx * 98u; // BLOCK_SIZE_BYTES = 98u;
961
+ let d_byte_base = block_byte_base + 0u;
962
+ let qs_byte_base = block_byte_base + 2u;
828
963
 
829
- fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
830
- for (var elem_idx = thread_id; elem_idx < TILE_SRC0_SHMEM; elem_idx += TOTAL_WORKGROUP_SIZE) {
831
- let tile_m = elem_idx / TILE_K;
832
- let tile_k = elem_idx % TILE_K;
833
- let global_m = offset_m + tile_m;
834
- let global_k = k_outer + tile_k;
964
+ let d = load_f16_as_f32_at_src0(d_byte_base);
835
965
 
836
- if (global_m >= params.m || global_k >= params.k) {
837
- shmem[elem_idx] = f16(0.0);
838
- continue;
839
- }
966
+ let sub_block = k_in_block / 32u;
967
+ let phase = (k_in_block / NQ) % 2u;
840
968
 
841
- let block_k = global_k / BLOCK_SIZE;
842
- let k_in_block = global_k % BLOCK_SIZE;
969
+ let qs_u32 = load_u32_at_src0(qs_byte_base + 8u * sub_block + 4u * phase);
970
+ let sign_u32 = load_u32_at_src0(qs_byte_base + 64u + 4u * sub_block);
971
+ let db = d * (0.5 + f32(sign_u32 >> 28u)) * 0.5;
843
972
 
844
- let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
845
- let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
846
- let d = load_f16_as_f32_at_src0(block_byte_base);
847
-
848
- let ib_pair = k_in_block / 32u;
849
- let in_pair = k_in_block % 32u;
850
- let l = in_pair / 8u;
851
- let in_l = in_pair % 8u;
852
- let k2 = in_l / 4u;
853
- let j = in_l % 4u;
854
-
855
- let ib = ib_pair * 2u;
856
- let sc_sign_off = block_byte_base + 2u + (ib + 32u) * 2u;
857
- let sc_sign = load_u32_at_src0(sc_sign_off);
858
- let db = d * (0.5 + f32(sc_sign >> 28u)) * 0.5;
859
- let is = (sc_sign >> (7u * l)) & 127u;
860
- let signs = get_byte(ksigns_iq2xs[is / 4u], is % 4u);
861
-
862
- let ig_word = load_u32_at_src0(block_byte_base + 2u + (ib * 2u + l) * 2u) & 0xFFFFu;
863
- let ig_byte = get_byte(ig_word, k2);
864
- let g = get_byte(iq3xxs_grid[ig_byte], j);
865
- let m = select(1.0, -1.0, (get_byte(kmask_iq2xs[k2], j) & signs) != 0u);
866
-
867
- shmem[elem_idx] = f16(db * f32(g) * m);
868
- }
869
- }
870
- #endif // INIT_SRC0_SHMEM_IQ3_XXS
973
+ let ig_0_3 = get_byte(qs_u32, 0);
974
+ let ig_4_7 = get_byte(qs_u32, 1);
975
+ let ig_8_11 = get_byte(qs_u32, 2);
976
+ let ig_12_15 = get_byte(qs_u32, 3);
871
977
 
872
- #ifdef INIT_SRC0_SHMEM_IQ3_S
873
- const BLOCK_SIZE = 256u;
874
- const BLOCK_SIZE_BYTES = 110u;
978
+ let gp0_is = (sign_u32 >> (14u * phase + 0u)) & 0x7Fu;
979
+ let gp1_is = (sign_u32 >> (14u * phase + 7u)) & 0x7Fu;
875
980
 
876
- fn init_shmem_src0(thread_id: u32, batch_offset: u32, offset_m: u32, k_outer: u32) {
877
- for (var elem_idx = thread_id; elem_idx < TILE_SRC0_SHMEM; elem_idx += TOTAL_WORKGROUP_SIZE) {
878
- let tile_m = elem_idx / TILE_K;
879
- let tile_k = elem_idx % TILE_K;
880
- let global_m = offset_m + tile_m;
881
- let global_k = k_outer + tile_k;
981
+ let gp0_signs = get_byte(ksigns_iq2xs[gp0_is / 4u], gp0_is % 4u);
982
+ let gp1_signs = get_byte(ksigns_iq2xs[gp1_is / 4u], gp1_is % 4u);
882
983
 
883
- if (global_m >= params.m || global_k >= params.k) {
884
- shmem[elem_idx] = f16(0.0);
885
- continue;
886
- }
984
+ let m_0_3_val4 = create_iq2_m4(gp0_signs, 0);
985
+ let m_4_7_val4 = create_iq2_m4(gp0_signs, 1);
986
+ let m_8_11_val4 = create_iq2_m4(gp1_signs, 0);
987
+ let m_12_15_val4 = create_iq2_m4(gp1_signs, 1);
887
988
 
888
- let block_k = global_k / BLOCK_SIZE;
889
- let k_in_block = global_k % BLOCK_SIZE;
890
-
891
- let src0_idx = batch_offset + global_m * params.stride_01 + block_k;
892
- let block_byte_base = src0_idx * BLOCK_SIZE_BYTES;
893
- let d = load_f16_as_f32_at_src0(block_byte_base);
894
-
895
- let ib = k_in_block / 64u;
896
- let rest = k_in_block % 64u;
897
- let k = rest / 32u;
898
- let in_k = rest % 32u;
899
- let l = in_k / 8u;
900
- let in_l = in_k % 8u;
901
- let k2 = in_l / 4u;
902
- let j = in_l % 4u;
903
-
904
- let scales_word = load_u32_at_src0(block_byte_base + 106u);
905
- let s = get_byte(scales_word, ib);
906
- let s_nib = select(s & 0xFu, (s >> 4u) & 0xFu, k != 0u);
907
- let dl = d * (1.0 + 2.0 * f32(s_nib));
908
-
909
- let qh_word = load_u32_at_src0(block_byte_base + 66u + (ib / 2u) * 4u);
910
- let qh_byte = get_byte(qh_word, (ib % 2u) * 2u + k);
989
+ let gw_0_3_val4 = create_iq_gw4(ig_0_3);
990
+ let gw_4_7_val4 = create_iq_gw4(ig_4_7);
991
+ let gw_8_11_val4 = create_iq_gw4(ig_8_11);
992
+ let gw_12_15_val4 = create_iq_gw4(ig_12_15);
911
993
 
912
- let ig_word = load_u32_at_src0(block_byte_base + 2u + (ib * 8u + k * 4u + l) * 2u) & 0xFFFFu;
913
- let ig_lo = get_byte(ig_word, 0u) | ((qh_byte << (8u - 2u * l)) & 256u);
914
- let ig_hi = get_byte(ig_word, 1u) | ((qh_byte << (7u - 2u * l)) & 256u);
915
- let ig = select(ig_lo, ig_hi, k2 != 0u);
916
-
917
- let signs_word = load_u32_at_src0(block_byte_base + 74u + (ib * 2u + k) * 4u);
918
- let signs = get_byte(signs_word, l);
919
-
920
- let g = get_byte(iq3s_grid[ig], j);
921
- let m = select(1.0, -1.0, (get_byte(kmask_iq2xs[k2], j) & signs) != 0u);
994
+ store_shmem_iquants(vec4<f16>(db * m_0_3_val4 * gw_0_3_val4), elem_idx + 0u);
995
+ store_shmem_iquants(vec4<f16>(db * m_4_7_val4 * gw_4_7_val4), elem_idx + 4u);
996
+ store_shmem_iquants(vec4<f16>(db * m_8_11_val4 * gw_8_11_val4), elem_idx + 8u);
997
+ store_shmem_iquants(vec4<f16>(db * m_12_15_val4 * gw_12_15_val4), elem_idx + 12u);
998
+ #endif // INIT_SRC0_SHMEM_IQ3_XXS
922
999
 
923
- shmem[elem_idx] = f16(dl * f32(g) * m);
1000
+ #if defined(INIT_SRC0_SHMEM_IQ3_S)
1001
+ let block_byte_base = src0_idx * 110u; // BLOCK_SIZE_BYTES = 110u;
1002
+ let d_byte_base = block_byte_base + 0u;
1003
+ let qs_byte_base = block_byte_base + 2u;
1004
+ let qh_byte_base = block_byte_base + 66u;
1005
+ let signs_byte_base = block_byte_base + 74u;
1006
+ let scales_byte_base = block_byte_base + 106u;
1007
+
1008
+ let d = load_f16_as_f32_at_src0(d_byte_base);
1009
+
1010
+ let sub_block = k_in_block / 32u;
1011
+ let phase = (k_in_block / NQ) % 2u;
1012
+
1013
+ let scale = (load_byte_at_src0_aligned(scales_byte_base + 1u * (sub_block / 2u)) >> (4u * (sub_block % 2u))) & 0xFu;
1014
+ let db = d * (1.0 + 2.0 * f32(scale));
1015
+
1016
+ let qs_u32 = load_u32_at_src0(qs_byte_base + 8u * sub_block + 4u * phase);
1017
+ let qh_u4 = (load_byte_at_src0_aligned(qh_byte_base + 1u * sub_block) >> (4u * phase)) & 0xFu;
1018
+ let signs_u16 = (load_u32_at_src0(signs_byte_base + 4u * sub_block + 2u * phase)) & 0xFFFFu;
1019
+
1020
+ let ig_0_3 = ((qs_u32 >> 0u) & 0xFFu) | ((qh_u4 & 0x1u) << 8u);
1021
+ let ig_4_7 = ((qs_u32 >> 8u) & 0xFFu) | ((qh_u4 & 0x2u) << 7u);
1022
+ let ig_8_11 = ((qs_u32 >> 16u) & 0xFFu) | ((qh_u4 & 0x4u) << 6u);
1023
+ let ig_12_15 = ((qs_u32 >> 24u) & 0xFFu) | ((qh_u4 & 0x8u) << 5u);
1024
+
1025
+ let gp0_signs = get_byte(signs_u16, 0);
1026
+ let gp1_signs = get_byte(signs_u16, 1);
1027
+
1028
+ let m_0_3_val4 = create_iq2_m4(gp0_signs, 0);
1029
+ let m_4_7_val4 = create_iq2_m4(gp0_signs, 1);
1030
+ let m_8_11_val4 = create_iq2_m4(gp1_signs, 0);
1031
+ let m_12_15_val4 = create_iq2_m4(gp1_signs, 1);
1032
+
1033
+ let gw_0_3_val4 = create_iq_gw4(ig_0_3);
1034
+ let gw_4_7_val4 = create_iq_gw4(ig_4_7);
1035
+ let gw_8_11_val4 = create_iq_gw4(ig_8_11);
1036
+ let gw_12_15_val4 = create_iq_gw4(ig_12_15);
1037
+
1038
+ store_shmem_iquants(vec4<f16>(db * m_0_3_val4 * gw_0_3_val4), elem_idx + 0u);
1039
+ store_shmem_iquants(vec4<f16>(db * m_4_7_val4 * gw_4_7_val4), elem_idx + 4u);
1040
+ store_shmem_iquants(vec4<f16>(db * m_8_11_val4 * gw_8_11_val4), elem_idx + 8u);
1041
+ store_shmem_iquants(vec4<f16>(db * m_12_15_val4 * gw_12_15_val4), elem_idx + 12u);
1042
+ #endif // INIT_SRC0_SHMEM_IQ3_S
924
1043
  }
925
1044
  }
926
- #endif // INIT_SRC0_SHMEM_IQ3_S
1045
+ #endif // i-quants (super block size: 256)