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
@@ -0,0 +1,1303 @@
1
+ #include "hmx-utils.h"
2
+ #include "hmx-queue.h"
3
+
4
+ // MXFP4 dequantization LUT: maps 4-bit index to fp16 mantissa value
5
+ // kvalues: 0, 0.5, 1, 1.5, 2, 3, 4, 6, 0, -0.5, -1, -1.5, -2, -3, -4, -6
6
+ static const __fp16 mxfp4_to_fp16_lut[64] __attribute__((aligned(VLEN))) = {
7
+ 0, 0, 0.5, 0, 1, 0, 1.5, 0, 2, 0, 3, 0, 4, 0, 6, 0, 0, 0, -0.5, 0, -1, 0, -1.5, 0, -2, 0, -3, 0, -4, 0, -6, 0,
8
+ };
9
+
10
+ static const __fp16 iq4_nl_to_fp16_lut[64] __attribute__((aligned(VLEN))) = {
11
+ -127, 0, -104, 0, -83, 0, -65, 0, -49, 0, -35, 0, -22, 0, -10, 0,
12
+ 1, 0, 13, 0, 25, 0, 38, 0, 53, 0, 69, 0, 89, 0, 113, 0,
13
+ };
14
+
15
+ // --- tiled format dequantizers ---
16
+
17
+ typedef struct {
18
+ struct htp_context * ctx;
19
+ struct htp_thread_trace * traces;
20
+ __fp16 * dst;
21
+ const uint8_t * src;
22
+
23
+ struct fastdiv_values n_k_tiles_div;
24
+ uint32_t n_k_tiles;
25
+ uint32_t n_tot_tiles;
26
+ uint32_t n_tiles_per_task;
27
+ uint32_t tile_size;
28
+ uint32_t aligned_tile_size;
29
+ uint32_t n_tasks;
30
+ uint32_t n_cols;
31
+ uint32_t k_block;
32
+ size_t row_stride;
33
+ uint32_t weight_type;
34
+ } tiled_dequantize_state_t;
35
+
36
+ // Dequantize a single tile from tiled weight data (already in VTCM) to tile-major FP16.
37
+ static void dequantize_tiled_weight_to_fp16_task_q4_0(
38
+ const tiled_dequantize_state_t *state,
39
+ uint32_t start_tile, uint32_t end_tile) {
40
+
41
+ const HVX_Vector mask_h4 = Q6_Vb_vsplat_R(0x0F);
42
+ const HVX_Vector i8 = Q6_Vb_vsplat_R(8);
43
+
44
+ for (uint32_t t = start_tile; t < end_tile; t++) {
45
+ const uint8_t * tile_src = state->src + t * state->aligned_tile_size;
46
+ __fp16 * dst_ptr = state->dst + t * HTP_MM_HMX_TILE_N_ELMS;
47
+
48
+ HVX_Vector v_sc = hvx_vmem(tile_src + 512);
49
+ HVX_Vector v_scale_duplicated = Q6_V_lo_W(Q6_W_vshuff_VVR(v_sc, v_sc, -2));
50
+
51
+ // Load all 4 groups in parallel
52
+ HVX_Vector vq0 = hvx_vmem(tile_src + 0 * 128);
53
+ HVX_Vector vq1 = hvx_vmem(tile_src + 1 * 128);
54
+ HVX_Vector vq2 = hvx_vmem(tile_src + 2 * 128);
55
+ HVX_Vector vq3 = hvx_vmem(tile_src + 3 * 128);
56
+
57
+ // Nibble extraction
58
+ HVX_Vector v_lo0 = Q6_V_vand_VV(vq0, mask_h4);
59
+ HVX_Vector v_hi0 = Q6_Vub_vlsr_VubR(vq0, 4);
60
+ HVX_Vector v_lo1 = Q6_V_vand_VV(vq1, mask_h4);
61
+ HVX_Vector v_hi1 = Q6_Vub_vlsr_VubR(vq1, 4);
62
+ HVX_Vector v_lo2 = Q6_V_vand_VV(vq2, mask_h4);
63
+ HVX_Vector v_hi2 = Q6_Vub_vlsr_VubR(vq2, 4);
64
+ HVX_Vector v_lo3 = Q6_V_vand_VV(vq3, mask_h4);
65
+ HVX_Vector v_hi3 = Q6_Vub_vlsr_VubR(vq3, 4);
66
+
67
+ // Offsetting (-8)
68
+ v_lo0 = Q6_Vb_vsub_VbVb(v_lo0, i8);
69
+ v_hi0 = Q6_Vb_vsub_VbVb(v_hi0, i8);
70
+ v_lo1 = Q6_Vb_vsub_VbVb(v_lo1, i8);
71
+ v_hi1 = Q6_Vb_vsub_VbVb(v_hi1, i8);
72
+ v_lo2 = Q6_Vb_vsub_VbVb(v_lo2, i8);
73
+ v_hi2 = Q6_Vb_vsub_VbVb(v_hi2, i8);
74
+ v_lo3 = Q6_Vb_vsub_VbVb(v_lo3, i8);
75
+ v_hi3 = Q6_Vb_vsub_VbVb(v_hi3, i8);
76
+
77
+ // Shuffling
78
+ HVX_VectorPair vp_shuf0 = Q6_W_vshuff_VVR(v_hi0, v_lo0, -1);
79
+ HVX_VectorPair vp_shuf1 = Q6_W_vshuff_VVR(v_hi1, v_lo1, -1);
80
+ HVX_VectorPair vp_shuf2 = Q6_W_vshuff_VVR(v_hi2, v_lo2, -1);
81
+ HVX_VectorPair vp_shuf3 = Q6_W_vshuff_VVR(v_hi3, v_lo3, -1);
82
+
83
+ // Unpack to 16-bit
84
+ HVX_VectorPair vp_int16_lo0 = Q6_Wh_vunpack_Vb(Q6_V_lo_W(vp_shuf0));
85
+ HVX_VectorPair vp_int16_hi0 = Q6_Wh_vunpack_Vb(Q6_V_hi_W(vp_shuf0));
86
+ HVX_VectorPair vp_int16_lo1 = Q6_Wh_vunpack_Vb(Q6_V_lo_W(vp_shuf1));
87
+ HVX_VectorPair vp_int16_hi1 = Q6_Wh_vunpack_Vb(Q6_V_hi_W(vp_shuf1));
88
+ HVX_VectorPair vp_int16_lo2 = Q6_Wh_vunpack_Vb(Q6_V_lo_W(vp_shuf2));
89
+ HVX_VectorPair vp_int16_hi2 = Q6_Wh_vunpack_Vb(Q6_V_hi_W(vp_shuf2));
90
+ HVX_VectorPair vp_int16_lo3 = Q6_Wh_vunpack_Vb(Q6_V_lo_W(vp_shuf3));
91
+ HVX_VectorPair vp_int16_hi3 = Q6_Wh_vunpack_Vb(Q6_V_hi_W(vp_shuf3));
92
+
93
+ // Convert and scale multiplication
94
+ HVX_Vector v_grp0_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_lo0)), v_scale_duplicated));
95
+ HVX_Vector v_grp0_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_lo0)), v_scale_duplicated));
96
+ HVX_Vector v_grp0_2 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_hi0)), v_scale_duplicated));
97
+ HVX_Vector v_grp0_3 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_hi0)), v_scale_duplicated));
98
+
99
+ HVX_Vector v_grp1_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_lo1)), v_scale_duplicated));
100
+ HVX_Vector v_grp1_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_lo1)), v_scale_duplicated));
101
+ HVX_Vector v_grp1_2 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_hi1)), v_scale_duplicated));
102
+ HVX_Vector v_grp1_3 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_hi1)), v_scale_duplicated));
103
+
104
+ HVX_Vector v_grp2_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_lo2)), v_scale_duplicated));
105
+ HVX_Vector v_grp2_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_lo2)), v_scale_duplicated));
106
+ HVX_Vector v_grp2_2 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_hi2)), v_scale_duplicated));
107
+ HVX_Vector v_grp2_3 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_hi2)), v_scale_duplicated));
108
+
109
+ HVX_Vector v_grp3_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_lo3)), v_scale_duplicated));
110
+ HVX_Vector v_grp3_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_lo3)), v_scale_duplicated));
111
+ HVX_Vector v_grp3_2 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_hi3)), v_scale_duplicated));
112
+ HVX_Vector v_grp3_3 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_hi3)), v_scale_duplicated));
113
+
114
+ hvx_vmem(dst_ptr + 0 * 64) = v_grp0_0;
115
+ hvx_vmem(dst_ptr + 1 * 64) = v_grp0_1;
116
+ hvx_vmem(dst_ptr + 2 * 64) = v_grp0_2;
117
+ hvx_vmem(dst_ptr + 3 * 64) = v_grp0_3;
118
+
119
+ hvx_vmem(dst_ptr + 4 * 64) = v_grp1_0;
120
+ hvx_vmem(dst_ptr + 5 * 64) = v_grp1_1;
121
+ hvx_vmem(dst_ptr + 6 * 64) = v_grp1_2;
122
+ hvx_vmem(dst_ptr + 7 * 64) = v_grp1_3;
123
+
124
+ hvx_vmem(dst_ptr + 8 * 64) = v_grp2_0;
125
+ hvx_vmem(dst_ptr + 9 * 64) = v_grp2_1;
126
+ hvx_vmem(dst_ptr + 10 * 64) = v_grp2_2;
127
+ hvx_vmem(dst_ptr + 11 * 64) = v_grp2_3;
128
+
129
+ hvx_vmem(dst_ptr + 12 * 64) = v_grp3_0;
130
+ hvx_vmem(dst_ptr + 13 * 64) = v_grp3_1;
131
+ hvx_vmem(dst_ptr + 14 * 64) = v_grp3_2;
132
+ hvx_vmem(dst_ptr + 15 * 64) = v_grp3_3;
133
+ }
134
+ }
135
+
136
+ static void dequantize_tiled_weight_to_fp16_task_q4_1(
137
+ const tiled_dequantize_state_t *state,
138
+ uint32_t start_tile, uint32_t end_tile) {
139
+
140
+ const HVX_Vector mask_h4 = Q6_Vb_vsplat_R(0x0F);
141
+
142
+ for (uint32_t t = start_tile; t < end_tile; t++) {
143
+ const uint8_t * tile_src = state->src + t * state->aligned_tile_size;
144
+ __fp16 * dst_ptr = state->dst + t * HTP_MM_HMX_TILE_N_ELMS;
145
+
146
+ HVX_Vector vscale_offset = hvx_vmem(tile_src + 512);
147
+ HVX_VectorPair dm_deal = Q6_W_vdeal_VVR(vscale_offset, vscale_offset, -2);
148
+ HVX_Vector vd = Q6_V_lo_W(dm_deal);
149
+ HVX_Vector vm = Q6_V_hi_W(dm_deal);
150
+
151
+ HVX_Vector v_scale_duplicated = Q6_V_lo_W(Q6_W_vshuff_VVR(vd, vd, -2));
152
+ HVX_Vector v_offset_duplicated = Q6_V_lo_W(Q6_W_vshuff_VVR(vm, vm, -2));
153
+
154
+ // Load all 4 groups in parallel
155
+ HVX_Vector vq0 = hvx_vmem(tile_src + 0 * 128);
156
+ HVX_Vector vq1 = hvx_vmem(tile_src + 1 * 128);
157
+ HVX_Vector vq2 = hvx_vmem(tile_src + 2 * 128);
158
+ HVX_Vector vq3 = hvx_vmem(tile_src + 3 * 128);
159
+
160
+ // Nibble extraction
161
+ HVX_Vector v_lo0 = Q6_V_vand_VV(vq0, mask_h4);
162
+ HVX_Vector v_hi0 = Q6_Vub_vlsr_VubR(vq0, 4);
163
+ HVX_Vector v_lo1 = Q6_V_vand_VV(vq1, mask_h4);
164
+ HVX_Vector v_hi1 = Q6_Vub_vlsr_VubR(vq1, 4);
165
+ HVX_Vector v_lo2 = Q6_V_vand_VV(vq2, mask_h4);
166
+ HVX_Vector v_hi2 = Q6_Vub_vlsr_VubR(vq2, 4);
167
+ HVX_Vector v_lo3 = Q6_V_vand_VV(vq3, mask_h4);
168
+ HVX_Vector v_hi3 = Q6_Vub_vlsr_VubR(vq3, 4);
169
+
170
+ // Shuffling
171
+ HVX_VectorPair vp_shuf0 = Q6_W_vshuff_VVR(v_hi0, v_lo0, -1);
172
+ HVX_VectorPair vp_shuf1 = Q6_W_vshuff_VVR(v_hi1, v_lo1, -1);
173
+ HVX_VectorPair vp_shuf2 = Q6_W_vshuff_VVR(v_hi2, v_lo2, -1);
174
+ HVX_VectorPair vp_shuf3 = Q6_W_vshuff_VVR(v_hi3, v_lo3, -1);
175
+
176
+ // Unpack to 16-bit
177
+ HVX_VectorPair vp_int16_lo0 = Q6_Wh_vunpack_Vb(Q6_V_lo_W(vp_shuf0));
178
+ HVX_VectorPair vp_int16_hi0 = Q6_Wh_vunpack_Vb(Q6_V_hi_W(vp_shuf0));
179
+ HVX_VectorPair vp_int16_lo1 = Q6_Wh_vunpack_Vb(Q6_V_lo_W(vp_shuf1));
180
+ HVX_VectorPair vp_int16_hi1 = Q6_Wh_vunpack_Vb(Q6_V_hi_W(vp_shuf1));
181
+ HVX_VectorPair vp_int16_lo2 = Q6_Wh_vunpack_Vb(Q6_V_lo_W(vp_shuf2));
182
+ HVX_VectorPair vp_int16_hi2 = Q6_Wh_vunpack_Vb(Q6_V_hi_W(vp_shuf2));
183
+ HVX_VectorPair vp_int16_lo3 = Q6_Wh_vunpack_Vb(Q6_V_lo_W(vp_shuf3));
184
+ HVX_VectorPair vp_int16_hi3 = Q6_Wh_vunpack_Vb(Q6_V_hi_W(vp_shuf3));
185
+
186
+ // Convert, multiply, add offset
187
+ HVX_Vector v_grp0_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_lo0)), v_scale_duplicated), v_offset_duplicated));
188
+ HVX_Vector v_grp0_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_lo0)), v_scale_duplicated), v_offset_duplicated));
189
+ HVX_Vector v_grp0_2 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_hi0)), v_scale_duplicated), v_offset_duplicated));
190
+ HVX_Vector v_grp0_3 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_hi0)), v_scale_duplicated), v_offset_duplicated));
191
+
192
+ HVX_Vector v_grp1_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_lo1)), v_scale_duplicated), v_offset_duplicated));
193
+ HVX_Vector v_grp1_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_lo1)), v_scale_duplicated), v_offset_duplicated));
194
+ HVX_Vector v_grp1_2 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_hi1)), v_scale_duplicated), v_offset_duplicated));
195
+ HVX_Vector v_grp1_3 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_hi1)), v_scale_duplicated), v_offset_duplicated));
196
+
197
+ HVX_Vector v_grp2_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_lo2)), v_scale_duplicated), v_offset_duplicated));
198
+ HVX_Vector v_grp2_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_lo2)), v_scale_duplicated), v_offset_duplicated));
199
+ HVX_Vector v_grp2_2 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_hi2)), v_scale_duplicated), v_offset_duplicated));
200
+ HVX_Vector v_grp2_3 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_hi2)), v_scale_duplicated), v_offset_duplicated));
201
+
202
+ HVX_Vector v_grp3_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_lo3)), v_scale_duplicated), v_offset_duplicated));
203
+ HVX_Vector v_grp3_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_lo3)), v_scale_duplicated), v_offset_duplicated));
204
+ HVX_Vector v_grp3_2 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_hi3)), v_scale_duplicated), v_offset_duplicated));
205
+ HVX_Vector v_grp3_3 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vadd_Vqf16Vhf(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_hi3)), v_scale_duplicated), v_offset_duplicated));
206
+
207
+ // Parallel Stores
208
+ hvx_vmem(dst_ptr + 0 * 64) = v_grp0_0;
209
+ hvx_vmem(dst_ptr + 1 * 64) = v_grp0_1;
210
+ hvx_vmem(dst_ptr + 2 * 64) = v_grp0_2;
211
+ hvx_vmem(dst_ptr + 3 * 64) = v_grp0_3;
212
+
213
+ hvx_vmem(dst_ptr + 4 * 64) = v_grp1_0;
214
+ hvx_vmem(dst_ptr + 5 * 64) = v_grp1_1;
215
+ hvx_vmem(dst_ptr + 6 * 64) = v_grp1_2;
216
+ hvx_vmem(dst_ptr + 7 * 64) = v_grp1_3;
217
+
218
+ hvx_vmem(dst_ptr + 8 * 64) = v_grp2_0;
219
+ hvx_vmem(dst_ptr + 9 * 64) = v_grp2_1;
220
+ hvx_vmem(dst_ptr + 10 * 64) = v_grp2_2;
221
+ hvx_vmem(dst_ptr + 11 * 64) = v_grp2_3;
222
+
223
+ hvx_vmem(dst_ptr + 12 * 64) = v_grp3_0;
224
+ hvx_vmem(dst_ptr + 13 * 64) = v_grp3_1;
225
+ hvx_vmem(dst_ptr + 14 * 64) = v_grp3_2;
226
+ hvx_vmem(dst_ptr + 15 * 64) = v_grp3_3;
227
+ }
228
+ }
229
+
230
+ static void dequantize_tiled_weight_to_fp16_task_iq4_nl(
231
+ const tiled_dequantize_state_t *state,
232
+ uint32_t start_tile, uint32_t end_tile) {
233
+
234
+ const HVX_Vector mask_h4 = Q6_Vb_vsplat_R(0x0F);
235
+ const HVX_Vector vlut_cvt = hvx_vmem(iq4_nl_to_fp16_lut);
236
+
237
+ for (uint32_t t = start_tile; t < end_tile; t++) {
238
+ const uint8_t * tile_src = state->src + t * state->aligned_tile_size;
239
+ __fp16 * dst_ptr = state->dst + t * HTP_MM_HMX_TILE_N_ELMS;
240
+
241
+ HVX_Vector v_sc = hvx_vmem(tile_src + 512);
242
+ HVX_Vector v_scale_duplicated = Q6_V_lo_W(Q6_W_vshuff_VVR(v_sc, v_sc, -2));
243
+
244
+ // Load all 4 groups in parallel
245
+ HVX_Vector vq0 = hvx_vmem(tile_src + 0 * 128);
246
+ HVX_Vector vq1 = hvx_vmem(tile_src + 1 * 128);
247
+ HVX_Vector vq2 = hvx_vmem(tile_src + 2 * 128);
248
+ HVX_Vector vq3 = hvx_vmem(tile_src + 3 * 128);
249
+
250
+ // Nibble extraction
251
+ HVX_Vector v_lo0 = Q6_V_vand_VV(vq0, mask_h4);
252
+ HVX_Vector v_hi0 = Q6_Vub_vlsr_VubR(vq0, 4);
253
+ HVX_Vector v_lo1 = Q6_V_vand_VV(vq1, mask_h4);
254
+ HVX_Vector v_hi1 = Q6_Vub_vlsr_VubR(vq1, 4);
255
+ HVX_Vector v_lo2 = Q6_V_vand_VV(vq2, mask_h4);
256
+ HVX_Vector v_hi2 = Q6_Vub_vlsr_VubR(vq2, 4);
257
+ HVX_Vector v_lo3 = Q6_V_vand_VV(vq3, mask_h4);
258
+ HVX_Vector v_hi3 = Q6_Vub_vlsr_VubR(vq3, 4);
259
+
260
+ // Shuffling
261
+ HVX_VectorPair vp_shuf0 = Q6_W_vshuff_VVR(v_hi0, v_lo0, -1);
262
+ HVX_VectorPair vp_shuf1 = Q6_W_vshuff_VVR(v_hi1, v_lo1, -1);
263
+ HVX_VectorPair vp_shuf2 = Q6_W_vshuff_VVR(v_hi2, v_lo2, -1);
264
+ HVX_VectorPair vp_shuf3 = Q6_W_vshuff_VVR(v_hi3, v_lo3, -1);
265
+
266
+ // Shuffle for LUT lookup
267
+ HVX_Vector v_q_lo0 = Q6_Vb_vshuff_Vb(Q6_V_lo_W(vp_shuf0));
268
+ HVX_Vector v_q_hi0 = Q6_Vb_vshuff_Vb(Q6_V_hi_W(vp_shuf0));
269
+ HVX_Vector v_q_lo1 = Q6_Vb_vshuff_Vb(Q6_V_lo_W(vp_shuf1));
270
+ HVX_Vector v_q_hi1 = Q6_Vb_vshuff_Vb(Q6_V_hi_W(vp_shuf1));
271
+ HVX_Vector v_q_lo2 = Q6_Vb_vshuff_Vb(Q6_V_lo_W(vp_shuf2));
272
+ HVX_Vector v_q_hi2 = Q6_Vb_vshuff_Vb(Q6_V_hi_W(vp_shuf2));
273
+ HVX_Vector v_q_lo3 = Q6_Vb_vshuff_Vb(Q6_V_lo_W(vp_shuf3));
274
+ HVX_Vector v_q_hi3 = Q6_Vb_vshuff_Vb(Q6_V_hi_W(vp_shuf3));
275
+
276
+ // LUT lookup
277
+ HVX_VectorPair vp_lo0 = Q6_Wh_vlut16_VbVhR(v_q_lo0, vlut_cvt, 0);
278
+ HVX_VectorPair vp_hi0 = Q6_Wh_vlut16_VbVhR(v_q_hi0, vlut_cvt, 0);
279
+ HVX_VectorPair vp_lo1 = Q6_Wh_vlut16_VbVhR(v_q_lo1, vlut_cvt, 0);
280
+ HVX_VectorPair vp_hi1 = Q6_Wh_vlut16_VbVhR(v_q_hi1, vlut_cvt, 0);
281
+ HVX_VectorPair vp_lo2 = Q6_Wh_vlut16_VbVhR(v_q_lo2, vlut_cvt, 0);
282
+ HVX_VectorPair vp_hi2 = Q6_Wh_vlut16_VbVhR(v_q_hi2, vlut_cvt, 0);
283
+ HVX_VectorPair vp_lo3 = Q6_Wh_vlut16_VbVhR(v_q_lo3, vlut_cvt, 0);
284
+ HVX_VectorPair vp_hi3 = Q6_Wh_vlut16_VbVhR(v_q_hi3, vlut_cvt, 0);
285
+
286
+ // Convert and scale multiplication
287
+ HVX_Vector v_grp0_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_lo_W(vp_lo0), v_scale_duplicated));
288
+ HVX_Vector v_grp0_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_hi_W(vp_lo0), v_scale_duplicated));
289
+ HVX_Vector v_grp0_2 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_lo_W(vp_hi0), v_scale_duplicated));
290
+ HVX_Vector v_grp0_3 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_hi_W(vp_hi0), v_scale_duplicated));
291
+
292
+ HVX_Vector v_grp1_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_lo_W(vp_lo1), v_scale_duplicated));
293
+ HVX_Vector v_grp1_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_hi_W(vp_lo1), v_scale_duplicated));
294
+ HVX_Vector v_grp1_2 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_lo_W(vp_hi1), v_scale_duplicated));
295
+ HVX_Vector v_grp1_3 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_hi_W(vp_hi1), v_scale_duplicated));
296
+
297
+ HVX_Vector v_grp2_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_lo_W(vp_lo2), v_scale_duplicated));
298
+ HVX_Vector v_grp2_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_hi_W(vp_lo2), v_scale_duplicated));
299
+ HVX_Vector v_grp2_2 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_lo_W(vp_hi2), v_scale_duplicated));
300
+ HVX_Vector v_grp2_3 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_hi_W(vp_hi2), v_scale_duplicated));
301
+
302
+ HVX_Vector v_grp3_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_lo_W(vp_lo3), v_scale_duplicated));
303
+ HVX_Vector v_grp3_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_hi_W(vp_lo3), v_scale_duplicated));
304
+ HVX_Vector v_grp3_2 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_lo_W(vp_hi3), v_scale_duplicated));
305
+ HVX_Vector v_grp3_3 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_hi_W(vp_hi3), v_scale_duplicated));
306
+
307
+ hvx_vmem(dst_ptr + 0 * 64) = v_grp0_0;
308
+ hvx_vmem(dst_ptr + 1 * 64) = v_grp0_1;
309
+ hvx_vmem(dst_ptr + 2 * 64) = v_grp0_2;
310
+ hvx_vmem(dst_ptr + 3 * 64) = v_grp0_3;
311
+
312
+ hvx_vmem(dst_ptr + 4 * 64) = v_grp1_0;
313
+ hvx_vmem(dst_ptr + 5 * 64) = v_grp1_1;
314
+ hvx_vmem(dst_ptr + 6 * 64) = v_grp1_2;
315
+ hvx_vmem(dst_ptr + 7 * 64) = v_grp1_3;
316
+
317
+ hvx_vmem(dst_ptr + 8 * 64) = v_grp2_0;
318
+ hvx_vmem(dst_ptr + 9 * 64) = v_grp2_1;
319
+ hvx_vmem(dst_ptr + 10 * 64) = v_grp2_2;
320
+ hvx_vmem(dst_ptr + 11 * 64) = v_grp2_3;
321
+
322
+ hvx_vmem(dst_ptr + 12 * 64) = v_grp3_0;
323
+ hvx_vmem(dst_ptr + 13 * 64) = v_grp3_1;
324
+ hvx_vmem(dst_ptr + 14 * 64) = v_grp3_2;
325
+ hvx_vmem(dst_ptr + 15 * 64) = v_grp3_3;
326
+ }
327
+ }
328
+
329
+ static void dequantize_tiled_weight_to_fp16_task_mxfp4(
330
+ const tiled_dequantize_state_t *state,
331
+ uint32_t start_tile, uint32_t end_tile) {
332
+
333
+ const HVX_Vector mask_h4 = Q6_Vb_vsplat_R(0x0F);
334
+ const HVX_Vector vlut_cvt = hvx_vmem(mxfp4_to_fp16_lut);
335
+
336
+ for (uint32_t t = start_tile; t < end_tile; t++) {
337
+ const uint8_t * tile_src = state->src + t * state->aligned_tile_size;
338
+ __fp16 * dst_ptr = state->dst + t * HTP_MM_HMX_TILE_N_ELMS;
339
+
340
+ HVX_Vector v = hvx_vmem(tile_src + 512);
341
+ HVX_Vector vh = Q6_V_lo_W(Q6_Wuh_vunpack_Vub(v));
342
+ vh = Q6_Vh_vsub_VhVh(vh, Q6_Vh_vsplat_R(112));
343
+ vh = Q6_Vh_vmax_VhVh(vh, Q6_V_vzero());
344
+ vh = Q6_Vh_vmin_VhVh(vh, Q6_Vh_vsplat_R(30));
345
+ vh = Q6_Vh_vasl_VhR(vh, 10);
346
+
347
+ HVX_Vector v_scale_duplicated = Q6_V_lo_W(Q6_W_vshuff_VVR(vh, vh, -2));
348
+
349
+ // Load all 4 groups in parallel
350
+ HVX_Vector vq0 = hvx_vmem(tile_src + 0 * 128);
351
+ HVX_Vector vq1 = hvx_vmem(tile_src + 1 * 128);
352
+ HVX_Vector vq2 = hvx_vmem(tile_src + 2 * 128);
353
+ HVX_Vector vq3 = hvx_vmem(tile_src + 3 * 128);
354
+
355
+ // Nibble extraction
356
+ HVX_Vector v_lo0 = Q6_V_vand_VV(vq0, mask_h4);
357
+ HVX_Vector v_hi0 = Q6_Vub_vlsr_VubR(vq0, 4);
358
+ HVX_Vector v_lo1 = Q6_V_vand_VV(vq1, mask_h4);
359
+ HVX_Vector v_hi1 = Q6_Vub_vlsr_VubR(vq1, 4);
360
+ HVX_Vector v_lo2 = Q6_V_vand_VV(vq2, mask_h4);
361
+ HVX_Vector v_hi2 = Q6_Vub_vlsr_VubR(vq2, 4);
362
+ HVX_Vector v_lo3 = Q6_V_vand_VV(vq3, mask_h4);
363
+ HVX_Vector v_hi3 = Q6_Vub_vlsr_VubR(vq3, 4);
364
+
365
+ // Shuffling
366
+ HVX_VectorPair vp_shuf0 = Q6_W_vshuff_VVR(v_hi0, v_lo0, -1);
367
+ HVX_VectorPair vp_shuf1 = Q6_W_vshuff_VVR(v_hi1, v_lo1, -1);
368
+ HVX_VectorPair vp_shuf2 = Q6_W_vshuff_VVR(v_hi2, v_lo2, -1);
369
+ HVX_VectorPair vp_shuf3 = Q6_W_vshuff_VVR(v_hi3, v_lo3, -1);
370
+
371
+ // Shuffle for LUT lookup
372
+ HVX_Vector v_q_lo0 = Q6_Vb_vshuff_Vb(Q6_V_lo_W(vp_shuf0));
373
+ HVX_Vector v_q_hi0 = Q6_Vb_vshuff_Vb(Q6_V_hi_W(vp_shuf0));
374
+ HVX_Vector v_q_lo1 = Q6_Vb_vshuff_Vb(Q6_V_lo_W(vp_shuf1));
375
+ HVX_Vector v_q_hi1 = Q6_Vb_vshuff_Vb(Q6_V_hi_W(vp_shuf1));
376
+ HVX_Vector v_q_lo2 = Q6_Vb_vshuff_Vb(Q6_V_lo_W(vp_shuf2));
377
+ HVX_Vector v_q_hi2 = Q6_Vb_vshuff_Vb(Q6_V_hi_W(vp_shuf2));
378
+ HVX_Vector v_q_lo3 = Q6_Vb_vshuff_Vb(Q6_V_lo_W(vp_shuf3));
379
+ HVX_Vector v_q_hi3 = Q6_Vb_vshuff_Vb(Q6_V_hi_W(vp_shuf3));
380
+
381
+ // LUT lookup
382
+ HVX_VectorPair vp_lo0 = Q6_Wh_vlut16_VbVhR(v_q_lo0, vlut_cvt, 0);
383
+ HVX_VectorPair vp_hi0 = Q6_Wh_vlut16_VbVhR(v_q_hi0, vlut_cvt, 0);
384
+ HVX_VectorPair vp_lo1 = Q6_Wh_vlut16_VbVhR(v_q_lo1, vlut_cvt, 0);
385
+ HVX_VectorPair vp_hi1 = Q6_Wh_vlut16_VbVhR(v_q_hi1, vlut_cvt, 0);
386
+ HVX_VectorPair vp_lo2 = Q6_Wh_vlut16_VbVhR(v_q_lo2, vlut_cvt, 0);
387
+ HVX_VectorPair vp_hi2 = Q6_Wh_vlut16_VbVhR(v_q_hi2, vlut_cvt, 0);
388
+ HVX_VectorPair vp_lo3 = Q6_Wh_vlut16_VbVhR(v_q_lo3, vlut_cvt, 0);
389
+ HVX_VectorPair vp_hi3 = Q6_Wh_vlut16_VbVhR(v_q_hi3, vlut_cvt, 0);
390
+
391
+ // Convert and scale multiplication
392
+ HVX_Vector v_grp0_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_lo_W(vp_lo0), v_scale_duplicated));
393
+ HVX_Vector v_grp0_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_hi_W(vp_lo0), v_scale_duplicated));
394
+ HVX_Vector v_grp0_2 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_lo_W(vp_hi0), v_scale_duplicated));
395
+ HVX_Vector v_grp0_3 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_hi_W(vp_hi0), v_scale_duplicated));
396
+
397
+ HVX_Vector v_grp1_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_lo_W(vp_lo1), v_scale_duplicated));
398
+ HVX_Vector v_grp1_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_hi_W(vp_lo1), v_scale_duplicated));
399
+ HVX_Vector v_grp1_2 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_lo_W(vp_hi1), v_scale_duplicated));
400
+ HVX_Vector v_grp1_3 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_hi_W(vp_hi1), v_scale_duplicated));
401
+
402
+ HVX_Vector v_grp2_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_lo_W(vp_lo2), v_scale_duplicated));
403
+ HVX_Vector v_grp2_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_hi_W(vp_lo2), v_scale_duplicated));
404
+ HVX_Vector v_grp2_2 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_lo_W(vp_hi2), v_scale_duplicated));
405
+ HVX_Vector v_grp2_3 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_hi_W(vp_hi2), v_scale_duplicated));
406
+
407
+ HVX_Vector v_grp3_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_lo_W(vp_lo3), v_scale_duplicated));
408
+ HVX_Vector v_grp3_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_hi_W(vp_lo3), v_scale_duplicated));
409
+ HVX_Vector v_grp3_2 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_lo_W(vp_hi3), v_scale_duplicated));
410
+ HVX_Vector v_grp3_3 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_V_hi_W(vp_hi3), v_scale_duplicated));
411
+
412
+ hvx_vmem(dst_ptr + 0 * 64) = v_grp0_0;
413
+ hvx_vmem(dst_ptr + 1 * 64) = v_grp0_1;
414
+ hvx_vmem(dst_ptr + 2 * 64) = v_grp0_2;
415
+ hvx_vmem(dst_ptr + 3 * 64) = v_grp0_3;
416
+
417
+ hvx_vmem(dst_ptr + 4 * 64) = v_grp1_0;
418
+ hvx_vmem(dst_ptr + 5 * 64) = v_grp1_1;
419
+ hvx_vmem(dst_ptr + 6 * 64) = v_grp1_2;
420
+ hvx_vmem(dst_ptr + 7 * 64) = v_grp1_3;
421
+
422
+ hvx_vmem(dst_ptr + 8 * 64) = v_grp2_0;
423
+ hvx_vmem(dst_ptr + 9 * 64) = v_grp2_1;
424
+ hvx_vmem(dst_ptr + 10 * 64) = v_grp2_2;
425
+ hvx_vmem(dst_ptr + 11 * 64) = v_grp2_3;
426
+
427
+ hvx_vmem(dst_ptr + 12 * 64) = v_grp3_0;
428
+ hvx_vmem(dst_ptr + 13 * 64) = v_grp3_1;
429
+ hvx_vmem(dst_ptr + 14 * 64) = v_grp3_2;
430
+ hvx_vmem(dst_ptr + 15 * 64) = v_grp3_3;
431
+ }
432
+ }
433
+
434
+ static void dequantize_tiled_weight_to_fp16_task_q8_0(
435
+ const tiled_dequantize_state_t *state,
436
+ uint32_t start_tile, uint32_t end_tile) {
437
+
438
+ for (uint32_t t = start_tile; t < end_tile; t++) {
439
+ const uint8_t * tile_src = state->src + t * state->aligned_tile_size;
440
+ __fp16 * dst_ptr = state->dst + t * HTP_MM_HMX_TILE_N_ELMS;
441
+
442
+ HVX_Vector v_sc = hvx_vmem(tile_src + 1024);
443
+ HVX_Vector v_scale_duplicated = Q6_V_lo_W(Q6_W_vshuff_VVR(v_sc, v_sc, -2));
444
+
445
+ // Load groups 0-3 in parallel
446
+ HVX_Vector vq0 = hvx_vmem(tile_src + 0 * 128);
447
+ HVX_Vector vq1 = hvx_vmem(tile_src + 1 * 128);
448
+ HVX_Vector vq2 = hvx_vmem(tile_src + 2 * 128);
449
+ HVX_Vector vq3 = hvx_vmem(tile_src + 3 * 128);
450
+
451
+ HVX_VectorPair vp_int16_0 = Q6_Wh_vunpack_Vb(vq0);
452
+ HVX_VectorPair vp_int16_1 = Q6_Wh_vunpack_Vb(vq1);
453
+ HVX_VectorPair vp_int16_2 = Q6_Wh_vunpack_Vb(vq2);
454
+ HVX_VectorPair vp_int16_3 = Q6_Wh_vunpack_Vb(vq3);
455
+
456
+ // Load groups 4-7 in parallel
457
+ HVX_Vector vq4 = hvx_vmem(tile_src + 4 * 128);
458
+ HVX_Vector vq5 = hvx_vmem(tile_src + 5 * 128);
459
+ HVX_Vector vq6 = hvx_vmem(tile_src + 6 * 128);
460
+ HVX_Vector vq7 = hvx_vmem(tile_src + 7 * 128);
461
+
462
+ HVX_VectorPair vp_int16_4 = Q6_Wh_vunpack_Vb(vq4);
463
+ HVX_VectorPair vp_int16_5 = Q6_Wh_vunpack_Vb(vq5);
464
+ HVX_VectorPair vp_int16_6 = Q6_Wh_vunpack_Vb(vq6);
465
+ HVX_VectorPair vp_int16_7 = Q6_Wh_vunpack_Vb(vq7);
466
+
467
+ // Convert and scale multiply for groups 0-3
468
+ HVX_Vector v_grp0_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_0)), v_scale_duplicated));
469
+ HVX_Vector v_grp0_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_0)), v_scale_duplicated));
470
+ HVX_Vector v_grp1_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_1)), v_scale_duplicated));
471
+ HVX_Vector v_grp1_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_1)), v_scale_duplicated));
472
+ HVX_Vector v_grp2_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_2)), v_scale_duplicated));
473
+ HVX_Vector v_grp2_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_2)), v_scale_duplicated));
474
+ HVX_Vector v_grp3_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_3)), v_scale_duplicated));
475
+ HVX_Vector v_grp3_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_3)), v_scale_duplicated));
476
+
477
+ // Store groups 0-3
478
+ hvx_vmem(dst_ptr + 0 * 64) = v_grp0_0;
479
+ hvx_vmem(dst_ptr + 1 * 64) = v_grp0_1;
480
+ hvx_vmem(dst_ptr + 2 * 64) = v_grp1_0;
481
+ hvx_vmem(dst_ptr + 3 * 64) = v_grp1_1;
482
+ hvx_vmem(dst_ptr + 4 * 64) = v_grp2_0;
483
+ hvx_vmem(dst_ptr + 5 * 64) = v_grp2_1;
484
+ hvx_vmem(dst_ptr + 6 * 64) = v_grp3_0;
485
+ hvx_vmem(dst_ptr + 7 * 64) = v_grp3_1;
486
+
487
+ // Convert and scale multiply for groups 4-7
488
+ HVX_Vector v_grp4_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_4)), v_scale_duplicated));
489
+ HVX_Vector v_grp4_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_4)), v_scale_duplicated));
490
+ HVX_Vector v_grp5_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_5)), v_scale_duplicated));
491
+ HVX_Vector v_grp5_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_5)), v_scale_duplicated));
492
+ HVX_Vector v_grp6_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_6)), v_scale_duplicated));
493
+ HVX_Vector v_grp6_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_6)), v_scale_duplicated));
494
+ HVX_Vector v_grp7_0 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_lo_W(vp_int16_7)), v_scale_duplicated));
495
+ HVX_Vector v_grp7_1 = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(Q6_Vhf_equals_Vh(Q6_V_hi_W(vp_int16_7)), v_scale_duplicated));
496
+
497
+ // Store groups 4-7
498
+ hvx_vmem(dst_ptr + 8 * 64) = v_grp4_0;
499
+ hvx_vmem(dst_ptr + 9 * 64) = v_grp4_1;
500
+ hvx_vmem(dst_ptr + 10 * 64) = v_grp5_0;
501
+ hvx_vmem(dst_ptr + 11 * 64) = v_grp5_1;
502
+ hvx_vmem(dst_ptr + 12 * 64) = v_grp6_0;
503
+ hvx_vmem(dst_ptr + 13 * 64) = v_grp6_1;
504
+ hvx_vmem(dst_ptr + 14 * 64) = v_grp7_0;
505
+ hvx_vmem(dst_ptr + 15 * 64) = v_grp7_1;
506
+ }
507
+ }
508
+
509
+ static __attribute__((noinline))
510
+ void convert_f16_weight_to_fp16_tiles_task(
511
+ const tiled_dequantize_state_t *state,
512
+ uint32_t start_tile, uint32_t end_tile) {
513
+
514
+ const uint32_t n_k_tiles = state->n_k_tiles;
515
+ const struct fastdiv_values n_k_tiles_div = state->n_k_tiles_div;
516
+
517
+ const HVX_Vector v_scat_base = hvx_vmem(hmx_transpose_scatter_offsets);
518
+ const HVX_Vector v_scat_step = Q6_V_vsplat_R(4);
519
+ const HVX_VectorPred q_mask64 = Q6_Q_vsetq_R(64);
520
+
521
+ unsigned ct = fastdiv((unsigned)start_tile, &n_k_tiles_div);
522
+ unsigned kt = fastmodulo((unsigned)start_tile, n_k_tiles, &n_k_tiles_div);
523
+
524
+ for (unsigned t = start_tile; t < (unsigned)end_tile; ) {
525
+ if (kt >= (unsigned)n_k_tiles) { kt = 0; ct++; }
526
+
527
+ __fp16 *tile_base = state->dst + t * HTP_MM_HMX_TILE_N_ELMS;
528
+ {
529
+ uint32_t byte_off = kt * 32 * sizeof(__fp16);
530
+
531
+ HVX_Vector v_off = v_scat_base;
532
+ for (uint32_t r = 0; r < HTP_MM_HMX_TILE_N_ROWS; r += 2) {
533
+ uint32_t row0 = ct * HTP_MM_HMX_TILE_N_COLS + r;
534
+ uint32_t row1 = row0 + 1;
535
+
536
+ const uint8_t *r0 = state->src + row0 * state->row_stride;
537
+ const uint8_t *r1 = state->src + row1 * state->row_stride;
538
+
539
+ HVX_Vector v0 = hvx_vmemu((const __fp16 *)(r0 + byte_off));
540
+ HVX_Vector v1 = (row1 < state->n_cols) ? hvx_vmemu((const __fp16 *)(r1 + byte_off)) : Q6_V_vzero();
541
+
542
+ Q6_vscatter_QRMVwV(q_mask64, (size_t)tile_base, HTP_MM_HMX_TILE_SIZE - 1, v_off, v0);
543
+ v_off = Q6_Vw_vadd_VwVw(v_off, v_scat_step);
544
+ Q6_vscatter_QRMVwV(q_mask64, (size_t)tile_base, HTP_MM_HMX_TILE_SIZE - 1, v_off, v1);
545
+ v_off = Q6_Vw_vadd_VwVw(v_off, v_scat_step);
546
+ }
547
+ }
548
+ ++t; ++kt;
549
+ }
550
+ }
551
+
552
+ static __attribute__((noinline))
553
+ void quantize_f32_weight_to_fp16_tiles_task(
554
+ const tiled_dequantize_state_t *state,
555
+ uint32_t start_tile, uint32_t end_tile) {
556
+
557
+ const uint32_t n_k_tiles = state->n_k_tiles;
558
+ const struct fastdiv_values n_k_tiles_div = state->n_k_tiles_div;
559
+
560
+ const HVX_Vector v_scat_base = hvx_vmem(hmx_transpose_scatter_offsets);
561
+ const HVX_Vector v_scat_step = Q6_V_vsplat_R(4);
562
+ const HVX_VectorPred q_mask64 = Q6_Q_vsetq_R(64);
563
+
564
+ unsigned ct = fastdiv((unsigned)start_tile, &n_k_tiles_div);
565
+ unsigned kt = fastmodulo((unsigned)start_tile, n_k_tiles, &n_k_tiles_div);
566
+
567
+ for (unsigned t = start_tile; t < (unsigned)end_tile; ) {
568
+ if (kt >= (unsigned)n_k_tiles) { kt = 0; ct++; }
569
+
570
+ __fp16 *tile_base = state->dst + t * HTP_MM_HMX_TILE_N_ELMS;
571
+ {
572
+ uint32_t byte_off = kt * 32 * sizeof(float);
573
+
574
+ HVX_Vector v_off = v_scat_base;
575
+ for (uint32_t r = 0; r < HTP_MM_HMX_TILE_N_ROWS; r += 2) {
576
+ uint32_t row0 = ct * HTP_MM_HMX_TILE_N_COLS + r;
577
+ uint32_t row1 = row0 + 1;
578
+
579
+ const uint8_t *r0 = state->src + row0 * state->row_stride;
580
+ const uint8_t *r1 = state->src + row1 * state->row_stride;
581
+
582
+ HVX_Vector v0_f32 = hvx_vmem((const float *)(r0 + byte_off));
583
+ HVX_Vector v1_f32 = (row1 < state->n_cols) ? hvx_vmem((const float *)(r1 + byte_off)) : Q6_V_vzero();
584
+
585
+ HVX_Vector v_out = hvx_vec_f32_to_f16(v0_f32, v1_f32);
586
+
587
+ Q6_vscatter_QRMVwV(q_mask64, (size_t)tile_base, HTP_MM_HMX_TILE_SIZE - 1, v_off, v_out);
588
+ v_off = Q6_Vw_vadd_VwVw(v_off, v_scat_step);
589
+
590
+ HVX_Vector v_out_hi = Q6_V_vror_VR(v_out, 64);
591
+ Q6_vscatter_QRMVwV(q_mask64, (size_t)tile_base, HTP_MM_HMX_TILE_SIZE - 1, v_off, v_out_hi);
592
+ v_off = Q6_Vw_vadd_VwVw(v_off, v_scat_step);
593
+ }
594
+ }
595
+ ++t; ++kt;
596
+ }
597
+ }
598
+
599
+ // --- End tiled dequantizers ---
600
+
601
+ // dot-chunk functions require external HMX lock
602
+
603
+ static void core_dot_chunk_fp16_short(__fp16 *restrict output, const __fp16 *restrict activation,
604
+ const __fp16 *restrict weight, const __fp16 *restrict scales,
605
+ uint32_t n_row_tiles, uint32_t n_col_tiles, uint32_t n_dot_tiles) {
606
+ __builtin_assume(n_row_tiles > 0);
607
+ __builtin_assume(n_col_tiles > 0);
608
+ __builtin_assume(n_dot_tiles > 0);
609
+ __builtin_assume(n_dot_tiles <= 32);
610
+
611
+ asm volatile(HMX_SET_BIAS("%0") :: "r"((unsigned int)scales));
612
+
613
+ const size_t dot_stride = n_dot_tiles * HTP_MM_HMX_TILE_N_ELMS;
614
+ const uint32_t range = 2048u * n_dot_tiles - 1;
615
+
616
+ for (uint32_t r = 0; r < n_row_tiles; ++r) {
617
+ const __fp16 *row_base = activation + r * dot_stride;
618
+ const __fp16 *col_base = weight;
619
+ __fp16 *out_tile = output + r * n_col_tiles * HTP_MM_HMX_TILE_N_ELMS;
620
+
621
+ for (size_t c = 0; c < n_col_tiles; ++c) {
622
+ asm volatile(HMX_CLRACC_F16());
623
+ asm volatile(HMX_LOAD_MPY_DEEP_F16("%1", "%2", "%0") : : "r"(range), "r"(row_base), "r"(col_base));
624
+ asm volatile(HMX_STORE_AFTER_F16("%0", "%1") : : "r"(out_tile), "r"(0) : "memory");
625
+ col_base += dot_stride;
626
+ out_tile += HTP_MM_HMX_TILE_N_ELMS;
627
+ }
628
+ }
629
+ }
630
+
631
+ static void core_dot_chunk_fp16(__fp16 *restrict output, const __fp16 *restrict activation,
632
+ const __fp16 *restrict weight, const __fp16 *restrict scales,
633
+ uint32_t n_row_tiles, uint32_t n_col_tiles, uint32_t n_dot_tiles) {
634
+ if (n_dot_tiles <= 32) {
635
+ core_dot_chunk_fp16_short(output, activation, weight, scales, n_row_tiles, n_col_tiles, n_dot_tiles);
636
+ return;
637
+ }
638
+ __builtin_assume(n_row_tiles > 0);
639
+ __builtin_assume(n_col_tiles > 0);
640
+ __builtin_assume(n_dot_tiles > 32);
641
+
642
+ asm volatile(HMX_SET_BIAS("%0") :: "r"((unsigned int)scales));
643
+
644
+ const size_t dot_stride = n_dot_tiles * HTP_MM_HMX_TILE_N_ELMS;
645
+
646
+ for (uint32_t r = 0; r < n_row_tiles; ++r) {
647
+ const __fp16 *row_base = activation + r * dot_stride;
648
+ const __fp16 *col_base = weight;
649
+ __fp16 *out_tile = output + r * n_col_tiles * HTP_MM_HMX_TILE_N_ELMS;
650
+
651
+ for (size_t c = 0; c < n_col_tiles; ++c) {
652
+ const __fp16 *row_tiles = row_base;
653
+ const __fp16 *col_tiles = col_base;
654
+
655
+ asm volatile(HMX_CLRACC_F16());
656
+
657
+ const uint32_t n_loops = n_dot_tiles / 32;
658
+ const uint32_t rem = n_dot_tiles % 32;
659
+
660
+ for (uint32_t l = 0; l < n_loops; ++l) {
661
+ asm volatile(HMX_LOAD_MPY_DEEP_F16("%1", "%2", "%0") : : "r"(65535), "r"(row_tiles), "r"(col_tiles));
662
+ row_tiles += 32 * HTP_MM_HMX_TILE_N_ELMS;
663
+ col_tiles += 32 * HTP_MM_HMX_TILE_N_ELMS;
664
+ }
665
+
666
+ if (rem > 0) {
667
+ const uint32_t range = 2048u * rem - 1;
668
+ asm volatile(HMX_LOAD_MPY_DEEP_F16("%1", "%2", "%0") : : "r"(range), "r"(row_tiles), "r"(col_tiles));
669
+ }
670
+
671
+ asm volatile(HMX_STORE_AFTER_F16("%0", "%1") : : "r"(out_tile), "r"(0) : "memory");
672
+
673
+ col_base += dot_stride;
674
+ out_tile += HTP_MM_HMX_TILE_N_ELMS;
675
+ }
676
+ }
677
+ }
678
+
679
+ static void core_mma_chunk_fp16_short(__fp16 *restrict c, const __fp16 *restrict a, const __fp16 *restrict b,
680
+ const __fp16 *restrict col_scales, const __fp16 *restrict eye_tile,
681
+ uint32_t n_row_tiles, uint32_t n_col_tiles, uint32_t n_dot_tiles, bool zero_init) {
682
+ __builtin_assume(n_row_tiles > 0);
683
+ __builtin_assume(n_col_tiles > 0);
684
+ __builtin_assume(n_dot_tiles > 0);
685
+ __builtin_assume(n_dot_tiles <= 32);
686
+
687
+ asm volatile(HMX_SET_BIAS("%0") :: "r"((unsigned int)col_scales));
688
+
689
+ const size_t dot_tile_stride = n_dot_tiles * HTP_MM_HMX_TILE_N_ELMS;
690
+ const uint32_t range = 2048u * n_dot_tiles - 1;
691
+
692
+ for (size_t i = 0; i < n_row_tiles; ++i) {
693
+ const __fp16 *row_base = a + i * dot_tile_stride;
694
+ __fp16 *res_base = c + i * n_col_tiles * HTP_MM_HMX_TILE_N_ELMS;
695
+ const __fp16 *col_base = b;
696
+ __fp16 *accum_tile = res_base;
697
+
698
+ for (size_t j = 0; j < n_col_tiles; ++j) {
699
+ asm volatile(HMX_CLRACC_F16());
700
+
701
+ if (!zero_init) {
702
+ asm volatile(HMX_LOAD_MPY_F16("%1", "%2", "%0") : : "r"(2047), "r"(accum_tile), "r"(eye_tile));
703
+ }
704
+
705
+ asm volatile(HMX_LOAD_MPY_DEEP_F16("%1", "%2", "%0") : : "r"(range), "r"(row_base), "r"(col_base));
706
+
707
+ asm volatile(HMX_STORE_AFTER_F16("%0", "%1") : : "r"(accum_tile), "r"(0) : "memory");
708
+
709
+ col_base += dot_tile_stride;
710
+ accum_tile += HTP_MM_HMX_TILE_N_ELMS;
711
+ }
712
+ }
713
+ }
714
+
715
+ static void core_mma_chunk_fp16(__fp16 *restrict c, const __fp16 *restrict a, const __fp16 *restrict b,
716
+ const __fp16 *restrict col_scales, const __fp16 *restrict eye_tile,
717
+ uint32_t n_row_tiles, uint32_t n_col_tiles, uint32_t n_dot_tiles, bool zero_init) {
718
+ if (n_dot_tiles <= 32) {
719
+ core_mma_chunk_fp16_short(c, a, b, col_scales, eye_tile, n_row_tiles, n_col_tiles, n_dot_tiles, zero_init);
720
+ return;
721
+ }
722
+ __builtin_assume(n_row_tiles > 0);
723
+ __builtin_assume(n_col_tiles > 0);
724
+ __builtin_assume(n_dot_tiles > 32);
725
+
726
+ asm volatile(HMX_SET_BIAS("%0") :: "r"((unsigned int)col_scales));
727
+
728
+ const size_t dot_tile_stride = n_dot_tiles * HTP_MM_HMX_TILE_N_ELMS;
729
+
730
+ for (size_t i = 0; i < n_row_tiles; ++i) {
731
+ const __fp16 *row_base = a + i * dot_tile_stride;
732
+ __fp16 *res_base = c + i * n_col_tiles * HTP_MM_HMX_TILE_N_ELMS;
733
+ const __fp16 *col_base = b;
734
+ __fp16 *accum_tile = res_base;
735
+
736
+ for (size_t j = 0; j < n_col_tiles; ++j) {
737
+ const __fp16 *col_tiles = col_base;
738
+ const __fp16 *row_tiles = row_base;
739
+
740
+ asm volatile(HMX_CLRACC_F16());
741
+
742
+ if (!zero_init) {
743
+ asm volatile(HMX_LOAD_MPY_F16("%1", "%2", "%0") : : "r"(2047), "r"(accum_tile), "r"(eye_tile));
744
+ }
745
+
746
+ const uint32_t n_loops = n_dot_tiles / 32;
747
+ const uint32_t rem = n_dot_tiles % 32;
748
+
749
+ for (uint32_t l = 0; l < n_loops; ++l) {
750
+ asm volatile(HMX_LOAD_MPY_DEEP_F16("%1", "%2", "%0") : : "r"(65535), "r"(row_tiles), "r"(col_tiles));
751
+ row_tiles += 32 * HTP_MM_HMX_TILE_N_ELMS;
752
+ col_tiles += 32 * HTP_MM_HMX_TILE_N_ELMS;
753
+ }
754
+
755
+ if (rem > 0) {
756
+ const uint32_t range = 2048u * rem - 1;
757
+ asm volatile(HMX_LOAD_MPY_DEEP_F16("%1", "%2", "%0") : : "r"(range), "r"(row_tiles), "r"(col_tiles));
758
+ }
759
+
760
+ asm volatile(HMX_STORE_AFTER_F16("%0", "%1") : : "r"(accum_tile), "r"(0) : "memory");
761
+
762
+ col_base += dot_tile_stride;
763
+ accum_tile += HTP_MM_HMX_TILE_N_ELMS;
764
+ }
765
+ }
766
+ }
767
+
768
+ // output : fp16 -> f32p
769
+
770
+ static void transfer_output_chunk_fp16_to_fp32(
771
+ float *restrict dst,
772
+ const float *restrict src2,
773
+ const __fp16 *restrict vtcm_src,
774
+ uint32_t start_row,
775
+ uint32_t n_rows,
776
+ uint32_t n_cols,
777
+ uint32_t dst_stride,
778
+ uint32_t src2_stride,
779
+ uint32_t dst_cols
780
+ ) {
781
+ assert(n_cols % HTP_MM_HMX_TILE_N_COLS == 0);
782
+ const size_t tile_row_stride = (n_cols / HTP_MM_HMX_TILE_N_COLS) * HTP_MM_HMX_TILE_N_ELMS;
783
+
784
+ const HVX_Vector one = hvx_vec_splat_f16(1.0);
785
+
786
+ const size_t limit_c = hex_smin(n_cols, dst_cols);
787
+ const size_t limit_c_aligned = (limit_c & ~31);
788
+
789
+ for (size_t r = 0; r < n_rows; r += 2) {
790
+ const size_t r_idx0 = start_row + r + 0;
791
+ const size_t r0 = r_idx0 / HTP_MM_HMX_TILE_N_ROWS;
792
+ const size_t r1 = (r_idx0 % HTP_MM_HMX_TILE_N_ROWS) / 2; // index of the row pair within the tile
793
+ const __fp16 *row_base = vtcm_src + r0 * tile_row_stride;
794
+ float *output_row_base = dst + r * dst_stride; // global memory row base for row r (and r+1)
795
+ const float *src2_row_base = src2 ? (src2 + r * src2_stride) : NULL;
796
+
797
+ #pragma unroll(4)
798
+ for (size_t c = 0; c < limit_c_aligned; c += HTP_MM_HMX_TILE_N_COLS) {
799
+ const size_t c0 = c / HTP_MM_HMX_TILE_N_COLS;
800
+ const __fp16 *tile = row_base + c0 * HTP_MM_HMX_TILE_N_ELMS;
801
+ HVX_Vector v = ((const HVX_Vector *) tile)[r1];
802
+ HVX_VectorPair vp = Q6_Wqf32_vmpy_VhfVhf(v, one);
803
+
804
+ HVX_Vector *pv_out0 = (HVX_Vector *) (output_row_base + c + 0);
805
+ HVX_Vector *pv_out1 = (HVX_Vector *) (output_row_base + c + dst_stride);
806
+
807
+ HVX_Vector v_out0 = Q6_Vsf_equals_Vqf32(Q6_V_lo_W(vp));
808
+ if (src2_row_base) {
809
+ HVX_Vector v_src2_0 = hvx_vmemu(src2_row_base + c + 0);
810
+ v_out0 = hvx_vec_add_f32_f32(v_out0, v_src2_0);
811
+ }
812
+ *pv_out0 = v_out0;
813
+
814
+ if (r + 1 < n_rows) {
815
+ HVX_Vector v_out1 = Q6_Vsf_equals_Vqf32(Q6_V_hi_W(vp));
816
+ if (src2_row_base) {
817
+ HVX_Vector v_src2_1 = hvx_vmemu(src2_row_base + c + src2_stride);
818
+ v_out1 = hvx_vec_add_f32_f32(v_out1, v_src2_1);
819
+ }
820
+ *pv_out1 = v_out1;
821
+ }
822
+ }
823
+
824
+ if (limit_c_aligned < limit_c) {
825
+ size_t c = limit_c_aligned;
826
+ size_t valid_c = limit_c - c;
827
+ const size_t c0 = c / HTP_MM_HMX_TILE_N_COLS;
828
+ const __fp16 *tile = row_base + c0 * HTP_MM_HMX_TILE_N_ELMS;
829
+ HVX_Vector v = ((const HVX_Vector *) tile)[r1];
830
+ HVX_VectorPair vp = Q6_Wqf32_vmpy_VhfVhf(v, one);
831
+
832
+ HVX_Vector v_out0 = Q6_Vsf_equals_Vqf32(Q6_V_lo_W(vp));
833
+ if (src2_row_base) {
834
+ HVX_Vector v_src2_0 = hvx_vmemu(src2_row_base + c + 0);
835
+ v_out0 = hvx_vec_add_f32_f32(v_out0, v_src2_0);
836
+ }
837
+ hvx_vec_store_u(output_row_base + c, valid_c * sizeof(float), v_out0);
838
+
839
+ if (r + 1 < n_rows) {
840
+ HVX_Vector v_out1 = Q6_Vsf_equals_Vqf32(Q6_V_hi_W(vp));
841
+ if (src2_row_base) {
842
+ HVX_Vector v_src2_1 = hvx_vmemu(src2_row_base + c + src2_stride);
843
+ v_out1 = hvx_vec_add_f32_f32(v_out1, v_src2_1);
844
+ }
845
+ hvx_vec_store_u(output_row_base + c + dst_stride, valid_c * sizeof(float), v_out1);
846
+ }
847
+ }
848
+ }
849
+ }
850
+
851
+ typedef struct {
852
+ const __fp16 *vtcm_src;
853
+ float *dst;
854
+ const float *src2;
855
+ uint32_t n_tasks;
856
+ uint32_t n_tot_chunks;
857
+ uint32_t n_chunks_per_task;
858
+ uint32_t n_cols;
859
+ uint32_t dst_stride; // DDR row stride
860
+ uint32_t src2_stride; // DDR row stride for residual
861
+ uint32_t dst_cols; // Actual output columns
862
+ struct htp_thread_trace * traces;
863
+ } output_transfer_task_state_t;
864
+
865
+ // activations : fp32 -> fp16
866
+
867
+ static void transfer_activation_chunk_fp32_to_fp16(__fp16 *restrict vtcm_dst, const float *restrict src, uint32_t n_rows, uint32_t k_block, uint32_t k_stride, uint32_t k_valid) {
868
+ const uint32_t n_rows_padded = hex_align_up(n_rows, HTP_MM_HMX_TILE_N_ROWS);
869
+ const uint32_t n_rows_tiled = (n_rows / HTP_MM_HMX_TILE_N_ROWS) * HTP_MM_HMX_TILE_N_ROWS;
870
+
871
+ uint32_t r = 0;
872
+
873
+ #pragma unroll(2)
874
+ for (r = 0; r < n_rows_tiled; r += 2) {
875
+ uint32_t r0 = r / HTP_MM_HMX_TILE_N_ROWS; // tile row index
876
+ uint32_t r1 = r % HTP_MM_HMX_TILE_N_ROWS; // intra-tile row idx
877
+
878
+ const float *ptr_in0 = src + (r + 0) * k_stride;
879
+ const float *ptr_in1 = src + (r + 1) * k_stride;
880
+
881
+ uint32_t c = 0;
882
+ for (; c + 32 <= k_valid; c += 32) {
883
+ HVX_Vector v0 = *(const HVX_Vector *)(ptr_in0 + c);
884
+ HVX_Vector v1 = *(const HVX_Vector *)(ptr_in1 + c);
885
+ HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
886
+
887
+ uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS; // tile column index
888
+ uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
889
+
890
+ HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
891
+ tile[r1 / 2] = v_out;
892
+ }
893
+ if (c < k_block) {
894
+ HVX_Vector v0 = *(const HVX_Vector *)(ptr_in0 + c);
895
+ HVX_Vector v1 = *(const HVX_Vector *)(ptr_in1 + c);
896
+
897
+ uint32_t rem = k_valid - c;
898
+ HVX_VectorPred mask = Q6_Q_vsetq2_R(rem > 0 ? rem * sizeof(float) : 0);
899
+ v0 = Q6_V_vmux_QVV(mask, v0, Q6_V_vzero());
900
+ v1 = Q6_V_vmux_QVV(mask, v1, Q6_V_vzero());
901
+
902
+ HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
903
+
904
+ uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS; // tile column index
905
+ uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
906
+
907
+ HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
908
+ tile[r1 / 2] = v_out;
909
+ }
910
+ }
911
+
912
+ for (; r < n_rows_padded; r += 2) {
913
+ uint32_t r0 = r / HTP_MM_HMX_TILE_N_ROWS; // tile row index
914
+ uint32_t r1 = r % HTP_MM_HMX_TILE_N_ROWS; // intra-tile row idx
915
+
916
+ const bool row0_valid = r < n_rows;
917
+ const bool row1_valid = (r + 1) < n_rows;
918
+
919
+ const float *ptr_in0 = row0_valid ? (src + (r + 0) * k_stride) : NULL;
920
+ const float *ptr_in1 = row1_valid ? (src + (r + 1) * k_stride) : NULL;
921
+
922
+ uint32_t c = 0;
923
+ for (; c + 32 <= k_valid; c += 32) {
924
+ HVX_Vector v0 = Q6_V_vzero();
925
+ HVX_Vector v1 = Q6_V_vzero();
926
+ if (row0_valid) v0 = *(const HVX_Vector *)(ptr_in0 + c);
927
+ if (row1_valid) v1 = *(const HVX_Vector *)(ptr_in1 + c);
928
+
929
+ HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
930
+
931
+ uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS; // tile column index
932
+ uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
933
+
934
+ HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
935
+ tile[r1 / 2] = v_out;
936
+ }
937
+ if (c < k_block) {
938
+ HVX_Vector v0 = Q6_V_vzero();
939
+ HVX_Vector v1 = Q6_V_vzero();
940
+ if (row0_valid) v0 = *(const HVX_Vector *)(ptr_in0 + c);
941
+ if (row1_valid) v1 = *(const HVX_Vector *)(ptr_in1 + c);
942
+
943
+ uint32_t rem = k_valid - c;
944
+ HVX_VectorPred mask = Q6_Q_vsetq2_R(rem > 0 ? rem * sizeof(float) : 0);
945
+ v0 = Q6_V_vmux_QVV(mask, v0, Q6_V_vzero());
946
+ v1 = Q6_V_vmux_QVV(mask, v1, Q6_V_vzero());
947
+
948
+ HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
949
+
950
+ uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS; // tile column index
951
+ uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
952
+
953
+ HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
954
+ tile[r1 / 2] = v_out;
955
+ }
956
+ }
957
+ }
958
+
959
+ static void transfer_activation_row_pair_fp32_to_fp16(
960
+ __fp16 *restrict vtcm_dst,
961
+ const float *restrict row0,
962
+ const float *restrict row1,
963
+ uint32_t r,
964
+ uint32_t k_block,
965
+ uint32_t k_valid,
966
+ bool row0_valid,
967
+ bool row1_valid) {
968
+
969
+ uint32_t r0 = r / HTP_MM_HMX_TILE_N_ROWS; // tile row index
970
+ uint32_t r1 = r % HTP_MM_HMX_TILE_N_ROWS; // intra-tile row idx
971
+
972
+ uint32_t c = 0;
973
+ for (; c + 32 <= k_valid; c += 32) {
974
+ HVX_Vector v0 = Q6_V_vzero();
975
+ HVX_Vector v1 = Q6_V_vzero();
976
+ if (row0_valid) v0 = *(const HVX_Vector *)(row0 + c);
977
+ if (row1_valid) v1 = *(const HVX_Vector *)(row1 + c);
978
+
979
+ HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
980
+
981
+ uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS; // tile column index
982
+ uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
983
+
984
+ HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
985
+ tile[r1 / 2] = v_out;
986
+ }
987
+ if (c < k_block) {
988
+ HVX_Vector v0 = Q6_V_vzero();
989
+ HVX_Vector v1 = Q6_V_vzero();
990
+ if (row0_valid) v0 = *(const HVX_Vector *)(row0 + c);
991
+ if (row1_valid) v1 = *(const HVX_Vector *)(row1 + c);
992
+
993
+ uint32_t rem = k_valid - c;
994
+ HVX_VectorPred mask = Q6_Q_vsetq2_R(rem > 0 ? rem * sizeof(float) : 0);
995
+ v0 = Q6_V_vmux_QVV(mask, v0, Q6_V_vzero());
996
+ v1 = Q6_V_vmux_QVV(mask, v1, Q6_V_vzero());
997
+
998
+ HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
999
+
1000
+ uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS; // tile column index
1001
+ uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
1002
+
1003
+ HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
1004
+ tile[r1 / 2] = v_out;
1005
+ }
1006
+ }
1007
+
1008
+ static void transfer_activation_chunk_fp32_to_fp16_gathered(
1009
+ __fp16 *restrict vtcm_dst,
1010
+ const float *restrict src,
1011
+ uint32_t start_row,
1012
+ uint32_t n_rows,
1013
+ uint32_t k_block,
1014
+ const struct mmid_row_mapping *matrix_rows,
1015
+ uint32_t cur_a,
1016
+ uint32_t mapping_stride,
1017
+ uint32_t ne11,
1018
+ const struct fastdiv_values * ne11_div,
1019
+ size_t nb11,
1020
+ size_t nb12,
1021
+ uint32_t cne1,
1022
+ uint32_t k_valid) {
1023
+ const uint32_t n_rows_padded = hex_align_up(n_rows, HTP_MM_HMX_TILE_N_ROWS);
1024
+ const uint32_t n_rows_tiled = (n_rows / HTP_MM_HMX_TILE_N_ROWS) * HTP_MM_HMX_TILE_N_ROWS;
1025
+
1026
+ uint32_t r = 0;
1027
+
1028
+ #pragma unroll(2)
1029
+ for (r = 0; r < n_rows_tiled; r += 2) {
1030
+ uint32_t r_idx0 = start_row + r + 0;
1031
+ uint32_t r_idx1 = start_row + r + 1;
1032
+ uint32_t r0 = r_idx0 / HTP_MM_HMX_TILE_N_ROWS; // tile row index
1033
+ uint32_t r1 = r_idx0 % HTP_MM_HMX_TILE_N_ROWS; // intra-tile row idx
1034
+
1035
+ struct mmid_row_mapping mapping0 = matrix_rows[cur_a * mapping_stride + r_idx0];
1036
+ struct mmid_row_mapping mapping1 = matrix_rows[cur_a * mapping_stride + r_idx1];
1037
+
1038
+ uint32_t i11_0 = fastmodulo(mapping0.i1, ne11, ne11_div);
1039
+ uint32_t i11_1 = fastmodulo(mapping1.i1, ne11, ne11_div);
1040
+
1041
+ const float *row0_ptr = (const float *) ((const uint8_t *) src + i11_0 * nb11 + mapping0.i2 * nb12);
1042
+ const float *row1_ptr = (const float *) ((const uint8_t *) src + i11_1 * nb11 + mapping1.i2 * nb12);
1043
+
1044
+ uint32_t c = 0;
1045
+ for (; c + 32 <= k_valid; c += 32) {
1046
+ HVX_Vector v0 = *(const HVX_Vector *)(row0_ptr + c);
1047
+ HVX_Vector v1 = *(const HVX_Vector *)(row1_ptr + c);
1048
+ HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
1049
+
1050
+ uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS;
1051
+ uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
1052
+
1053
+ HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
1054
+ tile[r1 / 2] = v_out;
1055
+ }
1056
+ if (c < k_block) {
1057
+ HVX_Vector v0 = *(const HVX_Vector *)(row0_ptr + c);
1058
+ HVX_Vector v1 = *(const HVX_Vector *)(row1_ptr + c);
1059
+
1060
+ uint32_t rem = k_valid - c;
1061
+ HVX_VectorPred mask = Q6_Q_vsetq2_R(rem > 0 ? rem * sizeof(float) : 0);
1062
+ v0 = Q6_V_vmux_QVV(mask, v0, Q6_V_vzero());
1063
+ v1 = Q6_V_vmux_QVV(mask, v1, Q6_V_vzero());
1064
+
1065
+ HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
1066
+
1067
+ uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS;
1068
+ uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
1069
+
1070
+ HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
1071
+ tile[r1 / 2] = v_out;
1072
+ }
1073
+ }
1074
+
1075
+ for (; r < n_rows_padded; r += 2) {
1076
+ uint32_t r_idx0 = start_row + r;
1077
+ uint32_t r0 = r_idx0 / HTP_MM_HMX_TILE_N_ROWS; // tile row index
1078
+ uint32_t r1 = r_idx0 % HTP_MM_HMX_TILE_N_ROWS; // intra-tile row idx
1079
+
1080
+ const bool row0_valid = (start_row + r + 0) < cne1;
1081
+ const bool row1_valid = (start_row + r + 1) < cne1;
1082
+
1083
+ const float *row0_ptr = NULL;
1084
+ const float *row1_ptr = NULL;
1085
+
1086
+ if (row0_valid) {
1087
+ struct mmid_row_mapping mapping0 = matrix_rows[cur_a * mapping_stride + (start_row + r + 0)];
1088
+ uint32_t i11_0 = fastmodulo(mapping0.i1, ne11, ne11_div);
1089
+ row0_ptr = (const float *) ((const uint8_t *) src + i11_0 * nb11 + mapping0.i2 * nb12);
1090
+ }
1091
+ if (row1_valid) {
1092
+ struct mmid_row_mapping mapping1 = matrix_rows[cur_a * mapping_stride + (start_row + r + 1)];
1093
+ uint32_t i11_1 = fastmodulo(mapping1.i1, ne11, ne11_div);
1094
+ row1_ptr = (const float *) ((const uint8_t *) src + i11_1 * nb11 + mapping1.i2 * nb12);
1095
+ }
1096
+
1097
+ uint32_t c = 0;
1098
+ for (; c + 32 <= k_valid; c += 32) {
1099
+ HVX_Vector v0 = Q6_V_vzero();
1100
+ HVX_Vector v1 = Q6_V_vzero();
1101
+ if (row0_valid) v0 = *(const HVX_Vector *)(row0_ptr + c);
1102
+ if (row1_valid) v1 = *(const HVX_Vector *)(row1_ptr + c);
1103
+
1104
+ HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
1105
+
1106
+ uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS;
1107
+ uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
1108
+
1109
+ HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
1110
+ tile[r1 / 2] = v_out;
1111
+ }
1112
+ if (c < k_block) {
1113
+ HVX_Vector v0 = Q6_V_vzero();
1114
+ HVX_Vector v1 = Q6_V_vzero();
1115
+ if (row0_valid) v0 = *(const HVX_Vector *)(row0_ptr + c);
1116
+ if (row1_valid) v1 = *(const HVX_Vector *)(row1_ptr + c);
1117
+
1118
+ uint32_t rem = k_valid - c;
1119
+ HVX_VectorPred mask = Q6_Q_vsetq2_R(rem > 0 ? rem * sizeof(float) : 0);
1120
+ v0 = Q6_V_vmux_QVV(mask, v0, Q6_V_vzero());
1121
+ v1 = Q6_V_vmux_QVV(mask, v1, Q6_V_vzero());
1122
+
1123
+ HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
1124
+
1125
+ uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS;
1126
+ uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
1127
+
1128
+ HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
1129
+ tile[r1 / 2] = v_out;
1130
+ }
1131
+ }
1132
+ }
1133
+
1134
+ static void transfer_activation_chunk_fp32_to_fp16_gathered_flat(
1135
+ __fp16 *restrict vtcm_dst,
1136
+ const float *restrict src,
1137
+ uint32_t start_row,
1138
+ uint32_t n_rows,
1139
+ uint32_t k_block,
1140
+ const struct mmid_row_mapping *matrix_rows,
1141
+ uint32_t cur_a,
1142
+ uint32_t mapping_stride,
1143
+ size_t nb12,
1144
+ uint32_t cne1,
1145
+ uint32_t k_valid) {
1146
+ const uint32_t n_rows_padded = hex_align_up(n_rows, HTP_MM_HMX_TILE_N_ROWS);
1147
+ const uint32_t n_rows_tiled = (n_rows / HTP_MM_HMX_TILE_N_ROWS) * HTP_MM_HMX_TILE_N_ROWS;
1148
+
1149
+ uint32_t r = 0;
1150
+
1151
+ #pragma unroll(2)
1152
+ for (r = 0; r < n_rows_tiled; r += 2) {
1153
+ uint32_t r_idx0 = start_row + r + 0;
1154
+ uint32_t r_idx1 = start_row + r + 1;
1155
+ uint32_t r0 = r_idx0 / HTP_MM_HMX_TILE_N_ROWS; // tile row index
1156
+ uint32_t r1 = r_idx0 % HTP_MM_HMX_TILE_N_ROWS; // intra-tile row idx
1157
+
1158
+ struct mmid_row_mapping mapping0 = matrix_rows[cur_a * mapping_stride + r_idx0];
1159
+ struct mmid_row_mapping mapping1 = matrix_rows[cur_a * mapping_stride + r_idx1];
1160
+
1161
+ const float *row0_ptr = (const float *) ((const uint8_t *) src + mapping0.i2 * nb12);
1162
+ const float *row1_ptr = (const float *) ((const uint8_t *) src + mapping1.i2 * nb12);
1163
+
1164
+ uint32_t c = 0;
1165
+ for (; c + 32 <= k_valid; c += 32) {
1166
+ HVX_Vector v0 = *(const HVX_Vector *)(row0_ptr + c);
1167
+ HVX_Vector v1 = *(const HVX_Vector *)(row1_ptr + c);
1168
+ HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
1169
+
1170
+ uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS;
1171
+ uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
1172
+
1173
+ HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
1174
+ tile[r1 / 2] = v_out;
1175
+ }
1176
+ if (c < k_block) {
1177
+ HVX_Vector v0 = *(const HVX_Vector *)(row0_ptr + c);
1178
+ HVX_Vector v1 = *(const HVX_Vector *)(row1_ptr + c);
1179
+
1180
+ uint32_t rem = k_valid - c;
1181
+ HVX_VectorPred mask = Q6_Q_vsetq2_R(rem > 0 ? rem * sizeof(float) : 0);
1182
+ v0 = Q6_V_vmux_QVV(mask, v0, Q6_V_vzero());
1183
+ v1 = Q6_V_vmux_QVV(mask, v1, Q6_V_vzero());
1184
+
1185
+ HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
1186
+
1187
+ uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS;
1188
+ uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
1189
+
1190
+ HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
1191
+ tile[r1 / 2] = v_out;
1192
+ }
1193
+ }
1194
+
1195
+ for (; r < n_rows_padded; r += 2) {
1196
+ uint32_t r_idx0 = start_row + r;
1197
+ uint32_t r0 = r_idx0 / HTP_MM_HMX_TILE_N_ROWS; // tile row index
1198
+ uint32_t r1 = r_idx0 % HTP_MM_HMX_TILE_N_ROWS; // intra-tile row idx
1199
+
1200
+ const bool row0_valid = (start_row + r + 0) < cne1;
1201
+ const bool row1_valid = (start_row + r + 1) < cne1;
1202
+
1203
+ const float *row0_ptr = NULL;
1204
+ const float *row1_ptr = NULL;
1205
+
1206
+ if (row0_valid) {
1207
+ struct mmid_row_mapping mapping0 = matrix_rows[cur_a * mapping_stride + (start_row + r + 0)];
1208
+ row0_ptr = (const float *) ((const uint8_t *) src + mapping0.i2 * nb12);
1209
+ }
1210
+ if (row1_valid) {
1211
+ struct mmid_row_mapping mapping1 = matrix_rows[cur_a * mapping_stride + (start_row + r + 1)];
1212
+ row1_ptr = (const float *) ((const uint8_t *) src + mapping1.i2 * nb12);
1213
+ }
1214
+
1215
+ uint32_t c = 0;
1216
+ for (; c + 32 <= k_valid; c += 32) {
1217
+ HVX_Vector v0 = Q6_V_vzero();
1218
+ HVX_Vector v1 = Q6_V_vzero();
1219
+ if (row0_valid) v0 = *(const HVX_Vector *)(row0_ptr + c);
1220
+ if (row1_valid) v1 = *(const HVX_Vector *)(row1_ptr + c);
1221
+
1222
+ HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
1223
+
1224
+ uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS;
1225
+ uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
1226
+
1227
+ HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
1228
+ tile[r1 / 2] = v_out;
1229
+ }
1230
+ if (c < k_block) {
1231
+ HVX_Vector v0 = Q6_V_vzero();
1232
+ HVX_Vector v1 = Q6_V_vzero();
1233
+ if (row0_valid) v0 = *(const HVX_Vector *)(row0_ptr + c);
1234
+ if (row1_valid) v1 = *(const HVX_Vector *)(row1_ptr + c);
1235
+
1236
+ uint32_t rem = k_valid - c;
1237
+ HVX_VectorPred mask = Q6_Q_vsetq2_R(rem > 0 ? rem * sizeof(float) : 0);
1238
+ v0 = Q6_V_vmux_QVV(mask, v0, Q6_V_vzero());
1239
+ v1 = Q6_V_vmux_QVV(mask, v1, Q6_V_vzero());
1240
+
1241
+ HVX_Vector v_out = hvx_vec_f32_to_f16_shuff(v0, v1);
1242
+
1243
+ uint32_t c0 = c / HTP_MM_HMX_TILE_N_COLS;
1244
+ uint32_t tile_idx = r0 * (k_block / HTP_MM_HMX_TILE_N_COLS) + c0;
1245
+
1246
+ HVX_Vector *tile = (HVX_Vector *) (vtcm_dst + tile_idx * HTP_MM_HMX_TILE_N_ELMS);
1247
+ tile[r1 / 2] = v_out;
1248
+ }
1249
+ }
1250
+ }
1251
+
1252
+ static void transfer_output_chunk_fp16_to_fp32_scattered(
1253
+ float *restrict dst,
1254
+ const __fp16 *restrict vtcm_src,
1255
+ uint32_t start_row,
1256
+ uint32_t n_rows,
1257
+ uint32_t n_cols,
1258
+ const struct mmid_row_mapping *matrix_rows,
1259
+ uint32_t cur_a,
1260
+ uint32_t mapping_stride,
1261
+ size_t dst_nb1,
1262
+ size_t dst_nb2,
1263
+ uint32_t cne1) {
1264
+ assert(n_cols % HTP_MM_HMX_TILE_N_COLS == 0);
1265
+ const size_t tile_row_stride = (n_cols / HTP_MM_HMX_TILE_N_COLS) * HTP_MM_HMX_TILE_N_ELMS;
1266
+
1267
+ const HVX_Vector one = hvx_vec_splat_f16(1.0);
1268
+
1269
+ for (size_t r = 0; r < n_rows; r += 2) {
1270
+ uint32_t r_idx0 = start_row + r + 0;
1271
+ uint32_t r_idx1 = start_row + r + 1;
1272
+ const size_t r0 = r_idx0 / HTP_MM_HMX_TILE_N_ROWS;
1273
+ const size_t r1 = (r_idx0 % HTP_MM_HMX_TILE_N_ROWS) / 2; // index of the row pair within the tile
1274
+ const __fp16 *row_base = vtcm_src + r0 * tile_row_stride;
1275
+
1276
+ if (r_idx0 >= cne1) break;
1277
+
1278
+ struct mmid_row_mapping mapping0 = matrix_rows[cur_a * mapping_stride + r_idx0];
1279
+ float *output_row0 = (float *) ((uint8_t *) dst + mapping0.i1 * dst_nb1 + mapping0.i2 * dst_nb2);
1280
+
1281
+ float *output_row1 = NULL;
1282
+ if (r_idx1 < cne1) {
1283
+ struct mmid_row_mapping mapping1 = matrix_rows[cur_a * mapping_stride + r_idx1];
1284
+ output_row1 = (float *) ((uint8_t *) dst + mapping1.i1 * dst_nb1 + mapping1.i2 * dst_nb2);
1285
+ }
1286
+
1287
+ #pragma unroll(4)
1288
+ for (size_t c = 0; c < (size_t)n_cols; c += HTP_MM_HMX_TILE_N_COLS) {
1289
+ const size_t c0 = c / HTP_MM_HMX_TILE_N_COLS;
1290
+ const __fp16 *tile = row_base + c0 * HTP_MM_HMX_TILE_N_ELMS;
1291
+ HVX_Vector v = ((const HVX_Vector *) tile)[r1];
1292
+ HVX_VectorPair vp = Q6_Wqf32_vmpy_VhfVhf(v, one);
1293
+
1294
+ HVX_Vector *pv_out0 = (HVX_Vector *) (output_row0 + c);
1295
+ HVX_Vector *pv_out1 = output_row1 ? (HVX_Vector *) (output_row1 + c) : NULL;
1296
+
1297
+ *pv_out0 = Q6_Vsf_equals_Vqf32(Q6_V_lo_W(vp));
1298
+ if (pv_out1) {
1299
+ *pv_out1 = Q6_Vsf_equals_Vqf32(Q6_V_hi_W(vp));
1300
+ }
1301
+ }
1302
+ }
1303
+ }