whispercpp 1.3.6 → 1.3.8

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (965) hide show
  1. checksums.yaml +4 -4
  2. data/.document +3 -0
  3. data/.rdoc_options +2 -0
  4. data/README.md +43 -9
  5. data/Rakefile +18 -3
  6. data/ext/dependencies.rb +10 -4
  7. data/ext/dependencies_for_windows.rb +17 -0
  8. data/ext/extconf.rb +20 -8
  9. data/ext/options.rb +54 -14
  10. data/ext/options_for_windows.rb +51 -0
  11. data/ext/ruby_whisper.c +35 -42
  12. data/ext/ruby_whisper.h +141 -0
  13. data/ext/ruby_whisper_context.c +157 -29
  14. data/ext/ruby_whisper_log_queue.c +180 -0
  15. data/ext/ruby_whisper_log_settable.h +46 -0
  16. data/ext/ruby_whisper_parakeet.c +49 -0
  17. data/ext/ruby_whisper_parakeet_context.c +304 -0
  18. data/ext/ruby_whisper_parakeet_context_params.c +117 -0
  19. data/ext/ruby_whisper_parakeet_model.c +84 -0
  20. data/ext/ruby_whisper_parakeet_params.c +548 -0
  21. data/ext/ruby_whisper_parakeet_segment.c +157 -0
  22. data/ext/ruby_whisper_parakeet_token.c +188 -0
  23. data/ext/ruby_whisper_parakeet_transcribe.cpp +58 -0
  24. data/ext/ruby_whisper_params.c +265 -73
  25. data/ext/ruby_whisper_segment.c +6 -6
  26. data/ext/ruby_whisper_transcribe.cpp +23 -15
  27. data/ext/ruby_whisper_vad_context.c +30 -10
  28. data/ext/ruby_whisper_vad_context_detect.cpp +8 -9
  29. data/ext/ruby_whisper_vad_params.c +4 -4
  30. data/ext/ruby_whisper_vad_segment.c +2 -2
  31. data/ext/sources/CMakeLists.txt +42 -3
  32. data/ext/sources/CMakePresets.json +95 -0
  33. data/ext/sources/cmake/parakeet-config.cmake.in +30 -0
  34. data/ext/sources/cmake/parakeet.pc.in +10 -0
  35. data/ext/sources/cmake/whisper.pc.in +2 -2
  36. data/ext/sources/examples/CMakeLists.txt +4 -2
  37. data/ext/sources/examples/bench/bench.cpp +1 -1
  38. data/ext/sources/examples/cli/cli.cpp +52 -10
  39. data/ext/sources/examples/common-ggml.cpp +4 -0
  40. data/ext/sources/examples/common-whisper.cpp +139 -67
  41. data/ext/sources/examples/common-whisper.h +11 -0
  42. data/ext/sources/examples/ffmpeg-transcode.cpp +211 -341
  43. data/ext/sources/examples/parakeet-cli/CMakeLists.txt +8 -0
  44. data/ext/sources/examples/parakeet-cli/parakeet-cli.cpp +243 -0
  45. data/ext/sources/examples/parakeet-quantize/CMakeLists.txt +7 -0
  46. data/ext/sources/examples/parakeet-quantize/parakeet-quantize.cpp +230 -0
  47. data/ext/sources/examples/server/server.cpp +199 -163
  48. data/ext/sources/examples/vad-speech-segments/speech.cpp +3 -2
  49. data/ext/sources/ggml/CMakeLists.txt +21 -14
  50. data/ext/sources/ggml/cmake/FindNCCL.cmake +36 -0
  51. data/ext/sources/ggml/cmake/ggml-config.cmake.in +12 -2
  52. data/ext/sources/ggml/include/ggml-alloc.h +1 -0
  53. data/ext/sources/ggml/include/ggml-backend.h +72 -10
  54. data/ext/sources/ggml/include/ggml-cuda.h +2 -2
  55. data/ext/sources/ggml/include/ggml-rpc.h +3 -3
  56. data/ext/sources/ggml/include/ggml-sycl.h +8 -0
  57. data/ext/sources/ggml/include/ggml.h +103 -9
  58. data/ext/sources/ggml/include/gguf.h +10 -2
  59. data/ext/sources/ggml/src/CMakeLists.txt +30 -6
  60. data/ext/sources/ggml/src/ggml-alloc.c +5 -1
  61. data/ext/sources/ggml/src/ggml-backend-impl.h +22 -2
  62. data/ext/sources/ggml/src/ggml-backend-meta.cpp +2266 -0
  63. data/ext/sources/ggml/src/ggml-backend-reg.cpp +12 -0
  64. data/ext/sources/ggml/src/ggml-backend.cpp +110 -9
  65. data/ext/sources/ggml/src/ggml-blas/ggml-blas.cpp +4 -0
  66. data/ext/sources/ggml/src/ggml-cann/aclnn_ops.cpp +672 -257
  67. data/ext/sources/ggml/src/ggml-cann/aclnn_ops.h +71 -0
  68. data/ext/sources/ggml/src/ggml-cann/common.h +20 -10
  69. data/ext/sources/ggml/src/ggml-cann/ggml-cann.cpp +211 -30
  70. data/ext/sources/ggml/src/ggml-common.h +24 -2
  71. data/ext/sources/ggml/src/ggml-cpu/CMakeLists.txt +59 -30
  72. data/ext/sources/ggml/src/ggml-cpu/amx/amx.cpp +2 -0
  73. data/ext/sources/ggml/src/ggml-cpu/amx/mmq.cpp +21 -22
  74. data/ext/sources/ggml/src/ggml-cpu/arch/arm/quants.c +194 -11
  75. data/ext/sources/ggml/src/ggml-cpu/arch/arm/repack.cpp +65 -0
  76. data/ext/sources/ggml/src/ggml-cpu/arch/loongarch/quants.c +151 -1
  77. data/ext/sources/ggml/src/ggml-cpu/arch/powerpc/quants.c +0 -1
  78. data/ext/sources/ggml/src/ggml-cpu/arch/riscv/quants.c +4279 -1292
  79. data/ext/sources/ggml/src/ggml-cpu/arch/riscv/repack.cpp +5 -35
  80. data/ext/sources/ggml/src/ggml-cpu/arch/s390/quants.c +0 -1
  81. data/ext/sources/ggml/src/ggml-cpu/arch/wasm/quants.c +72 -1
  82. data/ext/sources/ggml/src/ggml-cpu/arch/x86/quants.c +319 -31
  83. data/ext/sources/ggml/src/ggml-cpu/arch/x86/repack.cpp +1 -1
  84. data/ext/sources/ggml/src/ggml-cpu/arch-fallback.h +12 -2
  85. data/ext/sources/ggml/src/ggml-cpu/cmake/FindSMTIME.cmake +32 -0
  86. data/ext/sources/ggml/src/ggml-cpu/ggml-cpu-impl.h +10 -0
  87. data/ext/sources/ggml/src/ggml-cpu/ggml-cpu.c +109 -5
  88. data/ext/sources/ggml/src/ggml-cpu/ggml-cpu.cpp +2 -0
  89. data/ext/sources/ggml/src/ggml-cpu/kleidiai/kleidiai.cpp +146 -134
  90. data/ext/sources/ggml/src/ggml-cpu/llamafile/sgemm.cpp +107 -82
  91. data/ext/sources/ggml/src/ggml-cpu/ops.cpp +501 -119
  92. data/ext/sources/ggml/src/ggml-cpu/ops.h +3 -0
  93. data/ext/sources/ggml/src/ggml-cpu/quants.c +106 -0
  94. data/ext/sources/ggml/src/ggml-cpu/quants.h +6 -0
  95. data/ext/sources/ggml/src/ggml-cpu/repack.cpp +3 -0
  96. data/ext/sources/ggml/src/ggml-cpu/simd-gemm.h +91 -1
  97. data/ext/sources/ggml/src/ggml-cpu/simd-mappings.h +14 -16
  98. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime.cpp +1402 -687
  99. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime.h +8 -0
  100. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime1_kernels.cpp +597 -2766
  101. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime2_kernels.cpp +5768 -0
  102. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime_env.cpp +320 -0
  103. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime_env.h +55 -0
  104. data/ext/sources/ggml/src/ggml-cpu/spacemit/ime_kernels.h +182 -19
  105. data/ext/sources/ggml/src/ggml-cpu/spacemit/repack.cpp +1795 -0
  106. data/ext/sources/ggml/src/ggml-cpu/spacemit/repack.h +14 -0
  107. data/ext/sources/ggml/src/ggml-cpu/spacemit/rvv_kernels.cpp +3178 -0
  108. data/ext/sources/ggml/src/ggml-cpu/spacemit/rvv_kernels.h +95 -0
  109. data/ext/sources/ggml/src/ggml-cpu/spacemit/spine_barrier.h +34 -0
  110. data/ext/sources/ggml/src/ggml-cpu/spacemit/spine_mem_pool.cpp +760 -0
  111. data/ext/sources/ggml/src/ggml-cpu/spacemit/spine_mem_pool.h +32 -0
  112. data/ext/sources/ggml/src/ggml-cpu/spacemit/spine_tcm.h +409 -0
  113. data/ext/sources/ggml/src/ggml-cpu/vec.cpp +39 -55
  114. data/ext/sources/ggml/src/ggml-cpu/vec.h +225 -240
  115. data/ext/sources/ggml/src/ggml-cuda/CMakeLists.txt +17 -7
  116. data/ext/sources/ggml/src/ggml-cuda/allreduce.cu +971 -0
  117. data/ext/sources/ggml/src/ggml-cuda/allreduce.cuh +29 -0
  118. data/ext/sources/ggml/src/ggml-cuda/argsort.cu +62 -26
  119. data/ext/sources/ggml/src/ggml-cuda/binbcast.cu +134 -64
  120. data/ext/sources/ggml/src/ggml-cuda/binbcast.cuh +1 -0
  121. data/ext/sources/ggml/src/ggml-cuda/col2im-1d.cu +81 -0
  122. data/ext/sources/ggml/src/ggml-cuda/col2im-1d.cuh +3 -0
  123. data/ext/sources/ggml/src/ggml-cuda/common.cuh +246 -28
  124. data/ext/sources/ggml/src/ggml-cuda/concat.cu +134 -116
  125. data/ext/sources/ggml/src/ggml-cuda/conv-transpose-1d.cu +14 -12
  126. data/ext/sources/ggml/src/ggml-cuda/conv2d-transpose.cu +45 -21
  127. data/ext/sources/ggml/src/ggml-cuda/conv2d-transpose.cuh +1 -0
  128. data/ext/sources/ggml/src/ggml-cuda/convert.cu +139 -34
  129. data/ext/sources/ggml/src/ggml-cuda/convert.cuh +10 -0
  130. data/ext/sources/ggml/src/ggml-cuda/cpy.cu +88 -29
  131. data/ext/sources/ggml/src/ggml-cuda/dequantize.cuh +22 -0
  132. data/ext/sources/ggml/src/ggml-cuda/fattn-common.cuh +287 -49
  133. data/ext/sources/ggml/src/ggml-cuda/fattn-mma-f16.cuh +335 -130
  134. data/ext/sources/ggml/src/ggml-cuda/fattn-tile.cu +12 -0
  135. data/ext/sources/ggml/src/ggml-cuda/fattn-tile.cuh +127 -24
  136. data/ext/sources/ggml/src/ggml-cuda/fattn-vec.cuh +40 -15
  137. data/ext/sources/ggml/src/ggml-cuda/fattn-wmma-f16.cu +18 -9
  138. data/ext/sources/ggml/src/ggml-cuda/fattn.cu +169 -60
  139. data/ext/sources/ggml/src/ggml-cuda/fattn.cuh +2 -0
  140. data/ext/sources/ggml/src/ggml-cuda/fwht.cu +101 -0
  141. data/ext/sources/ggml/src/ggml-cuda/fwht.cuh +4 -0
  142. data/ext/sources/ggml/src/ggml-cuda/gated_delta_net.cu +109 -45
  143. data/ext/sources/ggml/src/ggml-cuda/gated_delta_net.cuh +10 -0
  144. data/ext/sources/ggml/src/ggml-cuda/getrows.cu +48 -23
  145. data/ext/sources/ggml/src/ggml-cuda/ggml-cuda.cu +2034 -2104
  146. data/ext/sources/ggml/src/ggml-cuda/im2col.cu +32 -29
  147. data/ext/sources/ggml/src/ggml-cuda/mean.cu +4 -2
  148. data/ext/sources/ggml/src/ggml-cuda/mma.cuh +242 -195
  149. data/ext/sources/ggml/src/ggml-cuda/mmf.cuh +3 -3
  150. data/ext/sources/ggml/src/ggml-cuda/mmq.cu +25 -12
  151. data/ext/sources/ggml/src/ggml-cuda/mmq.cuh +502 -423
  152. data/ext/sources/ggml/src/ggml-cuda/mmvf.cu +19 -12
  153. data/ext/sources/ggml/src/ggml-cuda/mmvq.cu +562 -97
  154. data/ext/sources/ggml/src/ggml-cuda/mmvq.cuh +6 -1
  155. data/ext/sources/ggml/src/ggml-cuda/norm.cu +36 -10
  156. data/ext/sources/ggml/src/ggml-cuda/out-prod.cu +66 -7
  157. data/ext/sources/ggml/src/ggml-cuda/quantize.cu +133 -26
  158. data/ext/sources/ggml/src/ggml-cuda/quantize.cuh +1 -1
  159. data/ext/sources/ggml/src/ggml-cuda/reduce_rows.cuh +5 -1
  160. data/ext/sources/ggml/src/ggml-cuda/rope.cu +11 -4
  161. data/ext/sources/ggml/src/ggml-cuda/scale.cu +4 -1
  162. data/ext/sources/ggml/src/ggml-cuda/set-rows.cu +78 -10
  163. data/ext/sources/ggml/src/ggml-cuda/snake.cu +72 -0
  164. data/ext/sources/ggml/src/ggml-cuda/snake.cuh +8 -0
  165. data/ext/sources/ggml/src/ggml-cuda/softcap.cu +4 -1
  166. data/ext/sources/ggml/src/ggml-cuda/ssm-conv.cu +45 -13
  167. data/ext/sources/ggml/src/ggml-cuda/ssm-conv.cuh +1 -1
  168. data/ext/sources/ggml/src/ggml-cuda/ssm-scan.cu +40 -18
  169. data/ext/sources/ggml/src/ggml-cuda/sumrows.cu +8 -4
  170. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_1-ncols2_16.cu +1 -0
  171. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_1-ncols2_32.cu +1 -0
  172. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_1-ncols2_8.cu +2 -0
  173. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_16-ncols2_2.cu +1 -0
  174. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_16-ncols2_4.cu +1 -0
  175. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_2-ncols2_16.cu +1 -0
  176. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_2-ncols2_32.cu +1 -0
  177. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_2-ncols2_4.cu +1 -0
  178. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_2-ncols2_8.cu +2 -0
  179. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_32-ncols2_2.cu +1 -0
  180. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_16.cu +1 -0
  181. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_2.cu +1 -0
  182. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_4.cu +1 -0
  183. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_8.cu +2 -0
  184. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_8-ncols2_2.cu +1 -0
  185. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_8-ncols2_4.cu +1 -0
  186. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_8-ncols2_8.cu +2 -0
  187. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-tile-instance-dkq192-dv128.cu +5 -0
  188. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-tile-instance-dkq320-dv256.cu +5 -0
  189. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-tile-instance-dkq512-dv512.cu +5 -0
  190. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-bf16.cu +7 -0
  191. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-f16.cu +7 -0
  192. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q4_0.cu +7 -0
  193. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q4_1.cu +7 -0
  194. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q5_0.cu +7 -0
  195. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q5_1.cu +7 -0
  196. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-bf16-q8_0.cu +7 -0
  197. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-f16-bf16.cu +7 -0
  198. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q4_0-bf16.cu +7 -0
  199. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q4_1-bf16.cu +7 -0
  200. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q5_0-bf16.cu +7 -0
  201. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q5_1-bf16.cu +7 -0
  202. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-vec-instance-q8_0-bf16.cu +7 -0
  203. data/ext/sources/ggml/src/ggml-cuda/template-instances/mmq-instance-nvfp4.cu +5 -0
  204. data/ext/sources/ggml/src/ggml-cuda/template-instances/mmq-instance-q1_0.cu +5 -0
  205. data/ext/sources/ggml/src/ggml-cuda/top-k.cu +5 -4
  206. data/ext/sources/ggml/src/ggml-cuda/topk-moe.cu +33 -24
  207. data/ext/sources/ggml/src/ggml-cuda/unary.cu +31 -2
  208. data/ext/sources/ggml/src/ggml-cuda/unary.cuh +2 -0
  209. data/ext/sources/ggml/src/ggml-cuda/vecdotq.cuh +80 -0
  210. data/ext/sources/ggml/src/ggml-cuda/vendors/cuda.h +7 -2
  211. data/ext/sources/ggml/src/ggml-cuda/vendors/hip.h +23 -4
  212. data/ext/sources/ggml/src/ggml-cuda/vendors/musa.h +4 -0
  213. data/ext/sources/ggml/src/ggml-hexagon/CMakeLists.txt +1 -5
  214. data/ext/sources/ggml/src/ggml-hexagon/ggml-hexagon.cpp +2788 -1762
  215. data/ext/sources/ggml/src/ggml-hexagon/htp/CMakeLists.txt +13 -4
  216. data/ext/sources/ggml/src/ggml-hexagon/htp/act-ops.c +53 -84
  217. data/ext/sources/ggml/src/ggml-hexagon/htp/argsort-ops.c +25 -12
  218. data/ext/sources/ggml/src/ggml-hexagon/htp/binary-ops.c +165 -184
  219. data/ext/sources/ggml/src/ggml-hexagon/htp/cmake-toolchain.cmake +17 -19
  220. data/ext/sources/ggml/src/ggml-hexagon/htp/concat-ops.c +277 -0
  221. data/ext/sources/ggml/src/ggml-hexagon/htp/cpy-ops.c +170 -127
  222. data/ext/sources/ggml/src/ggml-hexagon/htp/cumsum-ops.c +270 -0
  223. data/ext/sources/ggml/src/ggml-hexagon/htp/diag-ops.c +216 -0
  224. data/ext/sources/ggml/src/ggml-hexagon/htp/fill-ops.c +123 -0
  225. data/ext/sources/ggml/src/ggml-hexagon/htp/flash-attn-ops.c +1774 -396
  226. data/ext/sources/ggml/src/ggml-hexagon/htp/flash-attn-ops.h +303 -0
  227. data/ext/sources/ggml/src/ggml-hexagon/htp/gated-delta-net-ops.c +1148 -0
  228. data/ext/sources/ggml/src/ggml-hexagon/htp/get-rows-ops.c +148 -42
  229. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-common.h +80 -0
  230. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-dma.c +2 -2
  231. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-dma.h +255 -62
  232. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-dump.h +9 -0
  233. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-profile.h +64 -0
  234. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-utils.h +25 -21
  235. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-fa-kernels.h +555 -0
  236. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-mm-kernels-tiled.h +1303 -0
  237. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-queue.c +167 -0
  238. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-queue.h +157 -0
  239. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-utils.h +222 -0
  240. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-ctx.h +104 -13
  241. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-ops.h +222 -57
  242. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-vtcm.h +19 -0
  243. data/ext/sources/ggml/src/ggml-hexagon/htp/htp_iface.idl +10 -3
  244. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-base.h +78 -26
  245. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-copy.h +27 -10
  246. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-div.h +63 -23
  247. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-exp.h +48 -8
  248. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-fa-kernels.h +232 -0
  249. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-flash-attn.h +47 -0
  250. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-log.h +65 -0
  251. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-flat.h +1511 -0
  252. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-tiled.h +1200 -0
  253. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-pow.h +42 -0
  254. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-repl.h +74 -0
  255. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-sigmoid.h +40 -0
  256. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-sin-cos.h +90 -0
  257. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-utils.h +5 -8
  258. data/ext/sources/ggml/src/ggml-hexagon/htp/main.c +625 -816
  259. data/ext/sources/ggml/src/ggml-hexagon/htp/matmul-ops.c +3052 -2166
  260. data/ext/sources/ggml/src/ggml-hexagon/htp/matmul-ops.h +650 -0
  261. data/ext/sources/ggml/src/ggml-hexagon/htp/pad-ops.c +547 -0
  262. data/ext/sources/ggml/src/ggml-hexagon/htp/repeat-ops.c +148 -0
  263. data/ext/sources/ggml/src/ggml-hexagon/htp/rope-ops.c +337 -106
  264. data/ext/sources/ggml/src/ggml-hexagon/htp/set-rows-ops.c +59 -37
  265. data/ext/sources/ggml/src/ggml-hexagon/htp/softmax-ops.c +121 -133
  266. data/ext/sources/ggml/src/ggml-hexagon/htp/solve-tri-ops.c +267 -0
  267. data/ext/sources/ggml/src/ggml-hexagon/htp/ssm-conv.c +245 -151
  268. data/ext/sources/ggml/src/ggml-hexagon/htp/sum-rows-ops.c +6 -6
  269. data/ext/sources/ggml/src/ggml-hexagon/htp/unary-ops.c +719 -45
  270. data/ext/sources/ggml/src/ggml-hexagon/htp/worker-pool.c +15 -3
  271. data/ext/sources/ggml/src/ggml-hexagon/htp/worker-pool.h +8 -0
  272. data/ext/sources/ggml/src/ggml-hexagon/htp-opnode.h +390 -0
  273. data/ext/sources/ggml/src/ggml-hexagon/libggml-htp.inf +3 -5
  274. data/ext/sources/ggml/src/ggml-hip/CMakeLists.txt +27 -9
  275. data/ext/sources/ggml/src/ggml-impl.h +6 -1
  276. data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.cpp +207 -18
  277. data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.h +36 -2
  278. data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.m +186 -29
  279. data/ext/sources/ggml/src/ggml-metal/ggml-metal-impl.h +118 -0
  280. data/ext/sources/ggml/src/ggml-metal/ggml-metal-ops.cpp +322 -21
  281. data/ext/sources/ggml/src/ggml-metal/ggml-metal-ops.h +4 -0
  282. data/ext/sources/ggml/src/ggml-metal/ggml-metal.cpp +39 -26
  283. data/ext/sources/ggml/src/ggml-metal/ggml-metal.metal +1226 -467
  284. data/ext/sources/ggml/src/ggml-musa/CMakeLists.txt +5 -6
  285. data/ext/sources/ggml/src/ggml-opencl/CMakeLists.txt +67 -5
  286. data/ext/sources/ggml/src/ggml-opencl/fa_tune.h +92 -0
  287. data/ext/sources/ggml/src/ggml-opencl/ggml-opencl.cpp +16290 -6246
  288. data/ext/sources/ggml/src/ggml-opencl/kernels/concat.cl +67 -0
  289. data/ext/sources/ggml/src/ggml-opencl/kernels/cpy.cl +59 -0
  290. data/ext/sources/ggml/src/ggml-opencl/kernels/cvt.cl +1997 -92
  291. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f16.cl +81 -41
  292. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32.cl +88 -39
  293. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_f16.cl +1995 -96
  294. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_q4_0.cl +1615 -0
  295. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_q8_0.cl +1486 -0
  296. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_pre_f16.cl +156 -0
  297. data/ext/sources/ggml/src/ggml-opencl/kernels/gated_delta_net.cl +249 -0
  298. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_mxfp4_f32_ns.cl +374 -0
  299. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_0_f32_ns.cl +324 -0
  300. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_1_f32_ns.cl +326 -0
  301. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_k_f32_ns.cl +348 -0
  302. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_0_f32_ns.cl +328 -0
  303. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_1_f32_ns.cl +330 -0
  304. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_k_f32_ns.cl +356 -0
  305. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q6_k_f32_ns.cl +335 -0
  306. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_iq4_nl_f32.cl +150 -0
  307. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q1_0_f32.cl +94 -0
  308. data/ext/sources/ggml/src/ggml-opencl/kernels/{mul_mat_Ab_Bi_8x4.cl → gemm_noshuffle_q4_0_f32.cl} +1 -1
  309. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q4_k_f32.cl +172 -0
  310. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q5_0_f32.cl +131 -0
  311. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q5_1_f32.cl +134 -0
  312. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q5_k_f32.cl +176 -0
  313. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q6_k_f32.cl +140 -0
  314. data/ext/sources/ggml/src/ggml-opencl/kernels/{mul_mm_q8_0_f32_8x4.cl → gemm_noshuffle_q8_0_f32.cl} +1 -1
  315. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_xmem_f16_f32_os8.cl +233 -0
  316. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_mxfp4_f32_ns.cl +165 -0
  317. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q4_0_f32_ns.cl +120 -0
  318. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q4_1_f32_ns.cl +123 -0
  319. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q4_k_f32_ns.cl +155 -0
  320. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q5_0_f32_ns.cl +123 -0
  321. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q5_1_f32_ns.cl +125 -0
  322. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q5_k_f32_ns.cl +160 -0
  323. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_moe_q6_k_f32_ns.cl +141 -0
  324. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_iq4_nl_f32.cl +302 -0
  325. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q1_0_f32.cl +121 -0
  326. data/ext/sources/ggml/src/ggml-opencl/kernels/{gemv_noshuffle_general.cl → gemv_noshuffle_q4_0_f32.cl} +5 -5
  327. data/ext/sources/ggml/src/ggml-opencl/kernels/{gemv_noshuffle.cl → gemv_noshuffle_q4_0_f32_spec.cl} +5 -5
  328. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q4_k_f32.cl +318 -0
  329. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q5_0_f32.cl +291 -0
  330. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q5_1_f32.cl +294 -0
  331. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q5_k_f32.cl +326 -0
  332. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q6_k_f32.cl +293 -0
  333. data/ext/sources/ggml/src/ggml-opencl/kernels/{gemv_noshuffle_general_q8_0_f32.cl → gemv_noshuffle_q8_0_f32.cl} +1 -1
  334. data/ext/sources/ggml/src/ggml-opencl/kernels/get_rows.cl +15 -9
  335. data/ext/sources/ggml/src/ggml-opencl/kernels/moe_reorder_b.cl +30 -0
  336. data/ext/sources/ggml/src/ggml-opencl/kernels/moe_sort_by_expert.cl +82 -0
  337. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_iq4_nl_f32_l4_lm.cl +171 -0
  338. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q1_0_f32_l4_lm.cl +156 -0
  339. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q4_k_f32_l4_lm.cl +179 -0
  340. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q5_0_f32_l4_lm.cl +173 -0
  341. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q5_1_f32_l4_lm.cl +175 -0
  342. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q5_k_f32_l4_lm.cl +192 -0
  343. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_f16_f32_l4.cl +1149 -0
  344. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_iq4_nl_f32.cl +164 -0
  345. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_iq4_nl_f32_flat.cl +202 -0
  346. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q1_0_f32.cl +141 -0
  347. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q1_0_f32_flat.cl +190 -0
  348. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q4_k_f32_flat.cl +196 -0
  349. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_0_f32.cl +241 -0
  350. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_0_f32_flat.cl +243 -0
  351. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_1_f32.cl +243 -0
  352. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_1_f32_flat.cl +247 -0
  353. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_k_f32.cl +187 -0
  354. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q5_k_f32_flat.cl +203 -0
  355. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q6_k_f32_flat.cl +48 -64
  356. data/ext/sources/ggml/src/ggml-opencl/kernels/norm.cl +5 -2
  357. data/ext/sources/ggml/src/ggml-opencl/kernels/set_rows.cl +500 -0
  358. data/ext/sources/ggml/src/ggml-opencl/libdl.h +79 -0
  359. data/ext/sources/ggml/src/ggml-openvino/.clang-format +0 -5
  360. data/ext/sources/ggml/src/ggml-openvino/CMakeLists.txt +2 -4
  361. data/ext/sources/ggml/src/ggml-openvino/ggml-decoder.cpp +740 -127
  362. data/ext/sources/ggml/src/ggml-openvino/ggml-decoder.h +76 -23
  363. data/ext/sources/ggml/src/ggml-openvino/ggml-openvino-extra.cpp +75 -14
  364. data/ext/sources/ggml/src/ggml-openvino/ggml-openvino-extra.h +29 -8
  365. data/ext/sources/ggml/src/ggml-openvino/ggml-openvino.cpp +339 -69
  366. data/ext/sources/ggml/src/ggml-openvino/ggml-quants.cpp +330 -192
  367. data/ext/sources/ggml/src/ggml-openvino/ggml-quants.h +10 -4
  368. data/ext/sources/ggml/src/ggml-openvino/openvino/decoder.h +56 -16
  369. data/ext/sources/ggml/src/ggml-openvino/openvino/frontend.h +1 -1
  370. data/ext/sources/ggml/src/ggml-openvino/openvino/input_model.h +4 -4
  371. data/ext/sources/ggml/src/ggml-openvino/openvino/node_context.h +94 -37
  372. data/ext/sources/ggml/src/ggml-openvino/openvino/op/add_id.cpp +76 -0
  373. data/ext/sources/ggml/src/ggml-openvino/openvino/op/argsort.cpp +47 -0
  374. data/ext/sources/ggml/src/ggml-openvino/openvino/op/clamp.cpp +33 -0
  375. data/ext/sources/ggml/src/ggml-openvino/openvino/op/concat.cpp +48 -0
  376. data/ext/sources/ggml/src/ggml-openvino/openvino/op/cont.cpp +8 -16
  377. data/ext/sources/ggml/src/ggml-openvino/openvino/op/cpy.cpp +14 -1
  378. data/ext/sources/ggml/src/ggml-openvino/openvino/op/div.cpp +146 -0
  379. data/ext/sources/ggml/src/ggml-openvino/openvino/op/flash_attn_ext.cpp +108 -21
  380. data/ext/sources/ggml/src/ggml-openvino/openvino/op/gated_delta_net.cpp +282 -0
  381. data/ext/sources/ggml/src/ggml-openvino/openvino/op/gated_delta_net.hpp +65 -0
  382. data/ext/sources/ggml/src/ggml-openvino/openvino/op/get_rows.cpp +2 -9
  383. data/ext/sources/ggml/src/ggml-openvino/openvino/op/glu_geglu.cpp +21 -7
  384. data/ext/sources/ggml/src/ggml-openvino/openvino/op/glu_swiglu.cpp +41 -8
  385. data/ext/sources/ggml/src/ggml-openvino/openvino/op/im2col.cpp +120 -0
  386. data/ext/sources/ggml/src/ggml-openvino/openvino/op/l2_norm.cpp +44 -0
  387. data/ext/sources/ggml/src/ggml-openvino/openvino/op/mul_mat_id.cpp +226 -0
  388. data/ext/sources/ggml/src/ggml-openvino/openvino/op/mulmat.cpp +19 -9
  389. data/ext/sources/ggml/src/ggml-openvino/openvino/op/norm.cpp +58 -0
  390. data/ext/sources/ggml/src/ggml-openvino/openvino/op/pad.cpp +95 -0
  391. data/ext/sources/ggml/src/ggml-openvino/openvino/op/permute.cpp +58 -13
  392. data/ext/sources/ggml/src/ggml-openvino/openvino/op/repeat.cpp +74 -0
  393. data/ext/sources/ggml/src/ggml-openvino/openvino/op/reshape.cpp +13 -6
  394. data/ext/sources/ggml/src/ggml-openvino/openvino/op/rms_norm.cpp +1 -1
  395. data/ext/sources/ggml/src/ggml-openvino/openvino/op/rope.cpp +161 -39
  396. data/ext/sources/ggml/src/ggml-openvino/openvino/op/set_rows.cpp +3 -3
  397. data/ext/sources/ggml/src/ggml-openvino/openvino/op/softmax.cpp +126 -49
  398. data/ext/sources/ggml/src/ggml-openvino/openvino/op/ssm_conv.cpp +59 -0
  399. data/ext/sources/ggml/src/ggml-openvino/openvino/op/sum_rows.cpp +27 -0
  400. data/ext/sources/ggml/src/ggml-openvino/openvino/op/transpose.cpp +32 -1
  401. data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_silu.cpp +1 -1
  402. data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_softplus.cpp +38 -0
  403. data/ext/sources/ggml/src/ggml-openvino/openvino/op/view.cpp +90 -25
  404. data/ext/sources/ggml/src/ggml-openvino/openvino/op_table.cpp +41 -22
  405. data/ext/sources/ggml/src/ggml-openvino/openvino/op_table.h +18 -4
  406. data/ext/sources/ggml/src/ggml-openvino/openvino/pass/mark_decompression_convert_constant_folding.h +1 -1
  407. data/ext/sources/ggml/src/ggml-openvino/openvino/rt_info/weightless_caching_attributes.hpp +41 -0
  408. data/ext/sources/ggml/src/ggml-openvino/openvino/translate_session.cpp +70 -43
  409. data/ext/sources/ggml/src/ggml-openvino/openvino/translate_session.h +5 -4
  410. data/ext/sources/ggml/src/ggml-openvino/openvino/utils.cpp +612 -36
  411. data/ext/sources/ggml/src/ggml-openvino/openvino/utils.h +29 -26
  412. data/ext/sources/ggml/src/ggml-openvino/utils.cpp +460 -114
  413. data/ext/sources/ggml/src/ggml-openvino/utils.h +32 -9
  414. data/ext/sources/ggml/src/ggml-opt.cpp +1 -0
  415. data/ext/sources/ggml/src/ggml-quants.c +365 -114
  416. data/ext/sources/ggml/src/ggml-quants.h +6 -0
  417. data/ext/sources/ggml/src/ggml-rpc/CMakeLists.txt +24 -0
  418. data/ext/sources/ggml/src/ggml-rpc/ggml-rpc.cpp +167 -311
  419. data/ext/sources/ggml/src/ggml-rpc/transport.cpp +683 -0
  420. data/ext/sources/ggml/src/ggml-rpc/transport.h +34 -0
  421. data/ext/sources/ggml/src/ggml-sycl/CMakeLists.txt +50 -4
  422. data/ext/sources/ggml/src/ggml-sycl/add-id.cpp +1 -1
  423. data/ext/sources/ggml/src/ggml-sycl/backend.hpp +5 -1
  424. data/ext/sources/ggml/src/ggml-sycl/binbcast.cpp +12 -0
  425. data/ext/sources/ggml/src/ggml-sycl/col2im-1d.cpp +102 -0
  426. data/ext/sources/ggml/src/ggml-sycl/col2im-1d.hpp +8 -0
  427. data/ext/sources/ggml/src/ggml-sycl/common.cpp +72 -2
  428. data/ext/sources/ggml/src/ggml-sycl/common.hpp +59 -2
  429. data/ext/sources/ggml/src/ggml-sycl/concat.cpp +21 -1
  430. data/ext/sources/ggml/src/ggml-sycl/conv2d-dw.cpp +158 -0
  431. data/ext/sources/ggml/src/ggml-sycl/conv2d-dw.hpp +10 -0
  432. data/ext/sources/ggml/src/ggml-sycl/conv2d-transpose.cpp +125 -0
  433. data/ext/sources/ggml/src/ggml-sycl/conv2d-transpose.hpp +10 -0
  434. data/ext/sources/ggml/src/ggml-sycl/conv2d.cpp +150 -0
  435. data/ext/sources/ggml/src/ggml-sycl/conv2d.hpp +10 -0
  436. data/ext/sources/ggml/src/ggml-sycl/conv3d.cpp +224 -0
  437. data/ext/sources/ggml/src/ggml-sycl/conv3d.hpp +8 -0
  438. data/ext/sources/ggml/src/ggml-sycl/convert.cpp +121 -13
  439. data/ext/sources/ggml/src/ggml-sycl/convert.hpp +9 -0
  440. data/ext/sources/ggml/src/ggml-sycl/cpy.cpp +706 -0
  441. data/ext/sources/ggml/src/ggml-sycl/cpy.hpp +281 -0
  442. data/ext/sources/ggml/src/ggml-sycl/cross_entropy_loss.cpp +255 -0
  443. data/ext/sources/ggml/src/ggml-sycl/cross_entropy_loss.hpp +7 -0
  444. data/ext/sources/ggml/src/ggml-sycl/cumsum.cpp +148 -0
  445. data/ext/sources/ggml/src/ggml-sycl/cumsum.hpp +5 -0
  446. data/ext/sources/ggml/src/ggml-sycl/dequantize.hpp +678 -0
  447. data/ext/sources/ggml/src/ggml-sycl/diag.cpp +67 -0
  448. data/ext/sources/ggml/src/ggml-sycl/diag.hpp +5 -0
  449. data/ext/sources/ggml/src/ggml-sycl/dmmv.cpp +997 -244
  450. data/ext/sources/ggml/src/ggml-sycl/dpct/helper.hpp +15 -7
  451. data/ext/sources/ggml/src/ggml-sycl/element_wise.cpp +215 -204
  452. data/ext/sources/ggml/src/ggml-sycl/element_wise.hpp +2 -2
  453. data/ext/sources/ggml/src/ggml-sycl/fattn-buffers.cpp +56 -0
  454. data/ext/sources/ggml/src/ggml-sycl/fattn-buffers.hpp +63 -0
  455. data/ext/sources/ggml/src/ggml-sycl/fattn-common.hpp +7 -5
  456. data/ext/sources/ggml/src/ggml-sycl/fattn-tile.cpp +4 -0
  457. data/ext/sources/ggml/src/ggml-sycl/fattn-tile.hpp +76 -168
  458. data/ext/sources/ggml/src/ggml-sycl/fattn-vec.hpp +7 -0
  459. data/ext/sources/ggml/src/ggml-sycl/fattn.cpp +3 -1
  460. data/ext/sources/ggml/src/ggml-sycl/fill.cpp +55 -0
  461. data/ext/sources/ggml/src/ggml-sycl/fill.hpp +5 -0
  462. data/ext/sources/ggml/src/ggml-sycl/gated_delta_net.cpp +69 -31
  463. data/ext/sources/ggml/src/ggml-sycl/gated_delta_net.hpp +1 -0
  464. data/ext/sources/ggml/src/ggml-sycl/gemm.hpp +3 -0
  465. data/ext/sources/ggml/src/ggml-sycl/getrows.cpp +79 -3
  466. data/ext/sources/ggml/src/ggml-sycl/ggml-sycl.cpp +1758 -455
  467. data/ext/sources/ggml/src/ggml-sycl/im2col.cpp +353 -89
  468. data/ext/sources/ggml/src/ggml-sycl/im2col.hpp +5 -3
  469. data/ext/sources/ggml/src/ggml-sycl/mmvq.cpp +1542 -39
  470. data/ext/sources/ggml/src/ggml-sycl/mmvq.hpp +33 -0
  471. data/ext/sources/ggml/src/ggml-sycl/norm.cpp +103 -49
  472. data/ext/sources/ggml/src/ggml-sycl/outprod.cpp +45 -9
  473. data/ext/sources/ggml/src/ggml-sycl/pad.cpp +27 -27
  474. data/ext/sources/ggml/src/ggml-sycl/pool.cpp +185 -0
  475. data/ext/sources/ggml/src/ggml-sycl/pool.hpp +22 -0
  476. data/ext/sources/ggml/src/ggml-sycl/presets.hpp +3 -1
  477. data/ext/sources/ggml/src/ggml-sycl/quants.hpp +71 -0
  478. data/ext/sources/ggml/src/ggml-sycl/set_rows.cpp +17 -3
  479. data/ext/sources/ggml/src/ggml-sycl/softmax.cpp +9 -10
  480. data/ext/sources/ggml/src/ggml-sycl/solve_tri.cpp +172 -0
  481. data/ext/sources/ggml/src/ggml-sycl/solve_tri.hpp +8 -0
  482. data/ext/sources/ggml/src/ggml-sycl/ssm_conv.cpp +6 -1
  483. data/ext/sources/ggml/src/ggml-sycl/ssm_scan.cpp +156 -0
  484. data/ext/sources/ggml/src/ggml-sycl/ssm_scan.hpp +5 -0
  485. data/ext/sources/ggml/src/ggml-sycl/sycl_hw.cpp +62 -10
  486. data/ext/sources/ggml/src/ggml-sycl/sycl_hw.hpp +18 -6
  487. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-tile-instance-dkq512-dv512.cpp +6 -0
  488. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-f16.cpp +1 -0
  489. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q4_0.cpp +1 -0
  490. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q4_1.cpp +1 -0
  491. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q5_0.cpp +1 -0
  492. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q5_1.cpp +1 -0
  493. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-f16-q8_0.cpp +1 -0
  494. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-f16.cpp +1 -0
  495. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q4_0.cpp +1 -0
  496. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q4_1.cpp +1 -0
  497. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q5_0.cpp +1 -0
  498. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q5_1.cpp +1 -0
  499. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_0-q8_0.cpp +1 -0
  500. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-f16.cpp +1 -0
  501. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q4_0.cpp +1 -0
  502. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q4_1.cpp +1 -0
  503. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q5_0.cpp +1 -0
  504. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q5_1.cpp +1 -0
  505. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q4_1-q8_0.cpp +1 -0
  506. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-f16.cpp +1 -0
  507. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q4_0.cpp +1 -0
  508. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q4_1.cpp +1 -0
  509. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q5_0.cpp +1 -0
  510. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q5_1.cpp +1 -0
  511. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_0-q8_0.cpp +1 -0
  512. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-f16.cpp +1 -0
  513. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q4_0.cpp +1 -0
  514. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q4_1.cpp +1 -0
  515. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q5_0.cpp +1 -0
  516. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q5_1.cpp +1 -0
  517. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q5_1-q8_0.cpp +1 -0
  518. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-f16.cpp +1 -0
  519. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q4_0.cpp +1 -0
  520. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q4_1.cpp +1 -0
  521. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q5_0.cpp +1 -0
  522. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q5_1.cpp +1 -0
  523. data/ext/sources/ggml/src/ggml-sycl/template-instances/fattn-vec-instance-q8_0-q8_0.cpp +1 -0
  524. data/ext/sources/ggml/src/ggml-sycl/type.hpp +112 -0
  525. data/ext/sources/ggml/src/ggml-sycl/upscale.cpp +410 -0
  526. data/ext/sources/ggml/src/ggml-sycl/upscale.hpp +9 -0
  527. data/ext/sources/ggml/src/ggml-sycl/vecdotq.hpp +242 -45
  528. data/ext/sources/ggml/src/ggml-virtgpu/ggml-backend-buffer.cpp +4 -0
  529. data/ext/sources/ggml/src/ggml-virtgpu/ggml-backend-device.cpp +2 -0
  530. data/ext/sources/ggml/src/ggml-virtgpu/ggml-backend.cpp +2 -0
  531. data/ext/sources/ggml/src/ggml-virtgpu/virtgpu-shm.cpp +1 -0
  532. data/ext/sources/ggml/src/ggml-virtgpu/virtgpu.cpp +1 -0
  533. data/ext/sources/ggml/src/ggml-virtgpu/virtgpu.h +0 -2
  534. data/ext/sources/ggml/src/ggml-vulkan/CMakeLists.txt +16 -0
  535. data/ext/sources/ggml/src/ggml-vulkan/ggml-vulkan.cpp +2843 -700
  536. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/CMakeLists.txt +4 -0
  537. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/col2im_1d.comp +61 -0
  538. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/contig_copy.comp +6 -2
  539. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/conv2d_mm.comp +146 -13
  540. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/conv3d_mm.comp +431 -0
  541. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/copy.comp +3 -1
  542. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/copy_from_quant.comp +1 -1
  543. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/copy_to_quant.comp +25 -1
  544. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs.glsl +88 -0
  545. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs_cm2.glsl +643 -1
  546. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dequant_nvfp4.comp +32 -0
  547. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q1_0.comp +29 -0
  548. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/diag.comp +3 -4
  549. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/dot_product_funcs.glsl +27 -0
  550. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/feature-tests/coopmat2_decode_vector.comp +7 -0
  551. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp +198 -48
  552. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl +60 -59
  553. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp +116 -113
  554. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm2.comp +122 -31
  555. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_dequant.glsl +131 -0
  556. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_mmq_funcs.glsl +203 -0
  557. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/fwht.comp +115 -0
  558. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gated_delta_net.comp +125 -64
  559. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/generic_binary_head.glsl +0 -1
  560. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/generic_unary_head.glsl +21 -19
  561. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/get_rows_back.comp +25 -0
  562. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/glu_head.glsl +29 -1
  563. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/glu_main.glsl +17 -11
  564. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/im2col.comp +76 -54
  565. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/im2col_3d.comp +0 -1
  566. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/l2_norm.comp +4 -7
  567. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/log.comp +0 -1
  568. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec.comp +122 -27
  569. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_iface.glsl +6 -6
  570. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_q2_k.comp +1 -1
  571. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_q4_k.comp +1 -1
  572. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_q5_k.comp +1 -1
  573. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq.comp +22 -24
  574. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq_funcs.glsl +88 -55
  575. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp +42 -40
  576. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp +49 -15
  577. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl +222 -171
  578. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_funcs.glsl +8 -8
  579. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_shmem_types.glsl +24 -9
  580. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/multi_add.comp +0 -1
  581. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/norm.comp +10 -10
  582. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/repeat_back.comp +3 -3
  583. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/roll.comp +3 -3
  584. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/rope_funcs.glsl +5 -2
  585. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/rope_head.glsl +0 -1
  586. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/rope_params.glsl +3 -2
  587. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/snake.comp +49 -0
  588. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/ssm_conv.comp +11 -1
  589. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/tri.comp +3 -4
  590. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/types.glsl +79 -2
  591. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/unary.comp +168 -0
  592. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +282 -211
  593. data/ext/sources/ggml/src/ggml-webgpu/CMakeLists.txt +5 -2
  594. data/ext/sources/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp +2209 -283
  595. data/ext/sources/ggml/src/ggml-webgpu/ggml-webgpu.cpp +2618 -1416
  596. data/ext/sources/ggml/src/ggml-webgpu/pre_wgsl.hpp +37 -7
  597. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/add_id.wgsl +64 -0
  598. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/binary.wgsl +8 -7
  599. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/common_decls.tmpl +90 -95
  600. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/concat.wgsl +19 -1
  601. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/conv2d.wgsl +165 -0
  602. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/{cpy.tmpl.wgsl → cpy.wgsl} +25 -50
  603. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn.wgsl +107 -184
  604. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_quant_staging.tmpl +124 -0
  605. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_tile.wgsl +397 -0
  606. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_blk.wgsl +101 -0
  607. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_reduce.wgsl +84 -0
  608. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_split.wgsl +619 -0
  609. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/gated_delta_net.wgsl +149 -0
  610. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/get_rows.wgsl +204 -78
  611. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/glu.wgsl +155 -0
  612. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/im2col.wgsl +101 -0
  613. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_decls.tmpl +805 -526
  614. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_id.wgsl +195 -0
  615. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_id_gather.wgsl +52 -0
  616. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_id_vec.wgsl +154 -0
  617. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_reg_tile.wgsl +8 -6
  618. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_subgroup_matrix.wgsl +5 -1
  619. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec.wgsl +90 -413
  620. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_acc.tmpl +1553 -0
  621. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_q_acc.tmpl +297 -0
  622. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/quant_inner_loops.tmpl +21 -0
  623. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/quantize_q8.wgsl +178 -0
  624. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/rms_norm_mul.wgsl +152 -0
  625. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/{rope.tmpl.wgsl → rope.wgsl} +71 -142
  626. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/row_norm.wgsl +153 -0
  627. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/scale.wgsl +6 -4
  628. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/set.wgsl +109 -0
  629. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/set_rows.wgsl +2 -3
  630. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/set_rows_quant.wgsl +224 -0
  631. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/{soft_max.tmpl.wgsl → soft_max.wgsl} +106 -206
  632. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/solve_tri.wgsl +121 -0
  633. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/ssm_conv.wgsl +65 -0
  634. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/ssm_scan.wgsl +193 -0
  635. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/unary.wgsl +68 -48
  636. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/upscale.wgsl +240 -0
  637. data/ext/sources/ggml/src/ggml-zdnn/ggml-zdnn.cpp +18 -14
  638. data/ext/sources/ggml/src/ggml-zendnn/CMakeLists.txt +1 -1
  639. data/ext/sources/ggml/src/ggml-zendnn/ggml-zendnn.cpp +244 -10
  640. data/ext/sources/ggml/src/ggml.c +146 -42
  641. data/ext/sources/ggml/src/gguf.cpp +173 -28
  642. data/ext/sources/include/parakeet.h +342 -0
  643. data/ext/sources/include/whisper.h +31 -0
  644. data/ext/sources/media/matmul.png +0 -0
  645. data/ext/sources/src/CMakeLists.txt +23 -0
  646. data/ext/sources/src/parakeet-arch.h +188 -0
  647. data/ext/sources/src/parakeet.cpp +3838 -0
  648. data/ext/sources/src/whisper.cpp +220 -26
  649. data/extsources.rb +26 -10
  650. data/lib/whisper/log_settable.rb +33 -0
  651. data/lib/whisper/model/uri.rb +13 -8
  652. data/lib/whisper/output.rb +74 -0
  653. data/sig/whisper.rbs +417 -62
  654. data/test/helper.rb +2 -0
  655. data/test/jfk_reader/jfk_reader.c +50 -7
  656. data/test/test_callback.rb +1 -0
  657. data/test/test_package.rb +6 -5
  658. data/test/test_parakeet.rb +28 -0
  659. data/test/test_parakeet_callback.rb +107 -0
  660. data/test/test_parakeet_context.rb +116 -0
  661. data/test/test_parakeet_context_params.rb +24 -0
  662. data/test/test_parakeet_model.rb +21 -0
  663. data/test/test_parakeet_params.rb +78 -0
  664. data/test/test_parakeet_segment.rb +42 -0
  665. data/test/test_parakeet_token.rb +73 -0
  666. data/test/test_params.rb +2 -0
  667. data/test/test_vad.rb +9 -0
  668. data/test/test_vad_context.rb +2 -2
  669. data/test/test_vad_segment.rb +1 -1
  670. data/test/test_whisper.rb +24 -6
  671. data/whispercpp.gemspec +2 -2
  672. metadata +263 -304
  673. data/ext/sources/bindings/javascript/CMakeLists.txt +0 -41
  674. data/ext/sources/bindings/javascript/emscripten.cpp +0 -93
  675. data/ext/sources/bindings/javascript/libwhisper.worker.js +0 -1
  676. data/ext/sources/bindings/javascript/package.json +0 -26
  677. data/ext/sources/bindings/javascript/whisper.js +0 -19
  678. data/ext/sources/examples/addon.node/CMakeLists.txt +0 -31
  679. data/ext/sources/examples/addon.node/__test__/whisper.spec.js +0 -133
  680. data/ext/sources/examples/addon.node/addon.cpp +0 -557
  681. data/ext/sources/examples/addon.node/index.js +0 -59
  682. data/ext/sources/examples/addon.node/package.json +0 -16
  683. data/ext/sources/examples/addon.node/vad-example.js +0 -132
  684. data/ext/sources/examples/bench.wasm/CMakeLists.txt +0 -49
  685. data/ext/sources/examples/bench.wasm/emscripten.cpp +0 -87
  686. data/ext/sources/examples/bench.wasm/index-tmpl.html +0 -285
  687. data/ext/sources/examples/coi-serviceworker.js +0 -146
  688. data/ext/sources/examples/command/CMakeLists.txt +0 -10
  689. data/ext/sources/examples/command/command.cpp +0 -802
  690. data/ext/sources/examples/command/commands.txt +0 -9
  691. data/ext/sources/examples/command.wasm/CMakeLists.txt +0 -50
  692. data/ext/sources/examples/command.wasm/emscripten.cpp +0 -327
  693. data/ext/sources/examples/command.wasm/index-tmpl.html +0 -415
  694. data/ext/sources/examples/generate-karaoke.sh +0 -57
  695. data/ext/sources/examples/helpers.js +0 -191
  696. data/ext/sources/examples/livestream.sh +0 -112
  697. data/ext/sources/examples/lsp/CMakeLists.txt +0 -10
  698. data/ext/sources/examples/lsp/lsp.cpp +0 -471
  699. data/ext/sources/examples/lsp/whisper.vim +0 -362
  700. data/ext/sources/examples/python/test_whisper_processor.py +0 -7
  701. data/ext/sources/examples/python/whisper_processor.py +0 -54
  702. data/ext/sources/examples/server/bench.js +0 -29
  703. data/ext/sources/examples/server.py +0 -120
  704. data/ext/sources/examples/stream/CMakeLists.txt +0 -10
  705. data/ext/sources/examples/stream/stream.cpp +0 -437
  706. data/ext/sources/examples/stream.wasm/CMakeLists.txt +0 -49
  707. data/ext/sources/examples/stream.wasm/emscripten.cpp +0 -216
  708. data/ext/sources/examples/stream.wasm/index-tmpl.html +0 -491
  709. data/ext/sources/examples/sycl/CMakeLists.txt +0 -9
  710. data/ext/sources/examples/sycl/build.sh +0 -22
  711. data/ext/sources/examples/sycl/ls-sycl-device.cpp +0 -11
  712. data/ext/sources/examples/sycl/run-whisper.sh +0 -17
  713. data/ext/sources/examples/talk-llama/CMakeLists.txt +0 -48
  714. data/ext/sources/examples/talk-llama/eleven-labs.py +0 -80
  715. data/ext/sources/examples/talk-llama/llama-adapter.cpp +0 -488
  716. data/ext/sources/examples/talk-llama/llama-adapter.h +0 -89
  717. data/ext/sources/examples/talk-llama/llama-arch.cpp +0 -2877
  718. data/ext/sources/examples/talk-llama/llama-arch.h +0 -628
  719. data/ext/sources/examples/talk-llama/llama-batch.cpp +0 -919
  720. data/ext/sources/examples/talk-llama/llama-batch.h +0 -173
  721. data/ext/sources/examples/talk-llama/llama-chat.cpp +0 -896
  722. data/ext/sources/examples/talk-llama/llama-chat.h +0 -71
  723. data/ext/sources/examples/talk-llama/llama-context.cpp +0 -3633
  724. data/ext/sources/examples/talk-llama/llama-context.h +0 -359
  725. data/ext/sources/examples/talk-llama/llama-cparams.cpp +0 -5
  726. data/ext/sources/examples/talk-llama/llama-cparams.h +0 -47
  727. data/ext/sources/examples/talk-llama/llama-ext.h +0 -12
  728. data/ext/sources/examples/talk-llama/llama-grammar.cpp +0 -1464
  729. data/ext/sources/examples/talk-llama/llama-grammar.h +0 -194
  730. data/ext/sources/examples/talk-llama/llama-graph.cpp +0 -2735
  731. data/ext/sources/examples/talk-llama/llama-graph.h +0 -1031
  732. data/ext/sources/examples/talk-llama/llama-hparams.cpp +0 -258
  733. data/ext/sources/examples/talk-llama/llama-hparams.h +0 -353
  734. data/ext/sources/examples/talk-llama/llama-impl.cpp +0 -171
  735. data/ext/sources/examples/talk-llama/llama-impl.h +0 -75
  736. data/ext/sources/examples/talk-llama/llama-io.cpp +0 -15
  737. data/ext/sources/examples/talk-llama/llama-io.h +0 -35
  738. data/ext/sources/examples/talk-llama/llama-kv-cache-iswa.cpp +0 -330
  739. data/ext/sources/examples/talk-llama/llama-kv-cache-iswa.h +0 -137
  740. data/ext/sources/examples/talk-llama/llama-kv-cache.cpp +0 -2285
  741. data/ext/sources/examples/talk-llama/llama-kv-cache.h +0 -389
  742. data/ext/sources/examples/talk-llama/llama-kv-cells.h +0 -533
  743. data/ext/sources/examples/talk-llama/llama-memory-hybrid-iswa.cpp +0 -275
  744. data/ext/sources/examples/talk-llama/llama-memory-hybrid-iswa.h +0 -140
  745. data/ext/sources/examples/talk-llama/llama-memory-hybrid.cpp +0 -268
  746. data/ext/sources/examples/talk-llama/llama-memory-hybrid.h +0 -139
  747. data/ext/sources/examples/talk-llama/llama-memory-recurrent.cpp +0 -1165
  748. data/ext/sources/examples/talk-llama/llama-memory-recurrent.h +0 -182
  749. data/ext/sources/examples/talk-llama/llama-memory.cpp +0 -59
  750. data/ext/sources/examples/talk-llama/llama-memory.h +0 -122
  751. data/ext/sources/examples/talk-llama/llama-mmap.cpp +0 -752
  752. data/ext/sources/examples/talk-llama/llama-mmap.h +0 -73
  753. data/ext/sources/examples/talk-llama/llama-model-loader.cpp +0 -1655
  754. data/ext/sources/examples/talk-llama/llama-model-loader.h +0 -206
  755. data/ext/sources/examples/talk-llama/llama-model-saver.cpp +0 -299
  756. data/ext/sources/examples/talk-llama/llama-model-saver.h +0 -40
  757. data/ext/sources/examples/talk-llama/llama-model.cpp +0 -9056
  758. data/ext/sources/examples/talk-llama/llama-model.h +0 -597
  759. data/ext/sources/examples/talk-llama/llama-quant.cpp +0 -1304
  760. data/ext/sources/examples/talk-llama/llama-quant.h +0 -1
  761. data/ext/sources/examples/talk-llama/llama-sampler.cpp +0 -3885
  762. data/ext/sources/examples/talk-llama/llama-sampler.h +0 -42
  763. data/ext/sources/examples/talk-llama/llama-vocab.cpp +0 -3970
  764. data/ext/sources/examples/talk-llama/llama-vocab.h +0 -187
  765. data/ext/sources/examples/talk-llama/llama.cpp +0 -1194
  766. data/ext/sources/examples/talk-llama/llama.h +0 -1573
  767. data/ext/sources/examples/talk-llama/models/afmoe.cpp +0 -190
  768. data/ext/sources/examples/talk-llama/models/apertus.cpp +0 -125
  769. data/ext/sources/examples/talk-llama/models/arcee.cpp +0 -135
  770. data/ext/sources/examples/talk-llama/models/arctic.cpp +0 -137
  771. data/ext/sources/examples/talk-llama/models/arwkv7.cpp +0 -86
  772. data/ext/sources/examples/talk-llama/models/baichuan.cpp +0 -123
  773. data/ext/sources/examples/talk-llama/models/bailingmoe.cpp +0 -143
  774. data/ext/sources/examples/talk-llama/models/bailingmoe2.cpp +0 -133
  775. data/ext/sources/examples/talk-llama/models/bert.cpp +0 -184
  776. data/ext/sources/examples/talk-llama/models/bitnet.cpp +0 -145
  777. data/ext/sources/examples/talk-llama/models/bloom.cpp +0 -101
  778. data/ext/sources/examples/talk-llama/models/chameleon.cpp +0 -178
  779. data/ext/sources/examples/talk-llama/models/chatglm.cpp +0 -132
  780. data/ext/sources/examples/talk-llama/models/codeshell.cpp +0 -111
  781. data/ext/sources/examples/talk-llama/models/cogvlm.cpp +0 -102
  782. data/ext/sources/examples/talk-llama/models/cohere2-iswa.cpp +0 -134
  783. data/ext/sources/examples/talk-llama/models/command-r.cpp +0 -122
  784. data/ext/sources/examples/talk-llama/models/dbrx.cpp +0 -122
  785. data/ext/sources/examples/talk-llama/models/deci.cpp +0 -135
  786. data/ext/sources/examples/talk-llama/models/deepseek.cpp +0 -142
  787. data/ext/sources/examples/talk-llama/models/deepseek2.cpp +0 -262
  788. data/ext/sources/examples/talk-llama/models/delta-net-base.cpp +0 -445
  789. data/ext/sources/examples/talk-llama/models/dots1.cpp +0 -132
  790. data/ext/sources/examples/talk-llama/models/dream.cpp +0 -105
  791. data/ext/sources/examples/talk-llama/models/ernie4-5-moe.cpp +0 -148
  792. data/ext/sources/examples/talk-llama/models/ernie4-5.cpp +0 -110
  793. data/ext/sources/examples/talk-llama/models/eurobert.cpp +0 -97
  794. data/ext/sources/examples/talk-llama/models/exaone-moe.cpp +0 -145
  795. data/ext/sources/examples/talk-llama/models/exaone.cpp +0 -114
  796. data/ext/sources/examples/talk-llama/models/exaone4.cpp +0 -123
  797. data/ext/sources/examples/talk-llama/models/falcon-h1.cpp +0 -111
  798. data/ext/sources/examples/talk-llama/models/falcon.cpp +0 -120
  799. data/ext/sources/examples/talk-llama/models/gemma-embedding.cpp +0 -116
  800. data/ext/sources/examples/talk-llama/models/gemma.cpp +0 -112
  801. data/ext/sources/examples/talk-llama/models/gemma2-iswa.cpp +0 -128
  802. data/ext/sources/examples/talk-llama/models/gemma3.cpp +0 -155
  803. data/ext/sources/examples/talk-llama/models/gemma3n-iswa.cpp +0 -384
  804. data/ext/sources/examples/talk-llama/models/glm4-moe.cpp +0 -170
  805. data/ext/sources/examples/talk-llama/models/glm4.cpp +0 -157
  806. data/ext/sources/examples/talk-llama/models/gpt2.cpp +0 -105
  807. data/ext/sources/examples/talk-llama/models/gptneox.cpp +0 -144
  808. data/ext/sources/examples/talk-llama/models/granite-hybrid.cpp +0 -195
  809. data/ext/sources/examples/talk-llama/models/granite.cpp +0 -210
  810. data/ext/sources/examples/talk-llama/models/grok.cpp +0 -159
  811. data/ext/sources/examples/talk-llama/models/grovemoe.cpp +0 -139
  812. data/ext/sources/examples/talk-llama/models/hunyuan-dense.cpp +0 -132
  813. data/ext/sources/examples/talk-llama/models/hunyuan-moe.cpp +0 -153
  814. data/ext/sources/examples/talk-llama/models/internlm2.cpp +0 -120
  815. data/ext/sources/examples/talk-llama/models/jais.cpp +0 -86
  816. data/ext/sources/examples/talk-llama/models/jais2.cpp +0 -123
  817. data/ext/sources/examples/talk-llama/models/jamba.cpp +0 -106
  818. data/ext/sources/examples/talk-llama/models/kimi-linear.cpp +0 -381
  819. data/ext/sources/examples/talk-llama/models/lfm2.cpp +0 -196
  820. data/ext/sources/examples/talk-llama/models/llada-moe.cpp +0 -122
  821. data/ext/sources/examples/talk-llama/models/llada.cpp +0 -99
  822. data/ext/sources/examples/talk-llama/models/llama-iswa.cpp +0 -178
  823. data/ext/sources/examples/talk-llama/models/llama.cpp +0 -175
  824. data/ext/sources/examples/talk-llama/models/maincoder.cpp +0 -117
  825. data/ext/sources/examples/talk-llama/models/mamba-base.cpp +0 -289
  826. data/ext/sources/examples/talk-llama/models/mamba.cpp +0 -54
  827. data/ext/sources/examples/talk-llama/models/mimo2-iswa.cpp +0 -129
  828. data/ext/sources/examples/talk-llama/models/minicpm3.cpp +0 -200
  829. data/ext/sources/examples/talk-llama/models/minimax-m2.cpp +0 -123
  830. data/ext/sources/examples/talk-llama/models/mistral3.cpp +0 -160
  831. data/ext/sources/examples/talk-llama/models/models.h +0 -704
  832. data/ext/sources/examples/talk-llama/models/modern-bert.cpp +0 -109
  833. data/ext/sources/examples/talk-llama/models/mpt.cpp +0 -126
  834. data/ext/sources/examples/talk-llama/models/nemotron-h.cpp +0 -162
  835. data/ext/sources/examples/talk-llama/models/nemotron.cpp +0 -122
  836. data/ext/sources/examples/talk-llama/models/neo-bert.cpp +0 -104
  837. data/ext/sources/examples/talk-llama/models/olmo.cpp +0 -121
  838. data/ext/sources/examples/talk-llama/models/olmo2.cpp +0 -150
  839. data/ext/sources/examples/talk-llama/models/olmoe.cpp +0 -124
  840. data/ext/sources/examples/talk-llama/models/openai-moe-iswa.cpp +0 -127
  841. data/ext/sources/examples/talk-llama/models/openelm.cpp +0 -124
  842. data/ext/sources/examples/talk-llama/models/orion.cpp +0 -123
  843. data/ext/sources/examples/talk-llama/models/paddleocr.cpp +0 -122
  844. data/ext/sources/examples/talk-llama/models/pangu-embedded.cpp +0 -121
  845. data/ext/sources/examples/talk-llama/models/phi2.cpp +0 -121
  846. data/ext/sources/examples/talk-llama/models/phi3.cpp +0 -152
  847. data/ext/sources/examples/talk-llama/models/plamo.cpp +0 -110
  848. data/ext/sources/examples/talk-llama/models/plamo2.cpp +0 -320
  849. data/ext/sources/examples/talk-llama/models/plamo3.cpp +0 -128
  850. data/ext/sources/examples/talk-llama/models/plm.cpp +0 -169
  851. data/ext/sources/examples/talk-llama/models/qwen.cpp +0 -108
  852. data/ext/sources/examples/talk-llama/models/qwen2.cpp +0 -126
  853. data/ext/sources/examples/talk-llama/models/qwen2moe.cpp +0 -151
  854. data/ext/sources/examples/talk-llama/models/qwen2vl.cpp +0 -117
  855. data/ext/sources/examples/talk-llama/models/qwen3.cpp +0 -120
  856. data/ext/sources/examples/talk-llama/models/qwen35.cpp +0 -381
  857. data/ext/sources/examples/talk-llama/models/qwen35moe.cpp +0 -422
  858. data/ext/sources/examples/talk-llama/models/qwen3moe.cpp +0 -131
  859. data/ext/sources/examples/talk-llama/models/qwen3next.cpp +0 -525
  860. data/ext/sources/examples/talk-llama/models/qwen3vl-moe.cpp +0 -140
  861. data/ext/sources/examples/talk-llama/models/qwen3vl.cpp +0 -132
  862. data/ext/sources/examples/talk-llama/models/refact.cpp +0 -94
  863. data/ext/sources/examples/talk-llama/models/rnd1.cpp +0 -126
  864. data/ext/sources/examples/talk-llama/models/rwkv6-base.cpp +0 -164
  865. data/ext/sources/examples/talk-llama/models/rwkv6.cpp +0 -94
  866. data/ext/sources/examples/talk-llama/models/rwkv6qwen2.cpp +0 -86
  867. data/ext/sources/examples/talk-llama/models/rwkv7-base.cpp +0 -137
  868. data/ext/sources/examples/talk-llama/models/rwkv7.cpp +0 -90
  869. data/ext/sources/examples/talk-llama/models/seed-oss.cpp +0 -124
  870. data/ext/sources/examples/talk-llama/models/smallthinker.cpp +0 -126
  871. data/ext/sources/examples/talk-llama/models/smollm3.cpp +0 -128
  872. data/ext/sources/examples/talk-llama/models/stablelm.cpp +0 -146
  873. data/ext/sources/examples/talk-llama/models/starcoder.cpp +0 -100
  874. data/ext/sources/examples/talk-llama/models/starcoder2.cpp +0 -121
  875. data/ext/sources/examples/talk-llama/models/step35-iswa.cpp +0 -165
  876. data/ext/sources/examples/talk-llama/models/t5-dec.cpp +0 -166
  877. data/ext/sources/examples/talk-llama/models/t5-enc.cpp +0 -96
  878. data/ext/sources/examples/talk-llama/models/wavtokenizer-dec.cpp +0 -149
  879. data/ext/sources/examples/talk-llama/models/xverse.cpp +0 -108
  880. data/ext/sources/examples/talk-llama/prompts/talk-alpaca.txt +0 -23
  881. data/ext/sources/examples/talk-llama/speak +0 -40
  882. data/ext/sources/examples/talk-llama/speak.bat +0 -1
  883. data/ext/sources/examples/talk-llama/speak.ps1 +0 -14
  884. data/ext/sources/examples/talk-llama/talk-llama.cpp +0 -813
  885. data/ext/sources/examples/talk-llama/unicode-data.cpp +0 -7034
  886. data/ext/sources/examples/talk-llama/unicode-data.h +0 -20
  887. data/ext/sources/examples/talk-llama/unicode.cpp +0 -1103
  888. data/ext/sources/examples/talk-llama/unicode.h +0 -111
  889. data/ext/sources/examples/wchess/CMakeLists.txt +0 -10
  890. data/ext/sources/examples/wchess/libwchess/CMakeLists.txt +0 -19
  891. data/ext/sources/examples/wchess/libwchess/Chessboard.cpp +0 -803
  892. data/ext/sources/examples/wchess/libwchess/Chessboard.h +0 -33
  893. data/ext/sources/examples/wchess/libwchess/WChess.cpp +0 -193
  894. data/ext/sources/examples/wchess/libwchess/WChess.h +0 -63
  895. data/ext/sources/examples/wchess/libwchess/test-chessboard.cpp +0 -117
  896. data/ext/sources/examples/wchess/wchess.cmd/CMakeLists.txt +0 -8
  897. data/ext/sources/examples/wchess/wchess.cmd/wchess.cmd.cpp +0 -253
  898. data/ext/sources/examples/whisper.wasm/CMakeLists.txt +0 -50
  899. data/ext/sources/examples/whisper.wasm/emscripten.cpp +0 -118
  900. data/ext/sources/examples/whisper.wasm/index-tmpl.html +0 -659
  901. data/ext/sources/ggml/src/ggml-cuda/template-instances/generate_cu_files.py +0 -99
  902. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-msg.h +0 -155
  903. data/ext/sources/ggml/src/ggml-hexagon/op-desc.h +0 -153
  904. data/ext/sources/ggml/src/ggml-opencl/kernels/embed_kernel.py +0 -26
  905. data/ext/sources/ggml/src/ggml-openvino/openvino/pass/eliminate_zp.cpp +0 -123
  906. data/ext/sources/ggml/src/ggml-openvino/openvino/pass/eliminate_zp.h +0 -17
  907. data/ext/sources/ggml/src/ggml-virtgpu/regenerate_remoting.py +0 -333
  908. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/abs.comp +0 -21
  909. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/ceil.comp +0 -22
  910. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/clamp.comp +0 -17
  911. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/cos.comp +0 -17
  912. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/elu.comp +0 -27
  913. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/exp.comp +0 -21
  914. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/floor.comp +0 -22
  915. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu.comp +0 -25
  916. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu_erf.comp +0 -39
  917. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu_quick.comp +0 -23
  918. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/hardsigmoid.comp +0 -22
  919. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/hardswish.comp +0 -22
  920. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/leaky_relu.comp +0 -22
  921. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/neg.comp +0 -20
  922. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/relu.comp +0 -21
  923. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/round.comp +0 -29
  924. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/rte.glsl +0 -5
  925. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sgn.comp +0 -21
  926. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sigmoid.comp +0 -20
  927. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/silu.comp +0 -22
  928. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sin.comp +0 -17
  929. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/softplus.comp +0 -23
  930. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sqrt.comp +0 -17
  931. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/square.comp +0 -17
  932. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/step.comp +0 -22
  933. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/tanh.comp +0 -20
  934. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/trunc.comp +0 -22
  935. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/xielu.comp +0 -35
  936. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/embed_wgsl.py +0 -182
  937. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/glu.tmpl.wgsl +0 -323
  938. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat.wgsl +0 -718
  939. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/rms_norm.wgsl +0 -123
  940. data/ext/sources/tests/CMakeLists.txt +0 -112
  941. data/ext/sources/tests/earnings21/eval.mk +0 -58
  942. data/ext/sources/tests/earnings21/eval.py +0 -68
  943. data/ext/sources/tests/earnings21/normalizers/__init__.py +0 -2
  944. data/ext/sources/tests/earnings21/normalizers/basic.py +0 -80
  945. data/ext/sources/tests/earnings21/normalizers/english.json +0 -1741
  946. data/ext/sources/tests/earnings21/normalizers/english.py +0 -550
  947. data/ext/sources/tests/earnings21/requirements.txt +0 -6
  948. data/ext/sources/tests/en-0-ref.txt +0 -1
  949. data/ext/sources/tests/en-1-ref.txt +0 -1
  950. data/ext/sources/tests/en-2-ref.txt +0 -1
  951. data/ext/sources/tests/es-0-ref.txt +0 -1
  952. data/ext/sources/tests/librispeech/eval.mk +0 -39
  953. data/ext/sources/tests/librispeech/eval.py +0 -47
  954. data/ext/sources/tests/librispeech/normalizers/__init__.py +0 -2
  955. data/ext/sources/tests/librispeech/normalizers/basic.py +0 -80
  956. data/ext/sources/tests/librispeech/normalizers/english.json +0 -1741
  957. data/ext/sources/tests/librispeech/normalizers/english.py +0 -550
  958. data/ext/sources/tests/librispeech/requirements.txt +0 -6
  959. data/ext/sources/tests/run-tests.sh +0 -130
  960. data/ext/sources/tests/test-c.c +0 -3
  961. data/ext/sources/tests/test-vad-full.cpp +0 -56
  962. data/ext/sources/tests/test-vad.cpp +0 -83
  963. data/ext/sources/tests/test-whisper.js +0 -58
  964. data/lib/whisper/context.rb +0 -15
  965. data/lib/whisper/segment.rb +0 -58
@@ -0,0 +1,3178 @@
1
+ #include "rvv_kernels.h"
2
+
3
+ #include "common.h"
4
+ #include "ggml.h"
5
+ #include "ops.h"
6
+ #include "string.h"
7
+
8
+ #include <algorithm>
9
+ #include <cmath>
10
+ #include <cstdint>
11
+ #include <stdexcept>
12
+
13
+ #if !defined(__riscv_v) || !defined(__riscv_v_intrinsic)
14
+ # error "riscv v extension or v_intrinsic not enabled"
15
+ #else
16
+ # include <riscv_vector.h>
17
+ #endif
18
+
19
+ #if !defined(__riscv_zfh)
20
+ # error "riscv zfh extension not enabled"
21
+ #endif
22
+
23
+ #if defined(__GNUC__)
24
+ # pragma GCC diagnostic ignored "-Woverlength-strings"
25
+ # pragma GCC diagnostic ignored "-Wcast-qual"
26
+ # pragma GCC diagnostic ignored "-Wunused-parameter"
27
+ #endif
28
+
29
+ namespace spacemit_kernels::rvv {
30
+
31
+ namespace {
32
+
33
+ auto align_up(size_t value, size_t alignment) {
34
+ return (value + alignment - 1) / alignment * alignment;
35
+ }
36
+
37
+ static inline bool flash_attn_ext_supported_d_vlen1024_vf16(int64_t d) {
38
+ return d > 0 && d <= 128;
39
+ }
40
+
41
+ static inline bool flash_attn_ext_supported_shape_vlen1024_vf16(int64_t DK, int64_t DV) {
42
+ return flash_attn_ext_supported_d_vlen1024_vf16(DK) && flash_attn_ext_supported_d_vlen1024_vf16(DV);
43
+ }
44
+
45
+ static inline float reduce_sum_f32m4_vlen1024(vfloat32m4_t v, size_t vl) {
46
+ vfloat32m1_t s_v = __riscv_vfmv_v_f_f32m1(0.0f, 1);
47
+ s_v = __riscv_vfredusum_vs_f32m4_f32m1(v, s_v, vl);
48
+ return __riscv_vfmv_f_s_f32m1_f32(s_v);
49
+ }
50
+
51
+ static inline float reduce_sum_f32m2_vlen1024(vfloat32m2_t v, size_t vl) {
52
+ vfloat32m1_t s_v = __riscv_vfmv_v_f_f32m1(0.0f, 1);
53
+ s_v = __riscv_vfredusum_vs_f32m2_f32m1(v, s_v, vl);
54
+ return __riscv_vfmv_f_s_f32m1_f32(s_v);
55
+ }
56
+
57
+ // Adapted from ggml_v_expf_m2 in vec.h. This is accurate enough for softmax.
58
+ static inline vfloat32m2_t rvv_expf_approx_f32m2(vfloat32m2_t x, size_t vl) {
59
+ const vfloat32m2_t r = __riscv_vfmv_v_f_f32m2(0x1.8p23f, vl);
60
+ const vfloat32m2_t z = __riscv_vfmacc_vf_f32m2(r, 0x1.715476p+0f, x, vl);
61
+ const vfloat32m2_t n = __riscv_vfsub_vv_f32m2(z, r, vl);
62
+ const vfloat32m2_t b =
63
+ __riscv_vfnmsac_vf_f32m2(__riscv_vfnmsac_vf_f32m2(x, 0x1.62e4p-1f, n, vl), 0x1.7f7d1cp-20f, n, vl);
64
+ const vuint32m2_t e = __riscv_vsll_vx_u32m2(__riscv_vreinterpret_v_f32m2_u32m2(z), 23, vl);
65
+ const vfloat32m2_t k = __riscv_vreinterpret_v_u32m2_f32m2(__riscv_vadd_vx_u32m2(e, 0x3f800000, vl));
66
+ const vbool16_t c = __riscv_vmfgt_vf_f32m2_b16(__riscv_vfabs_v_f32m2(n, vl), 126.0f, vl);
67
+ const vfloat32m2_t u = __riscv_vfmul_vv_f32m2(b, b, vl);
68
+ const vfloat32m2_t j = __riscv_vfmacc_vv_f32m2(
69
+ __riscv_vfmul_vf_f32m2(b, 0x1.ffffecp-1f, vl),
70
+ __riscv_vfmacc_vv_f32m2(
71
+ __riscv_vfmacc_vf_f32m2(__riscv_vfmv_v_f_f32m2(0x1.fffdb6p-2f, vl), 0x1.555e66p-3f, b, vl),
72
+ __riscv_vfmacc_vf_f32m2(__riscv_vfmv_v_f_f32m2(0x1.573e2ep-5f, vl), 0x1.0e4020p-7f, b, vl), u, vl),
73
+ u, vl);
74
+
75
+ if (!__riscv_vcpop_m_b16(c, vl)) {
76
+ return __riscv_vfmacc_vv_f32m2(k, j, k, vl);
77
+ }
78
+
79
+ const vbool16_t dm = __riscv_vmfle_vf_f32m2_b16(n, 0.0f, vl);
80
+ const vuint32m2_t d = __riscv_vmerge_vxm_u32m2(__riscv_vmv_v_x_u32m2(0, vl), 0x82000000, dm, vl);
81
+ const vfloat32m2_t s1 = __riscv_vreinterpret_v_u32m2_f32m2(__riscv_vadd_vx_u32m2(d, 0x7f000000, vl));
82
+ const vfloat32m2_t s2 = __riscv_vreinterpret_v_u32m2_f32m2(__riscv_vsub_vv_u32m2(e, d, vl));
83
+ const vfloat32m2_t r1 =
84
+ __riscv_vmerge_vvm_f32m2(__riscv_vfmacc_vv_f32m2(k, k, j, vl),
85
+ __riscv_vfmul_vv_f32m2(__riscv_vfmacc_vv_f32m2(s2, s2, j, vl), s1, vl), c, vl);
86
+ return __riscv_vmerge_vvm_f32m2(r1, __riscv_vfmul_vv_f32m2(s1, s1, vl),
87
+ __riscv_vmfgt_vf_f32m2_b16(__riscv_vfabs_v_f32m2(n, vl), 192.0f, vl), vl);
88
+ }
89
+
90
+ static inline vfloat32m2_t rvv_tanh_approx_f32m2(vfloat32m2_t x, size_t vl) {
91
+ const vfloat32m2_t abs_x = __riscv_vfabs_v_f32m2(x, vl);
92
+ const vfloat32m2_t neg_2_abs = __riscv_vfmul_vf_f32m2(abs_x, -2.0f, vl);
93
+ const vfloat32m2_t exp_term = rvv_expf_approx_f32m2(neg_2_abs, vl);
94
+ const vfloat32m2_t numerator = __riscv_vfsub_vf_f32m2(exp_term, 1.0f, vl);
95
+ const vfloat32m2_t denominator = __riscv_vfadd_vf_f32m2(exp_term, 1.0f, vl);
96
+ const vfloat32m2_t tanh_abs = __riscv_vfneg_v_f32m2(__riscv_vfdiv_vv_f32m2(numerator, denominator, vl), vl);
97
+ const vbool16_t neg_mask = __riscv_vmflt_vf_f32m2_b16(x, 0.0f, vl);
98
+ const vfloat32m2_t tanh_neg = __riscv_vfneg_v_f32m2(tanh_abs, vl);
99
+ return __riscv_vmerge_vvm_f32m2(tanh_abs, tanh_neg, neg_mask, vl);
100
+ }
101
+
102
+ static void rvv_softcap_tanh_inplace_f32(float * dst, int64_t dst_stride, int64_t tile_rows, int64_t n, float softcap) {
103
+ for (int tq = 0; tq < tile_rows; ++tq, dst += dst_stride) {
104
+ float * dst_row = dst;
105
+ int64_t remaining = n;
106
+ while (remaining > 0) {
107
+ const size_t vl = __riscv_vsetvl_e32m2(remaining);
108
+ vfloat32m2_t v = __riscv_vle32_v_f32m2(dst_row, vl);
109
+ v = rvv_tanh_approx_f32m2(v, vl);
110
+ v = __riscv_vfmul_vf_f32m2(v, softcap, vl);
111
+ __riscv_vse32_v_f32m2(dst_row, v, vl);
112
+ dst_row += vl;
113
+ remaining -= vl;
114
+ }
115
+ }
116
+ }
117
+
118
+ static inline float rvv_softmax_exp_inplace_f32(float * dst, int64_t n, float max_value) {
119
+ float row_sum = 0.0f;
120
+ while (n > 0) {
121
+ const size_t vl = __riscv_vsetvl_e32m2(n);
122
+ vfloat32m2_t v = __riscv_vle32_v_f32m2(dst, vl);
123
+ v = __riscv_vfsub_vf_f32m2(v, max_value, vl);
124
+ v = rvv_expf_approx_f32m2(v, vl);
125
+ __riscv_vse32_v_f32m2(dst, v, vl);
126
+ row_sum += reduce_sum_f32m2_vlen1024(v, vl);
127
+ dst += vl;
128
+ n -= vl;
129
+ }
130
+ return row_sum;
131
+ }
132
+
133
+ static inline float rvv_add_max_inplace_f32(float * dst, const float * src, int64_t n) {
134
+ float max_val = -INFINITY;
135
+ while (n > 0) {
136
+ const size_t vl = __riscv_vsetvl_e32m4(n);
137
+ vfloat32m4_t vdst = __riscv_vle32_v_f32m4(dst, vl);
138
+ vfloat32m4_t vsrc = __riscv_vle32_v_f32m4(src, vl);
139
+ vdst = __riscv_vfadd_vv_f32m4(vdst, vsrc, vl);
140
+ __riscv_vse32_v_f32m4(dst, vdst, vl);
141
+
142
+ vfloat32m1_t seed = __riscv_vfmv_v_f_f32m1(max_val, 1);
143
+ seed = __riscv_vfredmax_vs_f32m4_f32m1(vdst, seed, vl);
144
+ max_val = __riscv_vfmv_f_s_f32m1_f32(seed);
145
+
146
+ dst += vl;
147
+ src += vl;
148
+ n -= vl;
149
+ }
150
+ return max_val;
151
+ }
152
+
153
+ static inline float rvv_softcap_add_max_inplace_f32(float * dst, const float * src, int64_t n, float softcap) {
154
+ if (softcap == 0.0f) {
155
+ return rvv_add_max_inplace_f32(dst, src, n);
156
+ }
157
+
158
+ float max_val = -INFINITY;
159
+ while (n > 0) {
160
+ const size_t vl = __riscv_vsetvl_e32m2(n);
161
+ vfloat32m2_t vdst = __riscv_vle32_v_f32m2(dst, vl);
162
+ vfloat32m2_t vsrc = __riscv_vle32_v_f32m2(src, vl);
163
+ vdst = rvv_tanh_approx_f32m2(vdst, vl);
164
+ vdst = __riscv_vfmul_vf_f32m2(vdst, softcap, vl);
165
+ vdst = __riscv_vfadd_vv_f32m2(vdst, vsrc, vl);
166
+ __riscv_vse32_v_f32m2(dst, vdst, vl);
167
+
168
+ vfloat32m1_t seed = __riscv_vfmv_v_f_f32m1(max_val, 1);
169
+ seed = __riscv_vfredmax_vs_f32m2_f32m1(vdst, seed, vl);
170
+ max_val = __riscv_vfmv_f_s_f32m1_f32(seed);
171
+
172
+ dst += vl;
173
+ src += vl;
174
+ n -= vl;
175
+ }
176
+ return max_val;
177
+ }
178
+
179
+ static inline void rvv_zero_f32(float * dst, int64_t n) {
180
+ while (n > 0) {
181
+ const size_t vl = __riscv_vsetvl_e32m4(n);
182
+ const vfloat32m4_t z = __riscv_vfmv_v_f_f32m4(0.0f, vl);
183
+ __riscv_vse32_v_f32m4(dst, z, vl);
184
+ dst += vl;
185
+ n -= vl;
186
+ }
187
+ }
188
+
189
+ static inline void rvv_scale_f32(float * dst, float scale, int64_t n) {
190
+ while (n > 0) {
191
+ const size_t vl = __riscv_vsetvl_e32m4(n);
192
+ vfloat32m4_t v = __riscv_vle32_v_f32m4(dst, vl);
193
+ v = __riscv_vfmul_vf_f32m4(v, scale, vl);
194
+ __riscv_vse32_v_f32m4(dst, v, vl);
195
+ dst += vl;
196
+ n -= vl;
197
+ }
198
+ }
199
+
200
+ static inline void rvv_add_inplace_f32(float * dst,
201
+ int64_t dst_stride,
202
+ const float * src,
203
+ int64_t src_stride,
204
+ int64_t tile_rows,
205
+ int64_t n) {
206
+ for (int tq = 0; tq < tile_rows; ++tq, dst += dst_stride, src += src_stride) {
207
+ int64_t remaining = n;
208
+ float * dst_row = dst;
209
+ const float * src_row = src;
210
+ while (remaining > 0) {
211
+ const size_t vl = __riscv_vsetvl_e32m4(remaining);
212
+ vfloat32m4_t vdst = __riscv_vle32_v_f32m4(dst_row, vl);
213
+ vfloat32m4_t vsrc = __riscv_vle32_v_f32m4(src_row, vl);
214
+ vdst = __riscv_vfadd_vv_f32m4(vdst, vsrc, vl);
215
+ __riscv_vse32_v_f32m4(dst_row, vdst, vl);
216
+ dst_row += vl;
217
+ src_row += vl;
218
+ remaining -= vl;
219
+ }
220
+ }
221
+ }
222
+
223
+ static inline float rvv_max_f32(const float * src, int64_t n) {
224
+ float max_val = -INFINITY;
225
+ while (n > 0) {
226
+ const size_t vl = __riscv_vsetvl_e32m4(n);
227
+ const vfloat32m4_t v = __riscv_vle32_v_f32m4(src, vl);
228
+ vfloat32m1_t seed = __riscv_vfmv_v_f_f32m1(max_val, 1);
229
+ seed = __riscv_vfredmax_vs_f32m4_f32m1(v, seed, vl);
230
+ max_val = __riscv_vfmv_f_s_f32m1_f32(seed);
231
+ src += vl;
232
+ n -= vl;
233
+ }
234
+ return max_val;
235
+ }
236
+
237
+ static void rvv_pack_f32_as_scaled_f16(void * dst,
238
+ int64_t dst_row_stride,
239
+ const void * src,
240
+ int64_t src_row_stride,
241
+ int64_t tile_rows,
242
+ int64_t n,
243
+ float scale) {
244
+ for (int tq = 0; tq < tile_rows; ++tq) {
245
+ const float * row_ptr = (const float *) ((const char *) src + tq * src_row_stride);
246
+ _Float16 * dst_row_ptr = (_Float16 *) ((char *) dst + tq * dst_row_stride);
247
+ int64_t remaining = n;
248
+ while (remaining > 0) {
249
+ const size_t vl = __riscv_vsetvl_e32m4(remaining);
250
+ vfloat32m4_t v32 = __riscv_vle32_v_f32m4(row_ptr, vl);
251
+ v32 = __riscv_vfmul_vf_f32m4(v32, scale, vl);
252
+ const vfloat16m2_t v16 = __riscv_vfncvt_f_f_w_f16m2(v32, vl);
253
+ __riscv_vse16_v_f16m2(dst_row_ptr, v16, vl);
254
+ dst_row_ptr += vl;
255
+ row_ptr += vl;
256
+ remaining -= vl;
257
+ }
258
+ }
259
+ }
260
+
261
+ static void rvv_pack_scaled_f16_as_f32(void * dst,
262
+ int64_t dst_row_stride,
263
+ const void * src,
264
+ int64_t src_row_stride,
265
+ int64_t tile_rows,
266
+ int64_t n,
267
+ float scale) {
268
+ for (int tq = 0; tq < tile_rows; ++tq) {
269
+ const _Float16 * row_ptr = (const _Float16 *) ((const char *) src + tq * src_row_stride);
270
+ float * dst_row_ptr = (float *) ((char *) dst + tq * dst_row_stride);
271
+ int64_t remaining = n;
272
+ while (remaining > 0) {
273
+ const size_t vl = __riscv_vsetvl_e16m2(remaining);
274
+ const vfloat16m2_t v16 = __riscv_vle16_v_f16m2(row_ptr, vl);
275
+ vfloat32m4_t v32 = __riscv_vfwcvt_f_f_v_f32m4(v16, vl);
276
+ v32 = __riscv_vfmul_vf_f32m4(v32, scale, vl);
277
+ __riscv_vse32_v_f32m4(dst_row_ptr, v32, vl);
278
+ dst_row_ptr += vl;
279
+ row_ptr += vl;
280
+ remaining -= vl;
281
+ }
282
+ }
283
+ }
284
+
285
+ static void rvv_pack_scaled_f32_as_f32(void * dst,
286
+ int64_t dst_row_stride,
287
+ const void * src,
288
+ int64_t src_row_stride,
289
+ int64_t tile_rows,
290
+ int64_t n,
291
+ float * scale) {
292
+ for (int tq = 0; tq < tile_rows; ++tq) {
293
+ const float * row_ptr = (const float *) ((const char *) src + tq * src_row_stride);
294
+ float * dst_row_ptr = (float *) ((char *) dst + tq * dst_row_stride);
295
+ int64_t remaining = n;
296
+ while (remaining > 0) {
297
+ const size_t vl = __riscv_vsetvl_e32m4(remaining);
298
+ vfloat32m4_t v32 = __riscv_vle32_v_f32m4(row_ptr, vl);
299
+ v32 = __riscv_vfmul_vf_f32m4(v32, scale[tq], vl);
300
+ __riscv_vse32_v_f32m4(dst_row_ptr, v32, vl);
301
+ dst_row_ptr += vl;
302
+ row_ptr += vl;
303
+ remaining -= vl;
304
+ }
305
+ }
306
+ }
307
+
308
+ static inline void rvv_transposed_s32_mn_to_nm(int8_t * dst,
309
+ int64_t n_dst_stride,
310
+ int8_t * src,
311
+ int64_t m_src_stride,
312
+ int64_t m,
313
+ int64_t n) {
314
+ int8_t * in = src;
315
+ int8_t * out = dst;
316
+
317
+ __asm__ volatile(
318
+ "vsetvli t0, zero, e32, m1, tu, mu \n\t"
319
+ "mul t3, t0, %[os0] \n\t"
320
+ "srli t2, %[isz0], 3 \n\t"
321
+ "blez t2, M1%= \n\t"
322
+
323
+ "LOOP_M8%=: \n\t"
324
+ "addi a1, %[dst], 0 \n\t"
325
+ "addi s1, %[src], 0 \n\t"
326
+ "add s2, %[src], %[is0] \n\t"
327
+ "add s3, s2, %[is0] \n\t"
328
+ "add s4, s3, %[is0] \n\t"
329
+ "add s5, s4, %[is0] \n\t"
330
+ "add s6, s5, %[is0] \n\t"
331
+ "add s7, s6, %[is0] \n\t"
332
+ "add s8, s7, %[is0] \n\t"
333
+ "addi t1, %[isz1], 0 \n\t"
334
+
335
+ "LOOP_M8N%=: \n\t"
336
+ "vsetvli t0, t1, e32, m1, tu, mu \n\t"
337
+ "sub t1, t1, t0 \n\t"
338
+ "vle32.v v0, (s1) \n\t"
339
+ "sh2add s1, t0, s1 \n\t"
340
+ "vle32.v v1, (s2) \n\t"
341
+ "sh2add s2, t0, s2 \n\t"
342
+ "vle32.v v2, (s3) \n\t"
343
+ "sh2add s3, t0, s3 \n\t"
344
+ "vle32.v v3, (s4) \n\t"
345
+ "sh2add s4, t0, s4 \n\t"
346
+ "vle32.v v4, (s5) \n\t"
347
+ "sh2add s5, t0, s5 \n\t"
348
+ "vle32.v v5, (s6) \n\t"
349
+ "sh2add s6, t0, s6 \n\t"
350
+ "vle32.v v6, (s7) \n\t"
351
+ "sh2add s7, t0, s7 \n\t"
352
+ "vle32.v v7, (s8) \n\t"
353
+ "sh2add s8, t0, s8 \n\t"
354
+ "vssseg8e32.v v0, (a1), %[os0] \n\t"
355
+ "add a1, a1, t3 \n\t"
356
+ "bnez t1, LOOP_M8N%= \n\t"
357
+ "sh3add %[src], %[is0], %[src] \n\t"
358
+ "addi %[dst], %[dst], 32 \n\t"
359
+ "addi t2, t2, -1 \n\t"
360
+ "bnez t2, LOOP_M8%= \n\t"
361
+
362
+ "M1%=: \n\t"
363
+ "andi t2, %[isz0], 7 \n\t"
364
+ "blez t2, END%= \n\t"
365
+
366
+ "LOOP_M1%=: \n\t"
367
+ "addi a1, %[dst], 0 \n\t"
368
+ "addi s1, %[src], 0 \n\t"
369
+ "addi t1, %[isz1], 0 \n\t"
370
+
371
+ "LOOP_M1N%=: \n\t"
372
+ "vsetvli t0, t1, e32, m1, tu, mu \n\t"
373
+ "sub t1, t1, t0 \n\t"
374
+ "vle32.v v0, (s1) \n\t"
375
+ "sh2add s1, t0, s1 \n\t"
376
+ "vsse32.v v0, (a1), %[os0] \n\t"
377
+ "add a1, a1, t3 \n\t"
378
+ "bnez t1, LOOP_M1N%= \n\t"
379
+ "add %[src], %[is0], %[src] \n\t"
380
+ "addi %[dst], %[dst], 4 \n\t"
381
+ "addi t2, t2, -1 \n\t"
382
+ "bnez t2, LOOP_M1%= \n\t"
383
+ "END%=: \n\t"
384
+
385
+ : [src] "+r"(in), [dst] "+r"(out), [isz0] "+r"(m)
386
+ : [isz1] "r"(n), [is0] "r"(m_src_stride), [os0] "r"(n_dst_stride)
387
+ : "cc", "t0", "t1", "t2", "t3", "s1", "s2", "s3", "s4", "s5", "s6", "s7", "s8", "a1");
388
+ }
389
+
390
+ static inline void rvv_transposed_s16_mn_to_nm(int8_t * dst,
391
+ int64_t n_dst_stride,
392
+ int8_t * src,
393
+ int64_t m_src_stride,
394
+ int64_t m,
395
+ int64_t n) {
396
+ int8_t * in = src;
397
+ int8_t * out = dst;
398
+
399
+ __asm__ volatile(
400
+ "vsetvli t0, zero, e16, m1, tu, mu \n\t"
401
+ "mul t3, t0, %[os0] \n\t"
402
+ "srli t2, %[isz0], 3 \n\t"
403
+ "blez t2, M1%= \n\t"
404
+
405
+ "LOOP_M8%=: \n\t"
406
+ "addi a1, %[dst], 0 \n\t"
407
+ "addi s1, %[src], 0 \n\t"
408
+ "add s2, %[src], %[is0] \n\t"
409
+ "add s3, s2, %[is0] \n\t"
410
+ "add s4, s3, %[is0] \n\t"
411
+ "add s5, s4, %[is0] \n\t"
412
+ "add s6, s5, %[is0] \n\t"
413
+ "add s7, s6, %[is0] \n\t"
414
+ "add s8, s7, %[is0] \n\t"
415
+ "addi t1, %[isz1], 0 \n\t"
416
+
417
+ "LOOP_M8N%=: \n\t"
418
+ "vsetvli t0, t1, e16, m1, tu, mu \n\t"
419
+ "sub t1, t1, t0 \n\t"
420
+ "vle16.v v0, (s1) \n\t"
421
+ "sh1add s1, t0, s1 \n\t"
422
+ "vle16.v v1, (s2) \n\t"
423
+ "sh1add s2, t0, s2 \n\t"
424
+ "vle16.v v2, (s3) \n\t"
425
+ "sh1add s3, t0, s3 \n\t"
426
+ "vle16.v v3, (s4) \n\t"
427
+ "sh1add s4, t0, s4 \n\t"
428
+ "vle16.v v4, (s5) \n\t"
429
+ "sh1add s5, t0, s5 \n\t"
430
+ "vle16.v v5, (s6) \n\t"
431
+ "sh1add s6, t0, s6 \n\t"
432
+ "vle16.v v6, (s7) \n\t"
433
+ "sh1add s7, t0, s7 \n\t"
434
+ "vle16.v v7, (s8) \n\t"
435
+ "sh1add s8, t0, s8 \n\t"
436
+ "vssseg8e16.v v0, (a1), %[os0] \n\t"
437
+ "add a1, a1, t3 \n\t"
438
+ "bnez t1, LOOP_M8N%= \n\t"
439
+ "sh3add %[src], %[is0], %[src] \n\t"
440
+ "addi %[dst], %[dst], 16 \n\t"
441
+ "addi t2, t2, -1 \n\t"
442
+ "bnez t2, LOOP_M8%= \n\t"
443
+
444
+ "M1%=: \n\t"
445
+ "andi t2, %[isz0], 7 \n\t"
446
+ "blez t2, END%= \n\t"
447
+
448
+ "LOOP_M1%=: \n\t"
449
+ "addi a1, %[dst], 0 \n\t"
450
+ "addi s1, %[src], 0 \n\t"
451
+ "addi t1, %[isz1], 0 \n\t"
452
+
453
+ "LOOP_M1N%=: \n\t"
454
+ "vsetvli t0, t1, e16, m1, tu, mu \n\t"
455
+ "sub t1, t1, t0 \n\t"
456
+ "vle16.v v0, (s1) \n\t"
457
+ "sh1add s1, t0, s1 \n\t"
458
+ "vsse16.v v0, (a1), %[os0] \n\t"
459
+ "add a1, a1, t3 \n\t"
460
+ "bnez t1, LOOP_M1N%= \n\t"
461
+ "add %[src], %[is0], %[src] \n\t"
462
+ "addi %[dst], %[dst], 2 \n\t"
463
+ "addi t2, t2, -1 \n\t"
464
+ "bnez t2, LOOP_M1%= \n\t"
465
+ "END%=: \n\t"
466
+
467
+ : [src] "+r"(in), [dst] "+r"(out), [isz0] "+r"(m)
468
+ : [isz1] "r"(n), [is0] "r"(m_src_stride), [os0] "r"(n_dst_stride)
469
+ : "cc", "t0", "t1", "t2", "t3", "s1", "s2", "s3", "s4", "s5", "s6", "s7", "s8", "a1");
470
+ }
471
+
472
+ static inline void rvv_qk_dot_tile_f16_x1(float * dst,
473
+ const _Float16 * q_row,
474
+ const _Float16 * k_pack,
475
+ int64_t dk,
476
+ int64_t kv_tile) {
477
+ const size_t vl = __riscv_vsetvl_e16m1(kv_tile);
478
+ vfloat32m2_t acc = __riscv_vfmv_v_f_f32m2(0.0f, vl);
479
+
480
+ for (int64_t d = 0; d < dk; ++d) {
481
+ const vfloat16m1_t k_vec = __riscv_vle16_v_f16m1(k_pack + d * ggml_fa_tile_config::KV, vl);
482
+ acc = __riscv_vfwmacc_vf_f32m2(acc, q_row[d], k_vec, vl);
483
+ }
484
+
485
+ __riscv_vse32_v_f32m2(dst, acc, vl);
486
+ }
487
+
488
+ static inline void rvv_qk_dot_tile_f16_x4(float * dst0,
489
+ float * dst1,
490
+ float * dst2,
491
+ float * dst3,
492
+ const _Float16 * q0,
493
+ const _Float16 * q1,
494
+ const _Float16 * q2,
495
+ const _Float16 * q3,
496
+ const _Float16 * k_pack,
497
+ int64_t dk,
498
+ int64_t kv_tile) {
499
+ const size_t vl = __riscv_vsetvl_e16m1(kv_tile);
500
+ vfloat32m2_t acc0 = __riscv_vfmv_v_f_f32m2(0.0f, vl);
501
+ vfloat32m2_t acc1 = __riscv_vfmv_v_f_f32m2(0.0f, vl);
502
+ vfloat32m2_t acc2 = __riscv_vfmv_v_f_f32m2(0.0f, vl);
503
+ vfloat32m2_t acc3 = __riscv_vfmv_v_f_f32m2(0.0f, vl);
504
+
505
+ for (int64_t d = 0; d < dk; ++d) {
506
+ const vfloat16m1_t k_vec = __riscv_vle16_v_f16m1(k_pack + d * ggml_fa_tile_config::KV, vl);
507
+ acc0 = __riscv_vfwmacc_vf_f32m2(acc0, q0[d], k_vec, vl);
508
+ acc1 = __riscv_vfwmacc_vf_f32m2(acc1, q1[d], k_vec, vl);
509
+ acc2 = __riscv_vfwmacc_vf_f32m2(acc2, q2[d], k_vec, vl);
510
+ acc3 = __riscv_vfwmacc_vf_f32m2(acc3, q3[d], k_vec, vl);
511
+ }
512
+
513
+ __riscv_vse32_v_f32m2(dst0, acc0, vl);
514
+ __riscv_vse32_v_f32m2(dst1, acc1, vl);
515
+ __riscv_vse32_v_f32m2(dst2, acc2, vl);
516
+ __riscv_vse32_v_f32m2(dst3, acc3, vl);
517
+ }
518
+
519
+ static inline void rvv_pv_accumulate_f16_x1(float * dst,
520
+ const float * prob,
521
+ const _Float16 * v_pack,
522
+ int64_t kv_tile,
523
+ int64_t dv) {
524
+ int64_t d_left = dv;
525
+ int64_t d_off = 0;
526
+
527
+ while (d_left > 0) {
528
+ const size_t vl = __riscv_vsetvl_e16m2(d_left);
529
+ vfloat32m4_t acc = __riscv_vle32_v_f32m4(dst + d_off, vl);
530
+
531
+ for (int64_t tk = 0; tk < kv_tile; ++tk) {
532
+ const vfloat16m2_t v16 = __riscv_vle16_v_f16m2(v_pack + tk * dv + d_off, vl);
533
+ const vfloat32m4_t v32 = __riscv_vfwcvt_f_f_v_f32m4(v16, vl);
534
+ acc = __riscv_vfmacc_vf_f32m4(acc, prob[tk], v32, vl);
535
+ }
536
+
537
+ __riscv_vse32_v_f32m4(dst + d_off, acc, vl);
538
+ d_left -= vl;
539
+ d_off += vl;
540
+ }
541
+ }
542
+
543
+ static inline void rvv_pv_accumulate_f16_x4(float * dst0,
544
+ float * dst1,
545
+ float * dst2,
546
+ float * dst3,
547
+ const float * prob0,
548
+ const float * prob1,
549
+ const float * prob2,
550
+ const float * prob3,
551
+ const _Float16 * v_pack,
552
+ int64_t kv_tile,
553
+ int64_t dv) {
554
+ int64_t d_left = dv;
555
+ int64_t d_off = 0;
556
+
557
+ while (d_left > 0) {
558
+ const size_t vl = __riscv_vsetvl_e16m2(d_left);
559
+ vfloat32m4_t acc0 = __riscv_vle32_v_f32m4(dst0 + d_off, vl);
560
+ vfloat32m4_t acc1 = __riscv_vle32_v_f32m4(dst1 + d_off, vl);
561
+ vfloat32m4_t acc2 = __riscv_vle32_v_f32m4(dst2 + d_off, vl);
562
+ vfloat32m4_t acc3 = __riscv_vle32_v_f32m4(dst3 + d_off, vl);
563
+
564
+ for (int64_t tk = 0; tk < kv_tile; ++tk) {
565
+ const vfloat16m2_t v16 = __riscv_vle16_v_f16m2(v_pack + tk * dv + d_off, vl);
566
+ const vfloat32m4_t v32 = __riscv_vfwcvt_f_f_v_f32m4(v16, vl);
567
+ acc0 = __riscv_vfmacc_vf_f32m4(acc0, prob0[tk], v32, vl);
568
+ acc1 = __riscv_vfmacc_vf_f32m4(acc1, prob1[tk], v32, vl);
569
+ acc2 = __riscv_vfmacc_vf_f32m4(acc2, prob2[tk], v32, vl);
570
+ acc3 = __riscv_vfmacc_vf_f32m4(acc3, prob3[tk], v32, vl);
571
+ }
572
+
573
+ __riscv_vse32_v_f32m4(dst0 + d_off, acc0, vl);
574
+ __riscv_vse32_v_f32m4(dst1 + d_off, acc1, vl);
575
+ __riscv_vse32_v_f32m4(dst2 + d_off, acc2, vl);
576
+ __riscv_vse32_v_f32m4(dst3 + d_off, acc3, vl);
577
+ d_left -= vl;
578
+ d_off += vl;
579
+ }
580
+ }
581
+
582
+ static inline void rvv_qk_dot_tile(float * dst,
583
+ const float * q_row,
584
+ const float * k_pack,
585
+ int64_t dk,
586
+ int64_t kv_tile,
587
+ float scale) {
588
+ const size_t vl = __riscv_vsetvl_e32m4(kv_tile);
589
+ vfloat32m4_t acc = __riscv_vfmv_v_f_f32m4(0.0f, vl);
590
+
591
+ for (int64_t d = 0; d < dk; ++d) {
592
+ const vfloat32m4_t k_vec = __riscv_vle32_v_f32m4(k_pack + d * kv_tile, vl);
593
+ acc = __riscv_vfmacc_vf_f32m4(acc, q_row[d] * scale, k_vec, vl);
594
+ }
595
+
596
+ __riscv_vse32_v_f32m4(dst, acc, vl);
597
+ }
598
+
599
+ static inline void rvv_pv_accumulate(float * dst,
600
+ const float * prob,
601
+ const float * v_pack,
602
+ int64_t kv_tile,
603
+ int64_t dv) {
604
+ int64_t d_left = dv;
605
+ int64_t d_off = 0;
606
+
607
+ while (d_left > 0) {
608
+ const size_t vl = __riscv_vsetvl_e32m4(d_left);
609
+ vfloat32m4_t acc = __riscv_vle32_v_f32m4(dst + d_off, vl);
610
+
611
+ for (int64_t tk = 0; tk < kv_tile; ++tk) {
612
+ const vfloat32m4_t v_vec = __riscv_vle32_v_f32m4(v_pack + tk * dv + d_off, vl);
613
+ acc = __riscv_vfmacc_vf_f32m4(acc, prob[tk], v_vec, vl);
614
+ }
615
+
616
+ __riscv_vse32_v_f32m4(dst + d_off, acc, vl);
617
+ d_left -= vl;
618
+ d_off += vl;
619
+ }
620
+ }
621
+
622
+ static void permute_transpose_impl(const ggml_tensor * src0,
623
+ ggml_tensor * dst,
624
+ int64_t batch,
625
+ int64_t m,
626
+ int64_t n,
627
+ int64_t batch_stride,
628
+ int64_t m_src_stride,
629
+ int64_t n_src_stride,
630
+ int64_t n_dst_stride,
631
+ int ith,
632
+ int nth) {
633
+ GGML_ASSERT(n_src_stride == sizeof(int32_t) || n_src_stride == sizeof(int16_t));
634
+
635
+ if (n_src_stride == sizeof(int32_t)) {
636
+ for (int64_t bi = ith; bi < batch; bi += nth) {
637
+ rvv_transposed_s32_mn_to_nm((int8_t *) ((char *) dst->data + bi * batch_stride), n_dst_stride,
638
+ (int8_t *) ((char *) src0->data + bi * batch_stride), m_src_stride, m, n);
639
+ }
640
+ } else if (n_src_stride == sizeof(int16_t)) {
641
+ for (int64_t bi = ith; bi < batch; bi += nth) {
642
+ rvv_transposed_s32_mn_to_nm((int8_t *) ((char *) dst->data + bi * batch_stride), n_dst_stride,
643
+ (int8_t *) ((char *) src0->data + bi * batch_stride), m_src_stride, m, n);
644
+ }
645
+ } else {
646
+ GGML_ABORT("not implemented");
647
+ }
648
+ }
649
+
650
+ template <size_t QLEN>
651
+ static void flash_attn_ext_f16_one_chunk_inner_vlen1024_vf16_mrow(float ** pq,
652
+ const char * k_data_row,
653
+ const char * v_data_row,
654
+ const ggml_fp16_t * mp,
655
+ float ** sinks,
656
+ float ** dst,
657
+ float scale,
658
+ float logit_softcap,
659
+ float slope,
660
+ int64_t nek1,
661
+ int64_t nbk1,
662
+ int64_t nbv1,
663
+ int64_t DV,
664
+ int64_t DK,
665
+ void * tcm_buffer,
666
+ size_t tcm_buffer_size) {
667
+ GGML_ASSERT(flash_attn_ext_supported_shape_vlen1024_vf16(DK, DV));
668
+ float S[QLEN] = { 0.0f }; // sum
669
+ float M[QLEN] = { -INFINITY }; // maximum KQ value
670
+
671
+ _Float16 * kq16_buffer = (_Float16 *) tcm_buffer;
672
+ _Float16 * qv_buffer = kq16_buffer + QLEN * DV;
673
+ const size_t qkv_temp_buffer_size = (QLEN * DV + QLEN * DK) * sizeof(_Float16);
674
+ char * kv_tile_buffer = (char *) (qv_buffer + QLEN * DK);
675
+
676
+ {
677
+ vfloat16m2_t VKQ16_v = __riscv_vfmv_v_f_f16m2(0.0f, DV);
678
+ for (int64_t i = 0; i < QLEN; ++i) {
679
+ __riscv_vse16_v_f16m2(kq16_buffer + i * DV, VKQ16_v, DV);
680
+ vfloat16m2_t Q_q_v = __riscv_vfncvt_f_f_w_f16m2(__riscv_vle32_v_f32m4(pq[i], DK), DK);
681
+ __riscv_vse16_v_f16m2(qv_buffer + i * DK, Q_q_v, DK);
682
+ }
683
+ }
684
+
685
+ const uintptr_t scratch_addr = reinterpret_cast<uintptr_t>(kv_tile_buffer);
686
+ const size_t scratch_size = tcm_buffer_size > qkv_temp_buffer_size ? tcm_buffer_size - qkv_temp_buffer_size : 0;
687
+ const uintptr_t kq_tile_addr = align_up(scratch_addr, alignof(float));
688
+ const size_t scratch_prefix = kq_tile_addr - scratch_addr;
689
+ const size_t packed_tile_size =
690
+ QLEN * sizeof(float) + DK * sizeof(_Float16) + DV * sizeof(_Float16) + sizeof(float);
691
+ const int64_t max_ic_tile_step = ((int64_t) __riscv_vsetvlmax_e16m1()) & ~((int64_t) 7);
692
+ const int64_t max_fit_by_tcm =
693
+ scratch_size > scratch_prefix ? (int64_t) ((scratch_size - scratch_prefix) / packed_tile_size) : 0;
694
+ const int64_t ic_tile_step = std::min(max_ic_tile_step, max_fit_by_tcm) & ~((int64_t) 7);
695
+
696
+ const uintptr_t k_tile_addr = kq_tile_addr + QLEN * ic_tile_step * sizeof(float);
697
+ const uintptr_t v_tile_addr = k_tile_addr + DK * ic_tile_step * sizeof(_Float16);
698
+ const uintptr_t mv_tile_addr = v_tile_addr + ic_tile_step * DV * sizeof(_Float16);
699
+
700
+ if (ic_tile_step >= 8) {
701
+ float * kq_tile_buffer = reinterpret_cast<float *>(kq_tile_addr);
702
+ _Float16 * k_tile_pack = reinterpret_cast<_Float16 *>(k_tile_addr);
703
+ _Float16 * v_tile_pack = reinterpret_cast<_Float16 *>(v_tile_addr);
704
+ float * mv_tile_pack = reinterpret_cast<float *>(mv_tile_addr);
705
+
706
+ const int64_t k_tile_byte_stride = ic_tile_step * (int64_t) sizeof(_Float16);
707
+
708
+ int64_t ic_step = 0;
709
+ for (int64_t ic = 0; ic < nek1; ++ic) {
710
+ const float mv = mp ? slope * ((_Float16 *) mp)[ic] : 0.0f;
711
+
712
+ if (mv != -INFINITY) {
713
+ const _Float16 * k_data = (const _Float16 *) (k_data_row + ic * nbk1);
714
+ const _Float16 * v_data = (const _Float16 *) (v_data_row + ic * nbv1);
715
+
716
+ const vfloat16m2_t k_data_v = __riscv_vle16_v_f16m2(k_data, DK);
717
+ const vfloat16m2_t v_data_v = __riscv_vle16_v_f16m2(v_data, DV);
718
+ __riscv_vsse16_v_f16m2(k_tile_pack + ic_step, k_tile_byte_stride, k_data_v, DK);
719
+ __riscv_vse16_v_f16m2(v_tile_pack + ic_step * DV, v_data_v, DV);
720
+ mv_tile_pack[ic_step] = mv;
721
+ ic_step++;
722
+ }
723
+
724
+ if (ic_step > 0 && (ic_step == ic_tile_step || ic == (nek1 - 1))) {
725
+ if constexpr (QLEN == 4) {
726
+ const size_t qk_vl = __riscv_vsetvl_e16m1(ic_step);
727
+ vfloat32m2_t qk_acc0 = __riscv_vfmv_v_f_f32m2(0.0f, qk_vl);
728
+ vfloat32m2_t qk_acc1 = __riscv_vfmv_v_f_f32m2(0.0f, qk_vl);
729
+ vfloat32m2_t qk_acc2 = __riscv_vfmv_v_f_f32m2(0.0f, qk_vl);
730
+ vfloat32m2_t qk_acc3 = __riscv_vfmv_v_f_f32m2(0.0f, qk_vl);
731
+
732
+ for (int64_t d = 0; d < DK; ++d) {
733
+ const vfloat16m1_t k_vec = __riscv_vle16_v_f16m1(k_tile_pack + d * ic_tile_step, qk_vl);
734
+ qk_acc0 = __riscv_vfwmacc_vf_f32m2(qk_acc0, qv_buffer[0 * DK + d], k_vec, qk_vl);
735
+ qk_acc1 = __riscv_vfwmacc_vf_f32m2(qk_acc1, qv_buffer[1 * DK + d], k_vec, qk_vl);
736
+ qk_acc2 = __riscv_vfwmacc_vf_f32m2(qk_acc2, qv_buffer[2 * DK + d], k_vec, qk_vl);
737
+ qk_acc3 = __riscv_vfwmacc_vf_f32m2(qk_acc3, qv_buffer[3 * DK + d], k_vec, qk_vl);
738
+ }
739
+
740
+ qk_acc0 = __riscv_vfmul_vf_f32m2(qk_acc0, scale, qk_vl);
741
+ qk_acc1 = __riscv_vfmul_vf_f32m2(qk_acc1, scale, qk_vl);
742
+ qk_acc2 = __riscv_vfmul_vf_f32m2(qk_acc2, scale, qk_vl);
743
+ qk_acc3 = __riscv_vfmul_vf_f32m2(qk_acc3, scale, qk_vl);
744
+
745
+ __riscv_vse32_v_f32m2(kq_tile_buffer + 0 * ic_tile_step, qk_acc0, qk_vl);
746
+ __riscv_vse32_v_f32m2(kq_tile_buffer + 1 * ic_tile_step, qk_acc1, qk_vl);
747
+ __riscv_vse32_v_f32m2(kq_tile_buffer + 2 * ic_tile_step, qk_acc2, qk_vl);
748
+ __riscv_vse32_v_f32m2(kq_tile_buffer + 3 * ic_tile_step, qk_acc3, qk_vl);
749
+ } else {
750
+ static_assert(QLEN == 2, "unsupported QLEN");
751
+
752
+ const size_t qk_vl = __riscv_vsetvl_e16m1(ic_step);
753
+ vfloat32m2_t qk_acc0 = __riscv_vfmv_v_f_f32m2(0.0f, qk_vl);
754
+ vfloat32m2_t qk_acc1 = __riscv_vfmv_v_f_f32m2(0.0f, qk_vl);
755
+
756
+ for (int64_t d = 0; d < DK; ++d) {
757
+ const vfloat16m1_t k_vec = __riscv_vle16_v_f16m1(k_tile_pack + d * ic_tile_step, qk_vl);
758
+ qk_acc0 = __riscv_vfwmacc_vf_f32m2(qk_acc0, qv_buffer[0 * DK + d], k_vec, qk_vl);
759
+ qk_acc1 = __riscv_vfwmacc_vf_f32m2(qk_acc1, qv_buffer[1 * DK + d], k_vec, qk_vl);
760
+ }
761
+
762
+ qk_acc0 = __riscv_vfmul_vf_f32m2(qk_acc0, scale, qk_vl);
763
+ qk_acc1 = __riscv_vfmul_vf_f32m2(qk_acc1, scale, qk_vl);
764
+
765
+ __riscv_vse32_v_f32m2(kq_tile_buffer + 0 * ic_tile_step, qk_acc0, qk_vl);
766
+ __riscv_vse32_v_f32m2(kq_tile_buffer + 1 * ic_tile_step, qk_acc1, qk_vl);
767
+ }
768
+
769
+ for (int i = 0; i < QLEN; ++i) {
770
+ float * row_ptr = kq_tile_buffer + i * ic_tile_step;
771
+ const float tile_max =
772
+ rvv_softcap_add_max_inplace_f32(row_ptr, mv_tile_pack, ic_step, logit_softcap);
773
+
774
+ const float Mold = M[i];
775
+
776
+ if (tile_max > Mold) {
777
+ const float ms = expf(Mold - tile_max);
778
+ M[i] = tile_max;
779
+ S[i] *= ms;
780
+
781
+ vfloat16m2_t VKQ16_v = __riscv_vle16_v_f16m2(kq16_buffer + i * DV, DV);
782
+ VKQ16_v = __riscv_vfmul_vf_f16m2(VKQ16_v, (_Float16) ms, DV);
783
+ __riscv_vse16_v_f16m2(kq16_buffer + i * DV, VKQ16_v, DV);
784
+ }
785
+
786
+ S[i] += rvv_softmax_exp_inplace_f32(row_ptr, ic_step, M[i]);
787
+ }
788
+
789
+ if constexpr (QLEN == 4) {
790
+ vfloat16m2_t pv_acc0 = __riscv_vle16_v_f16m2(kq16_buffer + 0 * DV, DV);
791
+ vfloat16m2_t pv_acc1 = __riscv_vle16_v_f16m2(kq16_buffer + 1 * DV, DV);
792
+ vfloat16m2_t pv_acc2 = __riscv_vle16_v_f16m2(kq16_buffer + 2 * DV, DV);
793
+ vfloat16m2_t pv_acc3 = __riscv_vle16_v_f16m2(kq16_buffer + 3 * DV, DV);
794
+
795
+ for (int64_t tk = 0; tk < ic_step; ++tk) {
796
+ const vfloat16m2_t v16 = __riscv_vle16_v_f16m2(v_tile_pack + tk * DV, DV);
797
+ pv_acc0 =
798
+ __riscv_vfmacc_vf_f16m2(pv_acc0, (_Float16) kq_tile_buffer[0 * ic_tile_step + tk], v16, DV);
799
+ pv_acc1 =
800
+ __riscv_vfmacc_vf_f16m2(pv_acc1, (_Float16) kq_tile_buffer[1 * ic_tile_step + tk], v16, DV);
801
+ pv_acc2 =
802
+ __riscv_vfmacc_vf_f16m2(pv_acc2, (_Float16) kq_tile_buffer[2 * ic_tile_step + tk], v16, DV);
803
+ pv_acc3 =
804
+ __riscv_vfmacc_vf_f16m2(pv_acc3, (_Float16) kq_tile_buffer[3 * ic_tile_step + tk], v16, DV);
805
+ }
806
+
807
+ __riscv_vse16_v_f16m2(kq16_buffer + 0 * DV, pv_acc0, DV);
808
+ __riscv_vse16_v_f16m2(kq16_buffer + 1 * DV, pv_acc1, DV);
809
+ __riscv_vse16_v_f16m2(kq16_buffer + 2 * DV, pv_acc2, DV);
810
+ __riscv_vse16_v_f16m2(kq16_buffer + 3 * DV, pv_acc3, DV);
811
+ } else {
812
+ static_assert(QLEN == 2, "unsupported QLEN");
813
+ vfloat16m2_t pv_acc0 = __riscv_vle16_v_f16m2(kq16_buffer + 0 * DV, DV);
814
+ vfloat16m2_t pv_acc1 = __riscv_vle16_v_f16m2(kq16_buffer + 1 * DV, DV);
815
+
816
+ for (int64_t tk = 0; tk < ic_step; ++tk) {
817
+ const vfloat16m2_t v16 = __riscv_vle16_v_f16m2(v_tile_pack + tk * DV, DV);
818
+ pv_acc0 =
819
+ __riscv_vfmacc_vf_f16m2(pv_acc0, (_Float16) kq_tile_buffer[0 * ic_tile_step + tk], v16, DV);
820
+ pv_acc1 =
821
+ __riscv_vfmacc_vf_f16m2(pv_acc1, (_Float16) kq_tile_buffer[1 * ic_tile_step + tk], v16, DV);
822
+ }
823
+
824
+ __riscv_vse16_v_f16m2(kq16_buffer + 0 * DV, pv_acc0, DV);
825
+ __riscv_vse16_v_f16m2(kq16_buffer + 1 * DV, pv_acc1, DV);
826
+ }
827
+
828
+ ic_step = 0;
829
+ }
830
+ }
831
+ } else {
832
+ for (int64_t ic = 0; ic < nek1; ++ic) {
833
+ const float mv = mp ? slope * ((_Float16 *) mp)[ic] : 0.0f;
834
+
835
+ const char * k_data = k_data_row + ic * nbk1;
836
+ const char * v_data = v_data_row + ic * nbv1;
837
+
838
+ vfloat16m2_t k_data_v;
839
+ vfloat16m2_t v_data_v;
840
+
841
+ if (mv != -INFINITY) {
842
+ k_data_v = __riscv_vle16_v_f16m2((_Float16 *) k_data, DK);
843
+ v_data_v = __riscv_vle16_v_f16m2((_Float16 *) v_data, DV);
844
+ } else {
845
+ continue;
846
+ }
847
+
848
+ for (int i = 0; i < QLEN; ++i) {
849
+ vfloat16m2_t Q_q_v = __riscv_vle16_v_f16m2(qv_buffer + i * DK, DK);
850
+ vfloat32m4_t qk_acc_v = __riscv_vfwmul_vv_f32m4(k_data_v, Q_q_v, DK);
851
+ float s = reduce_sum_f32m4_vlen1024(qk_acc_v, DK);
852
+ s = s * scale;
853
+ if (logit_softcap != 0.0f) {
854
+ s = logit_softcap * tanhf(s);
855
+ }
856
+ s += mv;
857
+
858
+ const float Mold = M[i];
859
+
860
+ float ms = 1.0f; // upon new higher max val, scale VKQ and KQ sum with this value
861
+ float vs = 1.0f; // post-softmax KQ value, expf(s - M)
862
+
863
+ vfloat16m2_t VKQ16_v = __riscv_vle16_v_f16m2(kq16_buffer + i * DV, DV);
864
+ if (s > M[i]) {
865
+ // s is new maximum, ms < 1.0f, vs == expf(s - s) == 1.0f
866
+ M[i] = s;
867
+ ms = expf(Mold - M[i]);
868
+
869
+ // V = V*expf(Mold - M)
870
+ VKQ16_v = __riscv_vfmul_vf_f16m2(VKQ16_v, ms, DV);
871
+ } else {
872
+ // no new maximum, ms == 1.0f, vs != 1.0f
873
+ vs = expf(s - M[i]);
874
+ }
875
+ VKQ16_v = __riscv_vfmacc_vf_f16m2(VKQ16_v, vs, v_data_v, DV);
876
+ __riscv_vse16_v_f16m2(kq16_buffer + i * DV, VKQ16_v, DV);
877
+ S[i] = S[i] * ms + vs; // scale and increment sum with partial sum
878
+ }
879
+ }
880
+ }
881
+
882
+ for (int i = 0; i < QLEN; ++i) {
883
+ vfloat16m2_t VKQ16_v = __riscv_vle16_v_f16m2(kq16_buffer + i * DV, DV);
884
+ vfloat32m4_t VKQ32_v = __riscv_vfwcvt_f_f_v_f32m4(VKQ16_v, DV);
885
+
886
+ // sinks
887
+ if (sinks[i]) {
888
+ const float s = *(sinks[i]);
889
+
890
+ float ms = 1.0f;
891
+ float vs = 1.0f;
892
+
893
+ if (s > M[i]) {
894
+ ms = expf(M[i] - s);
895
+ M[i] = s;
896
+ VKQ32_v = __riscv_vfmul_vf_f32m4(VKQ32_v, ms, DV);
897
+ } else {
898
+ vs = expf(s - M[i]);
899
+ }
900
+
901
+ S[i] = S[i] * ms + vs;
902
+ }
903
+
904
+ // V /= S
905
+ const float S_inv = S[i] == 0.0f ? 0.0f : 1.0f / S[i];
906
+
907
+ VKQ32_v = __riscv_vfmul_vf_f32m4(VKQ32_v, S_inv, DV);
908
+
909
+ __riscv_vse32_v_f32m4(dst[i], VKQ32_v, DV);
910
+ }
911
+ }
912
+
913
+ static void flash_attn_ext_f16_one_chunk_inner_vlen1024_vf16_m1(const float * pq,
914
+ const char * k_data_row,
915
+ const char * v_data_row,
916
+ const ggml_fp16_t * mp,
917
+ const float * sinks,
918
+ float * dst,
919
+ float scale,
920
+ float logit_softcap,
921
+ float slope,
922
+ int64_t nek1,
923
+ int64_t nbk1,
924
+ int64_t nbv1,
925
+ int64_t DV,
926
+ int64_t DK) {
927
+ GGML_ASSERT(flash_attn_ext_supported_shape_vlen1024_vf16(DK, DV));
928
+
929
+ float S = 0.0f; // sum
930
+ float M = -INFINITY; // maximum KQ value
931
+
932
+ vfloat16m2_t VKQ16_v = __riscv_vfmv_v_f_f16m2(0.0f, DV);
933
+
934
+ vfloat16m2_t Q_q_v = __riscv_vfncvt_f_f_w_f16m2(__riscv_vle32_v_f32m4(pq, DK), DK);
935
+
936
+ for (int64_t ic = 0; ic < nek1; ++ic) {
937
+ const float mv = mp ? slope * ((_Float16 *) mp)[ic] : 0.0f;
938
+ if (mv == -INFINITY) {
939
+ continue;
940
+ }
941
+
942
+ const char * k_data = k_data_row + ic * nbk1;
943
+
944
+ vfloat16m2_t k_data_v = __riscv_vle16_v_f16m2((_Float16 *) k_data, DK);
945
+
946
+ vfloat32m4_t qk_acc_v = __riscv_vfwmul_vv_f32m4(k_data_v, Q_q_v, DK);
947
+ float s = reduce_sum_f32m4_vlen1024(qk_acc_v, DK);
948
+
949
+ s = s * scale; // scale KQ value
950
+
951
+ if (logit_softcap != 0.0f) {
952
+ s = logit_softcap * tanhf(s);
953
+ }
954
+
955
+ s += mv; // apply mask
956
+
957
+ const float Mold = M;
958
+
959
+ float ms = 1.0f; // upon new higher max val, scale VKQ and KQ sum with this value
960
+ float vs = 1.0f; // post-softmax KQ value, expf(s - M)
961
+
962
+ const char * v_data = v_data_row + ic * nbv1;
963
+
964
+ vfloat16m2_t v_data_v = __riscv_vle16_v_f16m2((_Float16 *) v_data, DV);
965
+
966
+ if (s > M) {
967
+ // s is new maximum, ms < 1.0f, vs == expf(s - s) == 1.0f
968
+ M = s;
969
+ ms = expf(Mold - M);
970
+
971
+ // V = V*expf(Mold - M)
972
+ VKQ16_v = __riscv_vfmul_vf_f16m2(VKQ16_v, ms, DV);
973
+ } else {
974
+ // no new maximum, ms == 1.0f, vs != 1.0f
975
+ vs = expf(s - M);
976
+ }
977
+
978
+ VKQ16_v = __riscv_vfmacc_vf_f16m2(VKQ16_v, vs, v_data_v, DV);
979
+
980
+ S = S * ms + vs; // scale and increment sum with partial sum
981
+ }
982
+
983
+ vfloat32m4_t VKQ32_v = __riscv_vfwcvt_f_f_v_f32m4(VKQ16_v, DV);
984
+
985
+ // sinks
986
+ if (sinks) {
987
+ const float s = *sinks;
988
+
989
+ float ms = 1.0f;
990
+ float vs = 1.0f;
991
+
992
+ if (s > M) {
993
+ ms = expf(M - s);
994
+ M = s;
995
+ VKQ32_v = __riscv_vfmul_vf_f32m4(VKQ32_v, ms, DV);
996
+ } else {
997
+ vs = expf(s - M);
998
+ }
999
+
1000
+ S = S * ms + vs;
1001
+ }
1002
+
1003
+ // V /= S
1004
+ const float S_inv = S == 0.0f ? 0.0f : 1.0f / S;
1005
+
1006
+ VKQ32_v = __riscv_vfmul_vf_f32m4(VKQ32_v, S_inv, DV);
1007
+
1008
+ __riscv_vse32_v_f32m4(dst, VKQ32_v, DV);
1009
+ }
1010
+
1011
+ } // namespace
1012
+
1013
+ void memcpy1d(void * dst, const void * src, int64_t size) {
1014
+ size_t byte_size_all = size;
1015
+ size_t vlen = __riscv_vlenb() * 8;
1016
+ if (vlen == 256) {
1017
+ // 1024 bytes
1018
+ __asm__ volatile(
1019
+ //
1020
+ "srli t0, %[size], 10 \n\t"
1021
+ "blez t0, memcpy_tail%= \n\t"
1022
+ "vsetvli t1, x0, e8, m8, tu, mu \n\t"
1023
+ "memcpy_main_loop%=: \n\t"
1024
+ "addi t0, t0, -1 \n\t"
1025
+ "vle8.v v0, (%[s]) \n\t"
1026
+ "addi %[s], %[s], 256 \n\t"
1027
+ "vle8.v v8, (%[s]) \n\t"
1028
+ "addi %[s], %[s], 256 \n\t"
1029
+ "vle8.v v16, (%[s]) \n\t"
1030
+ "addi %[s], %[s], 256 \n\t"
1031
+ "vle8.v v24, (%[s]) \n\t"
1032
+ "addi %[s], %[s], 256 \n\t"
1033
+ //
1034
+ "vse8.v v0, (%[d]) \n\t"
1035
+ "addi %[d], %[d], 256 \n\t"
1036
+ "vse8.v v8, (%[d]) \n\t"
1037
+ "addi %[d], %[d], 256 \n\t"
1038
+ "vse8.v v16, (%[d]) \n\t"
1039
+ "addi %[d], %[d], 256 \n\t"
1040
+ "vse8.v v24, (%[d]) \n\t"
1041
+ "addi %[d], %[d], 256 \n\t"
1042
+ //
1043
+ "bnez t0, memcpy_main_loop%= \n\t"
1044
+ "memcpy_tail%=: \n\t"
1045
+ "andi t1, %[size], 1023 \n\t"
1046
+ "blez t1, out%= \n\t"
1047
+ "memcpy_tail_loop%=: \n\t"
1048
+ "vsetvli t0, t1, e8, m8, tu, mu \n\t"
1049
+ "sub t1, t1, t0 \n\t"
1050
+ "vle8.v v0, (%[s]) \n\t"
1051
+ "add %[s], %[s], t0 \n\t"
1052
+ "vse8.v v0, (%[d]) \n\t"
1053
+ "add %[d], %[d], t0 \n\t"
1054
+ "bnez t1, memcpy_tail_loop%= \n\t"
1055
+ "out%=: \n\t"
1056
+ : [s] "+r"(src), [d] "+r"(dst)
1057
+ : [size] "r"(byte_size_all)
1058
+ : "cc", "t0", "t1");
1059
+ } else if (vlen == 1024) {
1060
+ // 2048 bytes
1061
+ __asm__ volatile(
1062
+ //
1063
+ "srli t0, %[size], 11 \n\t"
1064
+ "blez t0, memcpy_tail%= \n\t"
1065
+ "vsetvli t1, x0, e8, m8, tu, mu \n\t"
1066
+ "addi t2, %[s], 1024 \n\t"
1067
+ "addi t3, %[d], 1024 \n\t"
1068
+ "li t5, 2048 \n\t"
1069
+ "memcpy_main_loop%=: \n\t"
1070
+ "addi t0, t0, -1 \n\t"
1071
+ "vle8.v v0, (%[s]) \n\t"
1072
+ "add %[s], %[s], t5 \n\t"
1073
+ "vle8.v v8, (t2) \n\t"
1074
+ "add t2, t2, t5 \n\t"
1075
+ //
1076
+ "vse8.v v0, (%[d]) \n\t"
1077
+ "add %[d], %[d], t5 \n\t"
1078
+ "vse8.v v8, (t3) \n\t"
1079
+ "add t3, t3, t5 \n\t"
1080
+ //
1081
+ "bnez t0, memcpy_main_loop%= \n\t"
1082
+ "memcpy_tail%=: \n\t"
1083
+ "andi t1, %[size], 2047 \n\t"
1084
+ "blez t1, out%= \n\t"
1085
+ "memcpy_tail_loop%=: \n\t"
1086
+ "vsetvli t0, t1, e8, m2, tu, mu \n\t"
1087
+ "sub t1, t1, t0 \n\t"
1088
+ "vle8.v v0, (%[s]) \n\t"
1089
+ "add %[s], %[s], t0 \n\t"
1090
+ "vse8.v v0, (%[d]) \n\t"
1091
+ "add %[d], %[d], t0 \n\t"
1092
+ "bnez t1, memcpy_tail_loop%= \n\t"
1093
+ "out%=: \n\t"
1094
+ : [s] "+r"(src), [d] "+r"(dst)
1095
+ : [size] "r"(byte_size_all)
1096
+ : "cc", "t0", "t1", "t2", "t3", "t5");
1097
+ } else {
1098
+ __asm__ volatile(
1099
+ //
1100
+ "add t1, %[size], zero \n\t"
1101
+ "memcpy_tail_loop%=: \n\t"
1102
+ "vsetvli t0, t1, e8, m8, tu, mu \n\t"
1103
+ "sub t1, t1, t0 \n\t"
1104
+ "vle8.v v0, (%[s]) \n\t"
1105
+ "add %[s], %[s], t0 \n\t"
1106
+ "vse8.v v0, (%[d]) \n\t"
1107
+ "add %[d], %[d], t0 \n\t"
1108
+ "bnez t1, memcpy_tail_loop%= \n\t"
1109
+ : [s] "+r"(src), [d] "+r"(dst)
1110
+ : [size] "r"(byte_size_all)
1111
+ : "cc", "t0", "t1", "t2", "t4", "t3");
1112
+ }
1113
+ }
1114
+
1115
+ void memcpy2d(void * dst, int64_t dst_stride, const void * src, int64_t src_stride, int64_t tile_rows, int64_t size) {
1116
+ for (int64_t i = 0; i < tile_rows; ++i) {
1117
+ memcpy1d((char *) dst + i * dst_stride, (const char *) src + i * src_stride, size);
1118
+ }
1119
+ }
1120
+
1121
+ void forward_flash_attn_ext_f16_one_chunk_vlen1024_vf16(const ggml_compute_params * params,
1122
+ ggml_tensor * dst,
1123
+ int ir0,
1124
+ int ir1,
1125
+ void * tcm_buffer,
1126
+ size_t tcm_buffer_size) {
1127
+ const ggml_tensor * q = dst->src[0];
1128
+ const ggml_tensor * k = dst->src[1];
1129
+ const ggml_tensor * v = dst->src[2];
1130
+ const ggml_tensor * mask = dst->src[3];
1131
+ const ggml_tensor * sinks = dst->src[4];
1132
+
1133
+ GGML_TENSOR_LOCALS(int64_t, neq, q, ne)
1134
+ GGML_TENSOR_LOCALS(size_t, nbq, q, nb)
1135
+ GGML_TENSOR_LOCALS(int64_t, nek, k, ne)
1136
+ GGML_TENSOR_LOCALS(size_t, nbk, k, nb)
1137
+ GGML_TENSOR_LOCALS(int64_t, nev, v, ne)
1138
+ GGML_TENSOR_LOCALS(size_t, nbv, v, nb)
1139
+ GGML_TENSOR_LOCALS(int64_t, ne, dst, ne)
1140
+ GGML_TENSOR_LOCALS(size_t, nb, dst, nb)
1141
+
1142
+ const int64_t DK = nek0;
1143
+ const int64_t DV = nev0;
1144
+ const int64_t N = neq1;
1145
+
1146
+ GGML_ASSERT(flash_attn_ext_supported_shape_vlen1024_vf16(DK, DV));
1147
+
1148
+ // broadcast factors
1149
+ const int64_t rk2 = neq2 / nek2;
1150
+ const int64_t rk3 = neq3 / nek3;
1151
+
1152
+ const int64_t rv2 = neq2 / nev2;
1153
+ const int64_t rv3 = neq3 / nev3;
1154
+
1155
+ // parallelize by q rows using ggml_vec_dot_f32
1156
+
1157
+ float scale = *((float *) dst->op_params + 0);
1158
+ float max_bias = *((float *) dst->op_params + 1);
1159
+ float logit_softcap = *((float *) dst->op_params + 2);
1160
+
1161
+ if (logit_softcap != 0) {
1162
+ scale /= logit_softcap;
1163
+ }
1164
+
1165
+ const uint32_t n_head = neq2;
1166
+ const uint32_t n_head_log2 = 1u << (uint32_t) floor(log2(n_head));
1167
+
1168
+ const float m0 = powf(2.0f, -(max_bias) / n_head_log2);
1169
+ const float m1 = powf(2.0f, -(max_bias / 2.0f) / n_head_log2);
1170
+
1171
+ const int KV_row_size = DK * sizeof(_Float16) + DV * sizeof(_Float16);
1172
+
1173
+ int ith = params->ith;
1174
+ int ir_step = 1;
1175
+ for (int ir = ir0; ir < ir1; ir += ir_step) {
1176
+ // q indices
1177
+ const int iq3 = ir / (neq2 * neq1);
1178
+ const int iq2 = (ir - iq3 * neq2 * neq1) / neq1;
1179
+ const int iq1 = (ir - iq3 * neq2 * neq1 - iq2 * neq1);
1180
+
1181
+ const int iq3_1 = (ir + 1) / (neq2 * neq1);
1182
+ const int iq2_1 = (ir + 1 - iq3_1 * neq2 * neq1) / neq1;
1183
+ const int iq1_1 = (ir + 1 - iq3_1 * neq2 * neq1 - iq2_1 * neq1);
1184
+
1185
+ const int iq3_2 = (ir + 2) / (neq2 * neq1);
1186
+ const int iq2_2 = (ir + 2 - iq3_2 * neq2 * neq1) / neq1;
1187
+ const int iq1_2 = (ir + 2 - iq3_2 * neq2 * neq1 - iq2_2 * neq1);
1188
+
1189
+ const int iq3_3 = (ir + 3) / (neq2 * neq1);
1190
+ const int iq2_3 = (ir + 3 - iq3_3 * neq2 * neq1) / neq1;
1191
+ const int iq1_3 = (ir + 3 - iq3_3 * neq2 * neq1 - iq2_3 * neq1);
1192
+
1193
+ const uint32_t h = iq2; // head index
1194
+ const float slope =
1195
+ (max_bias > 0.0f) ? h < n_head_log2 ? powf(m0, h + 1) : powf(m1, 2 * (h - n_head_log2) + 1) : 1.0f;
1196
+
1197
+ const ggml_fp16_t * mp =
1198
+ mask ? (ggml_fp16_t *) ((char *) mask->data + iq1 * mask->nb[1] + (iq2 % mask->ne[2]) * mask->nb[2] +
1199
+ (iq3 % mask->ne[3]) * mask->nb[3]) :
1200
+ NULL;
1201
+
1202
+ const bool mp_equal_2 = iq1_1 == iq1 && (iq2 % mask->ne[2]) == (iq2_1 % mask->ne[2]) &&
1203
+ (iq3 % mask->ne[3]) == (iq3_1 % mask->ne[3]);
1204
+
1205
+ const bool mp_equal_4 = mp_equal_2 && iq1_2 == iq1 && (iq2 % mask->ne[2]) == (iq2_2 % mask->ne[2]) &&
1206
+ (iq3 % mask->ne[3]) == (iq3_2 % mask->ne[3]) && iq1_3 == iq1 &&
1207
+ (iq2 % mask->ne[2]) == (iq2_3 % mask->ne[2]) &&
1208
+ (iq3 % mask->ne[3]) == (iq3_3 % mask->ne[3]);
1209
+
1210
+ // k indices
1211
+ const int ik3 = iq3 / rk3;
1212
+ const int ik2 = iq2 / rk2;
1213
+
1214
+ const int ik3_1 = iq3_1 / rk3;
1215
+ const int ik2_1 = iq2_1 / rk2;
1216
+
1217
+ const int ik3_2 = iq3_2 / rk3;
1218
+ const int ik2_2 = iq2_2 / rk2;
1219
+
1220
+ const int ik3_3 = iq3_3 / rk3;
1221
+ const int ik2_3 = iq2_3 / rk2;
1222
+
1223
+ // v indices
1224
+ const int iv3 = iq3 / rv3;
1225
+ const int iv2 = iq2 / rv2;
1226
+
1227
+ const int iv3_1 = iq3_1 / rv3;
1228
+ const int iv2_1 = iq2_1 / rv2;
1229
+
1230
+ const int iv3_2 = iq3_2 / rv3;
1231
+ const int iv2_2 = iq2_2 / rv2;
1232
+
1233
+ const int iv3_3 = iq3_3 / rv3;
1234
+ const int iv2_3 = iq2_3 / rv2;
1235
+
1236
+ const float * pq = (const float *) ((char *) q->data + (iq1 * nbq1 + iq2 * nbq2 + iq3 * nbq3));
1237
+
1238
+ std::array<float *, 4> pq_buffer;
1239
+ std::array<float *, 4> sinks_buffer;
1240
+ std::array<float *, 4> dst_buffer;
1241
+
1242
+ if (tcm_buffer != nullptr && 4 * KV_row_size < tcm_buffer_size && ir < (ir1 - 3) && mp_equal_4 &&
1243
+ ik3_3 == ik3 && ik2_3 == ik2 && iv3_3 == iv3 && iv2_3 == iv2 && ik3_2 == ik3 && ik2_2 == ik2 &&
1244
+ iv3_2 == iv3 && iv2_2 == iv2 && ik3_1 == ik3 && ik2_1 == ik2 && iv3_1 == iv3 && iv2_1 == iv2) {
1245
+ ir_step = 4;
1246
+
1247
+ pq_buffer[0] = (float *) ((char *) q->data + (iq1 * nbq1 + iq2 * nbq2 + iq3 * nbq3));
1248
+ pq_buffer[1] = (float *) ((char *) q->data + (iq1_1 * nbq1 + iq2_1 * nbq2 + iq3_1 * nbq3));
1249
+ pq_buffer[2] = (float *) ((char *) q->data + (iq1_2 * nbq1 + iq2_2 * nbq2 + iq3_2 * nbq3));
1250
+ pq_buffer[3] = (float *) ((char *) q->data + (iq1_3 * nbq1 + iq2_3 * nbq2 + iq3_3 * nbq3));
1251
+
1252
+ sinks_buffer[0] = sinks ? ((float *) ((char *) sinks->data)) + iq2 : nullptr;
1253
+ sinks_buffer[1] = sinks ? ((float *) ((char *) sinks->data)) + iq2_1 : nullptr;
1254
+ sinks_buffer[2] = sinks ? ((float *) ((char *) sinks->data)) + iq2_2 : nullptr;
1255
+ sinks_buffer[3] = sinks ? ((float *) ((char *) sinks->data)) + iq2_3 : nullptr;
1256
+
1257
+ dst_buffer[0] = (float *) ((char *) dst->data + (iq3 * ne2 * ne1 + iq2 + iq1 * ne1) * nb1);
1258
+ dst_buffer[1] = (float *) ((char *) dst->data + (iq3_1 * ne2 * ne1 + iq2_1 + iq1_1 * ne1) * nb1);
1259
+ dst_buffer[2] = (float *) ((char *) dst->data + (iq3_2 * ne2 * ne1 + iq2_2 + iq1_2 * ne1) * nb1);
1260
+ dst_buffer[3] = (float *) ((char *) dst->data + (iq3_3 * ne2 * ne1 + iq2_3 + iq1_3 * ne1) * nb1);
1261
+
1262
+ flash_attn_ext_f16_one_chunk_inner_vlen1024_vf16_mrow<4>( //
1263
+ pq_buffer.data(), //
1264
+ (const char *) k->data + (ik2 * nbk2 + ik3 * nbk3), //
1265
+ (const char *) v->data + (iv2 * nbv2 + iv3 * nbv3), //
1266
+ mp, //
1267
+ sinks_buffer.data(), //
1268
+ dst_buffer.data(), //
1269
+ scale, logit_softcap, slope, nek1, nbk1, nbv1, DV, DK, tcm_buffer, tcm_buffer_size);
1270
+ } else if (tcm_buffer != nullptr && 2 * KV_row_size < tcm_buffer_size && ir < (ir1 - 1) && mp_equal_2 &&
1271
+ ik3_1 == ik3 && ik2_1 == ik2 && iv3_1 == iv3 && iv2_1 == iv2) {
1272
+ ir_step = 2;
1273
+
1274
+ pq_buffer[0] = (float *) ((char *) q->data + (iq1 * nbq1 + iq2 * nbq2 + iq3 * nbq3));
1275
+ pq_buffer[1] = (float *) ((char *) q->data + (iq1_1 * nbq1 + iq2_1 * nbq2 + iq3_1 * nbq3));
1276
+
1277
+ sinks_buffer[0] = sinks ? ((float *) ((char *) sinks->data)) + iq2 : nullptr;
1278
+ sinks_buffer[1] = sinks ? ((float *) ((char *) sinks->data)) + iq2_1 : nullptr;
1279
+
1280
+ dst_buffer[0] = (float *) ((char *) dst->data + (iq3 * ne2 * ne1 + iq2 + iq1 * ne1) * nb1);
1281
+ dst_buffer[1] = (float *) ((char *) dst->data + (iq3_1 * ne2 * ne1 + iq2_1 + iq1_1 * ne1) * nb1);
1282
+
1283
+ flash_attn_ext_f16_one_chunk_inner_vlen1024_vf16_mrow<2>( //
1284
+ pq_buffer.data(), //
1285
+ (const char *) k->data + (ik2 * nbk2 + ik3 * nbk3), //
1286
+ (const char *) v->data + (iv2 * nbv2 + iv3 * nbv3), //
1287
+ mp, //
1288
+ sinks_buffer.data(), //
1289
+ dst_buffer.data(), //
1290
+ scale, logit_softcap, slope, nek1, nbk1, nbv1, DV, DK, tcm_buffer, tcm_buffer_size);
1291
+ } else {
1292
+ ir_step = 1;
1293
+ flash_attn_ext_f16_one_chunk_inner_vlen1024_vf16_m1( //
1294
+ pq, //
1295
+ (const char *) k->data + (ik2 * nbk2 + ik3 * nbk3), //
1296
+ (const char *) v->data + (iv2 * nbv2 + iv3 * nbv3), //
1297
+ mp, //
1298
+ sinks ? ((float *) ((char *) sinks->data)) + h : nullptr, //
1299
+ (float *) ((char *) dst->data + (iq3 * ne2 * ne1 + iq2 + iq1 * ne1) * nb1), //
1300
+ scale, logit_softcap, slope, nek1, nbk1, nbv1, DV, DK);
1301
+ }
1302
+ }
1303
+ }
1304
+
1305
+ void forward_flash_attn_ext_f16_tiled_vlen1024_vf16(const ggml_compute_params * params,
1306
+ ggml_tensor * dst,
1307
+ int ir0,
1308
+ int ir1,
1309
+ void * tcm_buffer,
1310
+ size_t tcm_buffer_size) {
1311
+ const ggml_tensor * q = dst->src[0];
1312
+ const ggml_tensor * k = dst->src[1];
1313
+ const ggml_tensor * v = dst->src[2];
1314
+ const ggml_tensor * mask = dst->src[3];
1315
+ const ggml_tensor * sinks = dst->src[4];
1316
+
1317
+ GGML_TENSOR_LOCALS(int64_t, neq, q, ne)
1318
+ GGML_TENSOR_LOCALS(size_t, nbq, q, nb)
1319
+ GGML_TENSOR_LOCALS(int64_t, nek, k, ne)
1320
+ GGML_TENSOR_LOCALS(size_t, nbk, k, nb)
1321
+ GGML_TENSOR_LOCALS(int64_t, nev, v, ne)
1322
+ GGML_TENSOR_LOCALS(size_t, nbv, v, nb)
1323
+ GGML_TENSOR_LOCALS(int64_t, ne, dst, ne)
1324
+ GGML_TENSOR_LOCALS(size_t, nb, dst, nb)
1325
+
1326
+ const int64_t DK = nek0;
1327
+ const int64_t DV = nev0;
1328
+ const int64_t N = neq1;
1329
+
1330
+ GGML_ASSERT(flash_attn_ext_supported_shape_vlen1024_vf16(DK, DV));
1331
+
1332
+ GGML_ASSERT(ne0 == DV);
1333
+ GGML_ASSERT(ne2 == N);
1334
+
1335
+ // input tensor rows must be contiguous
1336
+ GGML_ASSERT(nbq0 == ggml_type_size(q->type));
1337
+ GGML_ASSERT(nbk0 == ggml_type_size(k->type));
1338
+ GGML_ASSERT(nbv0 == ggml_type_size(v->type));
1339
+
1340
+ GGML_ASSERT(neq0 == DK);
1341
+ GGML_ASSERT(nek0 == DK);
1342
+ GGML_ASSERT(nev0 == DV);
1343
+
1344
+ GGML_ASSERT(neq1 == N);
1345
+
1346
+ // dst cannot be transposed or permuted
1347
+ GGML_ASSERT(nb0 == sizeof(float));
1348
+ GGML_ASSERT(nb0 <= nb1);
1349
+ GGML_ASSERT(nb1 <= nb2);
1350
+ GGML_ASSERT(nb2 <= nb3);
1351
+
1352
+ GGML_ASSERT(k->type == v->type);
1353
+ const ggml_type kv_type = k->type;
1354
+
1355
+ // broadcast factors
1356
+ const int64_t rk2 = neq2 / nek2;
1357
+ const int64_t rk3 = neq3 / nek3;
1358
+
1359
+ const int64_t rv2 = neq2 / nev2;
1360
+ const int64_t rv3 = neq3 / nev3;
1361
+
1362
+ float * param_list = (float *) dst->op_params;
1363
+ float scale = param_list[0];
1364
+ float max_bias = param_list[1];
1365
+ float logit_softcap = param_list[2];
1366
+
1367
+ if (logit_softcap != 0) {
1368
+ scale /= logit_softcap;
1369
+ }
1370
+
1371
+ const uint32_t n_head = neq2;
1372
+ const uint32_t n_head_log2 = 1u << (uint32_t) floor(log2(n_head));
1373
+
1374
+ const float m0 = powf(2.0f, -(max_bias) / n_head_log2);
1375
+ const float m1 = powf(2.0f, -(max_bias / 2.0f) / n_head_log2);
1376
+
1377
+ int ith = params->ith;
1378
+
1379
+ static constexpr int Q_TILE_SZ = ggml_fa_tile_config::Q;
1380
+ static constexpr int KV_TILE_SZ = ggml_fa_tile_config::KV;
1381
+
1382
+ // Per-thread scratch layout:
1383
+ // Q_f32: Q_TILE_SZ * DK
1384
+ // KQ: Q_TILE_SZ * KV_TILE_SZ
1385
+ // mask32: Q_TILE_SZ * KV_TILE_SZ
1386
+ // VKQ32: Q_TILE_SZ * DV
1387
+ // V32: KV_TILE_SZ * DV
1388
+ // K_f32: DK * KV_TILE_SZ (transposed K tile)
1389
+ float * base = (float *) params->wdata + ith * (Q_TILE_SZ * DK + 2 * Q_TILE_SZ * KV_TILE_SZ + Q_TILE_SZ * DV +
1390
+ KV_TILE_SZ * DV + KV_TILE_SZ * DK + CACHE_LINE_SIZE_F32);
1391
+ const size_t base_size =
1392
+ (Q_TILE_SZ * DK + 2 * Q_TILE_SZ * KV_TILE_SZ + Q_TILE_SZ * DV + KV_TILE_SZ * DV + KV_TILE_SZ * DK) *
1393
+ sizeof(float) +
1394
+ CACHE_LINE_SIZE_F32;
1395
+
1396
+ if (base_size <= tcm_buffer_size && tcm_buffer != nullptr) {
1397
+ base = (float *) tcm_buffer;
1398
+ }
1399
+
1400
+ float S_M_Buf[Q_TILE_SZ * 2]; // buffer to hold S, M, bias for one tile to reduce register pressure in main loop
1401
+ float * S = S_M_Buf;
1402
+ float * M = S_M_Buf + Q_TILE_SZ;
1403
+
1404
+ int ir = ir0;
1405
+ while (ir < ir1) {
1406
+ // q indices for the start of this tile
1407
+ const int iq3 = ir / (neq2 * neq1);
1408
+ const int iq2 = (ir - iq3 * neq2 * neq1) / neq1;
1409
+ const int iq1 = (ir - iq3 * neq2 * neq1 - iq2 * neq1);
1410
+
1411
+ // Number of valid rows in this tile:
1412
+ // - limited by tile size (Q_TILE_SZ)
1413
+ // - limited by chunk boundary (ir1 - ir)
1414
+ // - limited by head boundary (neq1 - iq1) to avoid crossing into next head
1415
+ const int tile_rows = MIN(Q_TILE_SZ, MIN((int) (ir1 - ir), (int) (neq1 - iq1)));
1416
+ GGML_ASSERT(tile_rows > 0);
1417
+
1418
+ const uint32_t h = iq2; // head index
1419
+ const float slope =
1420
+ (max_bias > 0.0f) ? h < n_head_log2 ? powf(m0, h + 1) : powf(m1, 2 * (h - n_head_log2) + 1) : 1.0f;
1421
+
1422
+ for (int i = 0; i < Q_TILE_SZ; ++i) {
1423
+ S[i] = 0.;
1424
+ M[i] = -INFINITY;
1425
+ }
1426
+
1427
+ float * Q_f32 = base;
1428
+ float * KQ = (float *) ((char *) base + Q_TILE_SZ * DK * sizeof(float));
1429
+ float * mask32 = KQ + Q_TILE_SZ * KV_TILE_SZ;
1430
+ float * VKQ32 = mask32 + Q_TILE_SZ * KV_TILE_SZ;
1431
+ float * V32 = VKQ32 + Q_TILE_SZ * DV;
1432
+ float * K_f32 = V32 + KV_TILE_SZ * DV;
1433
+ _Float16 * Q_f16 = (_Float16 *) Q_f32;
1434
+ _Float16 * V_f16 = (_Float16 *) V32;
1435
+ _Float16 * K_f16 = (_Float16 *) K_f32;
1436
+
1437
+ rvv_zero_f32(VKQ32, Q_TILE_SZ * DV);
1438
+
1439
+ // k indices
1440
+ const int ik3 = iq3 / rk3;
1441
+ const int ik2 = iq2 / rk2;
1442
+
1443
+ // v indices
1444
+ const int iv3 = iq3 / rv3;
1445
+ const int iv2 = iq2 / rv2;
1446
+
1447
+ const float * pq = (const float *) ((char *) q->data + (iq1 * nbq1 + iq2 * nbq2 + iq3 * nbq3));
1448
+ if (kv_type == GGML_TYPE_F16) {
1449
+ rvv_pack_f32_as_scaled_f16((uint8_t *) Q_f16, DK * sizeof(_Float16), (uint8_t *) pq, nbq1, tile_rows, DK,
1450
+ scale);
1451
+ } else {
1452
+ memcpy2d(Q_f32, DK * sizeof(float), pq, nbq1, tile_rows, DK * sizeof(float));
1453
+ }
1454
+
1455
+ for (int64_t ic = 0; ic < nek1; ic += KV_TILE_SZ) {
1456
+ const int kv_tile = (int) std::min((int64_t) KV_TILE_SZ, nek1 - ic);
1457
+
1458
+ rvv_zero_f32(K_f32, DK * KV_TILE_SZ);
1459
+ rvv_zero_f32(V32, KV_TILE_SZ * DV);
1460
+
1461
+ // skip the tile entirely if all the masks are -inf
1462
+ if (mask) {
1463
+ bool can_skip = true;
1464
+ const ggml_fp16_t * mp_row =
1465
+ (const ggml_fp16_t *) ((const char *) mask->data + iq1 * mask->nb[1] +
1466
+ (iq2 % mask->ne[2]) * mask->nb[2] + (iq3 % mask->ne[3]) * mask->nb[3]);
1467
+ rvv_pack_scaled_f16_as_f32(mask32, KV_TILE_SZ * sizeof(float), mp_row + ic, mask->nb[1], tile_rows,
1468
+ kv_tile, slope);
1469
+
1470
+ for (int tq = 0; tq < tile_rows; tq++) {
1471
+ for (int tk = 0; tk < kv_tile; tk++) {
1472
+ if (mask32[tq * KV_TILE_SZ + tk] != -INFINITY) {
1473
+ can_skip = false;
1474
+ }
1475
+ }
1476
+ // Pad remaining mask entries with -inf
1477
+ for (int tk = kv_tile; tk < KV_TILE_SZ; tk++) {
1478
+ mask32[tq * KV_TILE_SZ + tk] = -INFINITY;
1479
+ }
1480
+ }
1481
+
1482
+ if (can_skip) {
1483
+ continue;
1484
+ }
1485
+ }
1486
+
1487
+ if (kv_type == GGML_TYPE_F16) {
1488
+ rvv_transposed_s16_mn_to_nm((int8_t *) K_f16, KV_TILE_SZ * sizeof(_Float16),
1489
+ (int8_t *) k->data + ic * nbk1 + ik2 * nbk2 + ik3 * nbk3, nbk1, kv_tile,
1490
+ DK);
1491
+
1492
+ int tq = 0;
1493
+ for (; tq + 3 < tile_rows; tq += 4) {
1494
+ rvv_qk_dot_tile_f16_x4(KQ + (tq + 0) * KV_TILE_SZ, KQ + (tq + 1) * KV_TILE_SZ,
1495
+ KQ + (tq + 2) * KV_TILE_SZ, KQ + (tq + 3) * KV_TILE_SZ,
1496
+ Q_f16 + (tq + 0) * DK, Q_f16 + (tq + 1) * DK, Q_f16 + (tq + 2) * DK,
1497
+ Q_f16 + (tq + 3) * DK, K_f16, DK, kv_tile);
1498
+ }
1499
+ for (; tq < tile_rows; ++tq) {
1500
+ rvv_qk_dot_tile_f16_x1(KQ + tq * KV_TILE_SZ, Q_f16 + tq * DK, K_f16, DK, kv_tile);
1501
+ }
1502
+ } else {
1503
+ for (int tk = 0; tk < kv_tile; tk++) {
1504
+ const char * k_data = (const char *) k->data + (ic + tk) * nbk1 + ik2 * nbk2 + ik3 * nbk3;
1505
+ float * k_col = K_f32 + tk;
1506
+ const float * k_src = (const float *) k_data;
1507
+ for (int64_t dk = 0; dk < DK; ++dk) {
1508
+ k_col[dk * KV_TILE_SZ] = k_src[dk];
1509
+ }
1510
+ }
1511
+
1512
+ for (int tq = 0; tq < tile_rows; ++tq) {
1513
+ rvv_qk_dot_tile(KQ + tq * KV_TILE_SZ, Q_f32 + tq * DK, K_f32, DK, KV_TILE_SZ, scale);
1514
+ }
1515
+ }
1516
+
1517
+ // Set padded KQ entries to -inf so softmax gives them zero weight
1518
+ if (kv_tile < KV_TILE_SZ) {
1519
+ for (int tq = 0; tq < tile_rows; tq++) {
1520
+ for (int tk = kv_tile; tk < KV_TILE_SZ; tk++) {
1521
+ KQ[tq * KV_TILE_SZ + tk] = -INFINITY;
1522
+ }
1523
+ }
1524
+ }
1525
+
1526
+ if (logit_softcap != 0.0f) {
1527
+ rvv_softcap_tanh_inplace_f32(KQ, KV_TILE_SZ, tile_rows, KV_TILE_SZ, logit_softcap);
1528
+ }
1529
+
1530
+ if (mask) {
1531
+ rvv_add_inplace_f32(KQ, KV_TILE_SZ, mask32, KV_TILE_SZ, tile_rows, KV_TILE_SZ);
1532
+ }
1533
+
1534
+ bool skip[Q_TILE_SZ] = {};
1535
+
1536
+ for (int tq = 0; tq < tile_rows; tq++) {
1537
+ float * kq_row = KQ + tq * KV_TILE_SZ;
1538
+
1539
+ const float tile_max = rvv_max_f32(kq_row, KV_TILE_SZ);
1540
+
1541
+ if (tile_max == -INFINITY) {
1542
+ skip[tq] = true;
1543
+ continue;
1544
+ }
1545
+
1546
+ const float Mold = M[tq];
1547
+ const float Mnew = fmaxf(Mold, tile_max);
1548
+
1549
+ if (Mnew > Mold) {
1550
+ const float ms = expf(Mold - Mnew);
1551
+ rvv_scale_f32(VKQ32 + tq * DV, ms, DV);
1552
+ S[tq] *= ms;
1553
+ }
1554
+ M[tq] = Mnew;
1555
+
1556
+ S[tq] += rvv_softmax_exp_inplace_f32(kq_row, KV_TILE_SZ, Mnew);
1557
+ }
1558
+
1559
+ // Pack V as contiguous [KV_TILE_SZ][DV].
1560
+ if (kv_type == GGML_TYPE_F16) {
1561
+ const char * v_data = (const char *) v->data + ic * nbv1 + iv2 * nbv2 + iv3 * nbv3;
1562
+ memcpy2d(V_f16, DV * sizeof(_Float16), v_data, nbv1, kv_tile, DV * sizeof(_Float16));
1563
+
1564
+ int tq = 0;
1565
+ for (; tq + 3 < tile_rows; tq += 4) {
1566
+ if (skip[tq + 0] || skip[tq + 1] || skip[tq + 2] || skip[tq + 3]) {
1567
+ for (int i = 0; i < 4; ++i) {
1568
+ if (!skip[tq + i]) {
1569
+ rvv_pv_accumulate_f16_x1(VKQ32 + (tq + i) * DV, KQ + (tq + i) * KV_TILE_SZ, V_f16,
1570
+ KV_TILE_SZ, DV);
1571
+ }
1572
+ }
1573
+ continue;
1574
+ }
1575
+
1576
+ rvv_pv_accumulate_f16_x4(VKQ32 + (tq + 0) * DV, VKQ32 + (tq + 1) * DV, VKQ32 + (tq + 2) * DV,
1577
+ VKQ32 + (tq + 3) * DV, KQ + (tq + 0) * KV_TILE_SZ,
1578
+ KQ + (tq + 1) * KV_TILE_SZ, KQ + (tq + 2) * KV_TILE_SZ,
1579
+ KQ + (tq + 3) * KV_TILE_SZ, V_f16, KV_TILE_SZ, DV);
1580
+ }
1581
+ for (; tq < tile_rows; ++tq) {
1582
+ if (!skip[tq]) {
1583
+ rvv_pv_accumulate_f16_x1(VKQ32 + tq * DV, KQ + tq * KV_TILE_SZ, V_f16, KV_TILE_SZ, DV);
1584
+ }
1585
+ }
1586
+ } else {
1587
+ const char * v_data = (const char *) v->data + ic * nbv1 + iv2 * nbv2 + iv3 * nbv3;
1588
+ memcpy2d(V32, DV * sizeof(float), v_data, nbv1, kv_tile, DV * sizeof(float));
1589
+
1590
+ for (int tq = 0; tq < tile_rows; ++tq) {
1591
+ if (!skip[tq]) {
1592
+ rvv_pv_accumulate(VKQ32 + tq * DV, KQ + tq * KV_TILE_SZ, V32, KV_TILE_SZ, DV);
1593
+ }
1594
+ }
1595
+ }
1596
+ }
1597
+
1598
+ // sinks (apply only to valid rows in the tile)
1599
+ if (sinks) {
1600
+ const float s = ((float *) ((char *) sinks->data))[h];
1601
+
1602
+ for (int tq = 0; tq < tile_rows; tq++) {
1603
+ float ms = 1.0f;
1604
+ float vs = 1.0f;
1605
+
1606
+ if (s > M[tq]) {
1607
+ ms = expf(M[tq] - s);
1608
+ rvv_scale_f32(VKQ32 + tq * DV, ms, DV);
1609
+ } else {
1610
+ vs = expf(s - M[tq]);
1611
+ }
1612
+
1613
+ float S_temp = S[tq] * ms + vs;
1614
+ S[tq] = S_temp == 0.0f ? 0.0f : 1.0f / S_temp;
1615
+ }
1616
+ } else {
1617
+ for (int tq = 0; tq < tile_rows; tq++) {
1618
+ const float S_inv = S[tq] == 0.0f ? 0.0f : 1.0f / S[tq];
1619
+ S[tq] = S_inv;
1620
+ }
1621
+ }
1622
+
1623
+ float * dst_ptr = (float *) ((char *) dst->data + (iq3 * ne2 * ne1 + iq2 + (iq1) *ne1) * nb1);
1624
+ rvv_pack_scaled_f32_as_f32(dst_ptr, nb1 * ne1, VKQ32, DV * sizeof(float), tile_rows, DV, S);
1625
+
1626
+ ir += tile_rows;
1627
+ }
1628
+ }
1629
+
1630
+ void forward_rms_norm_f32(ggml_compute_params * params, ggml_tensor * op) {
1631
+ const ggml_tensor * src0 = op->src[0];
1632
+ ggml_tensor * dst = op;
1633
+ GGML_ASSERT(ggml_are_same_shape(src0, dst));
1634
+ GGML_ASSERT(src0->nb[0] == sizeof(float));
1635
+
1636
+ int ith = params->ith;
1637
+ int nth = params->nth;
1638
+
1639
+ GGML_TENSOR_UNARY_OP_LOCALS
1640
+
1641
+ float epsilon = *((float *) dst->op_params);
1642
+
1643
+ GGML_ASSERT(epsilon > 0.0f);
1644
+
1645
+ auto * input = (char *) src0->data;
1646
+ auto * output = (char *) dst->data;
1647
+
1648
+ const auto hidden_size = ne00;
1649
+ const auto task_count = ne01 * ne02 * ne03;
1650
+ const auto task_per_thread = (task_count + nth - 1) / nth;
1651
+
1652
+ const auto task_begin = ith * task_per_thread;
1653
+ const auto task_end = std::min((ith + 1) * task_per_thread, task_count);
1654
+
1655
+ for (auto task_idx = task_begin; task_idx < task_end; task_idx++) {
1656
+ int64_t i03 = task_idx / (ne02 * ne01);
1657
+ int64_t i02 = (task_idx - i03 * ne02 * ne01) / ne01;
1658
+ int64_t i01 = (task_idx - i03 * ne02 * ne01 - i02 * ne01);
1659
+
1660
+ auto * p_input = (float *) (input + i01 * nb01 + i02 * nb02 + i03 * nb03);
1661
+ auto * p_output = (float *) (output + i01 * nb1 + i02 * nb2 + i03 * nb3);
1662
+ auto * p_temp_output = p_output;
1663
+
1664
+ size_t gvl = __riscv_vsetvlmax_e32m4();
1665
+ vfloat32m4_t sum_sq = __riscv_vfmv_v_f_f32m4(0.f, gvl);
1666
+ int64_t length = hidden_size;
1667
+ while (length > 0) {
1668
+ gvl = __riscv_vsetvl_e32m4(length);
1669
+ vfloat32m4_t src_data = __riscv_vle32_v_f32m4(p_input, gvl);
1670
+ sum_sq = __riscv_vfmacc_vv_f32m4(sum_sq, src_data, src_data, gvl);
1671
+ __riscv_vse32_v_f32m4(p_temp_output, src_data, gvl);
1672
+
1673
+ p_input += gvl;
1674
+ p_temp_output += gvl;
1675
+ length -= gvl;
1676
+ }
1677
+
1678
+ gvl = __riscv_vsetvlmax_e32m1();
1679
+ vfloat32m1_t zero_v = __riscv_vfmv_v_f_f32m1(0.f, gvl);
1680
+ vfloat32m1_t mean_square_v =
1681
+ __riscv_vfadd_vv_f32m1(__riscv_vget_v_f32m4_f32m1(sum_sq, 0), __riscv_vget_v_f32m4_f32m1(sum_sq, 1), gvl);
1682
+
1683
+ mean_square_v = __riscv_vfadd_vv_f32m1(mean_square_v, __riscv_vget_v_f32m4_f32m1(sum_sq, 2), gvl);
1684
+ mean_square_v = __riscv_vfadd_vv_f32m1(mean_square_v, __riscv_vget_v_f32m4_f32m1(sum_sq, 3), gvl);
1685
+ mean_square_v = __riscv_vfredusum_vs_f32m1_f32m1(mean_square_v, zero_v, gvl);
1686
+
1687
+ float mean_square = __riscv_vfmv_f_s_f32m1_f32(mean_square_v);
1688
+ mean_square /= hidden_size;
1689
+
1690
+ mean_square = sqrt(mean_square + epsilon);
1691
+
1692
+ mean_square = 1.0f / mean_square;
1693
+ length = hidden_size;
1694
+ p_temp_output = p_output;
1695
+
1696
+ while (length > 0) {
1697
+ gvl = __riscv_vsetvl_e32m4(length);
1698
+ vfloat32m4_t src_data = __riscv_vle32_v_f32m4(p_temp_output, gvl);
1699
+ src_data = __riscv_vfmul_vf_f32m4(src_data, mean_square, gvl);
1700
+ __riscv_vse32_v_f32m4(p_output, src_data, gvl);
1701
+ p_temp_output += gvl;
1702
+ p_output += gvl;
1703
+ length -= gvl;
1704
+ }
1705
+ }
1706
+ }
1707
+
1708
+ template <size_t MB_ROWS>
1709
+ void quantize_a_nrow_i8_ref(size_t blk_len, const float * a_ptr, size_t count_k, uint8_t * quant_a_ptr) {
1710
+ int64_t a_blk_stride = q8_blk_size(blk_len, true);
1711
+ int64_t a_nrow_block_stride = a_blk_stride * MB_ROWS;
1712
+ for (size_t k = 0; k < count_k; k += blk_len, quant_a_ptr += a_nrow_block_stride) {
1713
+ float * scale_a_ptr = reinterpret_cast<float *>(quant_a_ptr);
1714
+ int16_t * a_sum_ptr = reinterpret_cast<int16_t *>(quant_a_ptr + sizeof(float) * MB_ROWS);
1715
+ int8_t * quant_a_blk =
1716
+ reinterpret_cast<int8_t *>(quant_a_ptr + sizeof(float) * MB_ROWS + sizeof(int16_t) * MB_ROWS);
1717
+
1718
+ for (size_t row = 0; row < MB_ROWS; row++) {
1719
+ float max_abs_a = 0.0f;
1720
+ for (size_t bk = 0; bk < blk_len; bk++) {
1721
+ max_abs_a = std::max(max_abs_a, std::abs(a_ptr[row * count_k + k + bk]));
1722
+ }
1723
+
1724
+ float rep_scale_a = ((1 << 7) - 1) / max_abs_a;
1725
+ scale_a_ptr[row] = 1 / rep_scale_a;
1726
+
1727
+ int16_t a_sum = 0;
1728
+ for (size_t bk = 0; bk < blk_len; bk++) {
1729
+ const int8_t quantized = static_cast<int8_t>(
1730
+ std::clamp(std::nearbyintf(a_ptr[row * count_k + k + bk] * rep_scale_a), -128.0f, 127.0f));
1731
+ quant_a_blk[row * blk_len + bk] = quantized;
1732
+ a_sum += quantized;
1733
+ }
1734
+ a_sum_ptr[row] = -a_sum;
1735
+ }
1736
+ }
1737
+ }
1738
+
1739
+ template <size_t MB_ROWS>
1740
+ void quantize_a_nrow_i8_hp_ref(size_t blk_len, const float * a_ptr, size_t count_k, uint8_t * quant_a_ptr) {
1741
+ constexpr size_t k_subblk_len = 32;
1742
+ const size_t subblk_count = blk_len / k_subblk_len;
1743
+
1744
+ GGML_ASSERT(blk_len == 256);
1745
+
1746
+ float scale_temp[8] = { 0.0f };
1747
+ int64_t a_blk_stride = q8_hp_blk_size(blk_len, true, true);
1748
+ int64_t a_nrow_block_stride = a_blk_stride * MB_ROWS;
1749
+ int64_t a_subblk_stride = q8_hp_blk_size(k_subblk_len, false, false) * MB_ROWS;
1750
+
1751
+ for (size_t k = 0; k < count_k; k += blk_len, quant_a_ptr += a_nrow_block_stride) {
1752
+ _Float16 * a_sum_ptr = reinterpret_cast<_Float16 *>(quant_a_ptr + a_subblk_stride * subblk_count);
1753
+
1754
+ float scale_avg = 0.0f;
1755
+ for (size_t kk = 0; kk < subblk_count; kk++) {
1756
+ float max_abs_a = 0.0f;
1757
+ for (size_t row = 0; row < MB_ROWS; row++) {
1758
+ for (size_t bk = 0; bk < k_subblk_len; bk++) {
1759
+ max_abs_a = std::max(max_abs_a, std::abs(a_ptr[row * count_k + k + bk + kk * k_subblk_len]));
1760
+ }
1761
+ }
1762
+ scale_temp[kk] = max_abs_a / ((1 << 7) - 1);
1763
+ scale_avg += scale_temp[kk];
1764
+ }
1765
+
1766
+ scale_avg /= subblk_count;
1767
+ float scale_factor = 1.0f / scale_avg;
1768
+
1769
+ _Float16 * scale_avg_ptr =
1770
+ reinterpret_cast<_Float16 *>(quant_a_ptr + a_nrow_block_stride - sizeof(_Float16) * MB_ROWS);
1771
+ scale_avg_ptr[0] = scale_avg;
1772
+
1773
+ for (size_t kk = 0; kk < subblk_count; kk++) {
1774
+ uint8_t * a_subblk_base = quant_a_ptr + kk * a_subblk_stride;
1775
+ _Float16 * scale_a_ptr = reinterpret_cast<_Float16 *>(a_subblk_base);
1776
+ int8_t * quant_a_blk = reinterpret_cast<int8_t *>(a_subblk_base + sizeof(_Float16) * MB_ROWS);
1777
+
1778
+ scale_a_ptr[0] = static_cast<_Float16>(scale_temp[kk] * scale_factor);
1779
+
1780
+ const float rep_scale_a = 1.0f / scale_temp[kk];
1781
+
1782
+ for (size_t row = 0; row < MB_ROWS; row++) {
1783
+ int16_t a_sum = 0;
1784
+ for (size_t bk = 0; bk < k_subblk_len; bk++) {
1785
+ const int8_t quantized = static_cast<int8_t>(
1786
+ std::clamp(std::nearbyintf(a_ptr[row * count_k + k + bk + kk * k_subblk_len] * rep_scale_a),
1787
+ -128.0f, 127.0f));
1788
+ quant_a_blk[row * k_subblk_len + bk] = quantized;
1789
+ a_sum += quantized;
1790
+ }
1791
+ a_sum_ptr[row * subblk_count + kk] = static_cast<_Float16>(-a_sum) * static_cast<_Float16>(8.0f);
1792
+ }
1793
+ }
1794
+ }
1795
+ }
1796
+
1797
+ template <size_t MB_ROWS>
1798
+ void quantize_a_nrow_i8k_ref(size_t blk_len, const float * a_ptr, size_t count_k, uint8_t * quant_a_ptr) {
1799
+ int64_t a_blk_stride = q8k_blk_size(256);
1800
+ int64_t a_nrow_block_stride = a_blk_stride * MB_ROWS;
1801
+ int64_t a_sum_size = 256 / 16;
1802
+
1803
+ for (size_t k = 0; k < count_k; k += blk_len, quant_a_ptr += a_nrow_block_stride) {
1804
+ float * scale_a_ptr = reinterpret_cast<float *>(quant_a_ptr);
1805
+ int16_t * a_sum_ptr = reinterpret_cast<int16_t *>(quant_a_ptr + sizeof(float) * MB_ROWS);
1806
+ int8_t * quant_a_blk =
1807
+ reinterpret_cast<int8_t *>(quant_a_ptr + sizeof(float) * MB_ROWS + sizeof(int16_t) * a_sum_size * MB_ROWS);
1808
+
1809
+ for (size_t row = 0; row < MB_ROWS; row++) {
1810
+ float max_a = 0.0f;
1811
+ float max_abs_a = 0.0f;
1812
+ for (size_t bk = 0; bk < blk_len; bk++) {
1813
+ float ax = std::abs(a_ptr[row * count_k + k + bk]);
1814
+ if (ax > max_abs_a) {
1815
+ max_abs_a = ax;
1816
+ max_a = a_ptr[row * count_k + k + bk];
1817
+ }
1818
+ }
1819
+
1820
+ if (!max_abs_a) {
1821
+ scale_a_ptr[row] = 0;
1822
+ for (size_t bki = 0; bki < a_sum_size; bki++) {
1823
+ for (size_t bk = bki * 16; bk < (bki + 1) * 16; bk++) {
1824
+ quant_a_blk[row * blk_len + bk] = 0;
1825
+ }
1826
+ a_sum_ptr[row * a_sum_size + bki] = 0;
1827
+ }
1828
+ continue;
1829
+ }
1830
+
1831
+ float rep_scale_a = ((1 << 7) - 1) / max_abs_a;
1832
+ scale_a_ptr[row] = 1 / rep_scale_a;
1833
+
1834
+ for (size_t bki = 0; bki < a_sum_size; bki++) {
1835
+ int16_t a_sum = 0;
1836
+ for (size_t bk = bki * 16; bk < (bki + 1) * 16; bk++) {
1837
+ const int8_t quantized = static_cast<int8_t>(
1838
+ std::clamp(std::nearbyintf(a_ptr[row * count_k + k + bk] * rep_scale_a), -128.0f, 127.0f));
1839
+ quant_a_blk[row * blk_len + bk] = quantized;
1840
+ a_sum += quantized;
1841
+ }
1842
+ a_sum_ptr[row * a_sum_size + bki] = -a_sum;
1843
+ }
1844
+ }
1845
+ }
1846
+ }
1847
+
1848
+ void quantize_a_row_i8(size_t blk_len, const float * a_ptr, size_t count_k, uint8_t * quant_a_ptr) {
1849
+ GGML_ASSERT(blk_len == 32);
1850
+ int64_t a_blk_stride = q8_blk_size(blk_len, true);
1851
+ size_t vlenb = __riscv_vlenb();
1852
+
1853
+ if (vlenb == 128) {
1854
+ for (size_t k = 0; k < count_k; k += blk_len, quant_a_ptr += a_blk_stride) {
1855
+ float * scale_a_ptr = reinterpret_cast<float *>(quant_a_ptr);
1856
+ int16_t * a_sum_ptr = reinterpret_cast<int16_t *>(quant_a_ptr + sizeof(float));
1857
+ int8_t * quant_a_blk = reinterpret_cast<int8_t *>(quant_a_ptr + sizeof(float) + sizeof(int16_t));
1858
+
1859
+ size_t vl = __riscv_vsetvl_e32m1(blk_len);
1860
+ vfloat32m1_t v_a = __riscv_vle32_v_f32m1(a_ptr + k, vl);
1861
+ vfloat32m1_t v_a_abs = __riscv_vfabs_v_f32m1(v_a, vl);
1862
+
1863
+ vfloat32m1_t tmp = __riscv_vfmv_v_f_f32m1(0.0f, vl);
1864
+ vfloat32m1_t v_a_max = __riscv_vfredmax_vs_f32m1_f32m1(v_a_abs, tmp, vl);
1865
+ float max_abs_a = __riscv_vfmv_f_s_f32m1_f32(v_a_max);
1866
+
1867
+ float scale_a = max_abs_a / ((1 << 7) - 1);
1868
+ float rep_scale_a = scale_a ? 1.0f / scale_a : 0.0f;
1869
+ scale_a_ptr[0] = scale_a;
1870
+
1871
+ vfloat32m1_t v_a_scale = __riscv_vfmul_vf_f32m1(v_a, rep_scale_a, vl);
1872
+ vint16mf2_t v_a_quant = __riscv_vfncvt_x_f_w_i16mf2(v_a_scale, vl);
1873
+ vint8mf4_t v_a_quant_i8 = __riscv_vncvt_x_x_w_i8mf4(v_a_quant, vl);
1874
+
1875
+ vint16m1_t tmp_sum = __riscv_vmv_v_x_i16m1(0, vl);
1876
+ vint16m1_t v_a_sum = __riscv_vwredsum_vs_i8mf4_i16m1(v_a_quant_i8, tmp_sum, vl);
1877
+ int16_t a_sum = __riscv_vmv_x_s_i16m1_i16(v_a_sum);
1878
+ a_sum_ptr[0] = -a_sum;
1879
+
1880
+ __riscv_vse8_v_i8mf4(quant_a_blk, v_a_quant_i8, vl);
1881
+ }
1882
+ } else if (vlenb == 32) {
1883
+ for (size_t k = 0; k < count_k; k += blk_len, quant_a_ptr += a_blk_stride) {
1884
+ float * scale_a_ptr = reinterpret_cast<float *>(quant_a_ptr);
1885
+ int16_t * a_sum_ptr = reinterpret_cast<int16_t *>(quant_a_ptr + sizeof(float));
1886
+ int8_t * quant_a_blk = reinterpret_cast<int8_t *>(quant_a_ptr + sizeof(float) + sizeof(int16_t));
1887
+
1888
+ size_t vl = __riscv_vsetvl_e32m4(blk_len);
1889
+ vfloat32m4_t v_a = __riscv_vle32_v_f32m4(a_ptr + k, vl);
1890
+ vfloat32m4_t v_a_abs = __riscv_vfabs_v_f32m4(v_a, vl);
1891
+
1892
+ vfloat32m1_t tmp = __riscv_vfmv_v_f_f32m1(0.0f, vl);
1893
+ vfloat32m1_t v_a_max = __riscv_vfredmax_vs_f32m4_f32m1(v_a_abs, tmp, vl);
1894
+ float max_abs_a = __riscv_vfmv_f_s_f32m1_f32(v_a_max);
1895
+
1896
+ float scale_a = max_abs_a / ((1 << 7) - 1);
1897
+ float rep_scale_a = scale_a ? 1.0f / scale_a : 0.0f;
1898
+ scale_a_ptr[0] = scale_a;
1899
+
1900
+ vfloat32m4_t v_a_scale = __riscv_vfmul_vf_f32m4(v_a, rep_scale_a, vl);
1901
+ vint16m2_t v_a_quant = __riscv_vfncvt_x_f_w_i16m2(v_a_scale, vl);
1902
+ vint8m1_t v_a_quant_i8 = __riscv_vncvt_x_x_w_i8m1(v_a_quant, vl);
1903
+
1904
+ vint16m1_t tmp_sum = __riscv_vmv_v_x_i16m1(0, vl);
1905
+ vint16m1_t v_a_sum = __riscv_vwredsum_vs_i8m1_i16m1(v_a_quant_i8, tmp_sum, vl);
1906
+ int16_t a_sum = __riscv_vmv_x_s_i16m1_i16(v_a_sum);
1907
+ a_sum_ptr[0] = -a_sum;
1908
+
1909
+ __riscv_vse8_v_i8m1(quant_a_blk, v_a_quant_i8, vl);
1910
+ }
1911
+ } else {
1912
+ quantize_a_nrow_i8_ref<1>(blk_len, a_ptr, count_k, quant_a_ptr);
1913
+ }
1914
+ }
1915
+
1916
+ void quantize_a_4row_i8(size_t blk_len, const float * a_ptr, size_t count_k, uint8_t * quant_a_ptr) {
1917
+ GGML_ASSERT(blk_len == 32);
1918
+ int64_t a_blk_stride = q8_blk_size(blk_len, true);
1919
+ int64_t a_nrow_block_stride = a_blk_stride * 4;
1920
+ size_t vlenb = __riscv_vlenb();
1921
+
1922
+ if (vlenb == 128) {
1923
+ for (size_t k = 0; k < count_k; k += blk_len, quant_a_ptr += a_nrow_block_stride) {
1924
+ float * scale_a_ptr = reinterpret_cast<float *>(quant_a_ptr);
1925
+ int16_t * a_sum_ptr = reinterpret_cast<int16_t *>(quant_a_ptr + sizeof(float) * 4);
1926
+ int8_t * quant_a_blk = reinterpret_cast<int8_t *>(quant_a_ptr + sizeof(float) * 4 + sizeof(int16_t) * 4);
1927
+
1928
+ for (size_t mi = 0; mi < 4; mi++) {
1929
+ size_t vl = __riscv_vsetvl_e32m1(blk_len);
1930
+ vfloat32m1_t v_a = __riscv_vle32_v_f32m1(a_ptr + mi * count_k + k, vl);
1931
+ vfloat32m1_t v_a_abs = __riscv_vfabs_v_f32m1(v_a, vl);
1932
+
1933
+ vfloat32m1_t tmp = __riscv_vfmv_v_f_f32m1(0.0f, vl);
1934
+ vfloat32m1_t v_a_max = __riscv_vfredmax_vs_f32m1_f32m1(v_a_abs, tmp, vl);
1935
+ float max_abs_a = __riscv_vfmv_f_s_f32m1_f32(v_a_max);
1936
+
1937
+ float scale_a = max_abs_a / ((1 << 7) - 1);
1938
+ float rep_scale_a = scale_a ? 1.0f / scale_a : 0.0f;
1939
+ scale_a_ptr[mi] = scale_a;
1940
+
1941
+ vfloat32m1_t v_a_scale = __riscv_vfmul_vf_f32m1(v_a, rep_scale_a, vl);
1942
+ vint16mf2_t v_a_quant = __riscv_vfncvt_x_f_w_i16mf2(v_a_scale, vl);
1943
+ vint8mf4_t v_a_quant_i8 = __riscv_vncvt_x_x_w_i8mf4(v_a_quant, vl);
1944
+
1945
+ vint16m1_t tmp_sum = __riscv_vmv_v_x_i16m1(0, vl);
1946
+ vint16m1_t v_a_sum = __riscv_vwredsum_vs_i8mf4_i16m1(v_a_quant_i8, tmp_sum, vl);
1947
+ int16_t a_sum = __riscv_vmv_x_s_i16m1_i16(v_a_sum);
1948
+ a_sum_ptr[mi] = -a_sum;
1949
+
1950
+ __riscv_vse8_v_i8mf4(quant_a_blk + mi * blk_len, v_a_quant_i8, vl);
1951
+ }
1952
+ }
1953
+ } else if (vlenb == 32) {
1954
+ for (size_t k = 0; k < count_k; k += blk_len, quant_a_ptr += a_nrow_block_stride) {
1955
+ float * scale_a_ptr = reinterpret_cast<float *>(quant_a_ptr);
1956
+ int16_t * a_sum_ptr = reinterpret_cast<int16_t *>(quant_a_ptr + sizeof(float) * 4);
1957
+ int8_t * quant_a_blk = reinterpret_cast<int8_t *>(quant_a_ptr + sizeof(float) * 4 + sizeof(int16_t) * 4);
1958
+
1959
+ for (size_t mi = 0; mi < 4; mi++) {
1960
+ size_t vl = __riscv_vsetvl_e32m4(blk_len);
1961
+ vfloat32m4_t v_a = __riscv_vle32_v_f32m4(a_ptr + mi * count_k + k, vl);
1962
+ vfloat32m4_t v_a_abs = __riscv_vfabs_v_f32m4(v_a, vl);
1963
+
1964
+ vfloat32m1_t tmp = __riscv_vfmv_v_f_f32m1(0.0f, vl);
1965
+ vfloat32m1_t v_a_max = __riscv_vfredmax_vs_f32m4_f32m1(v_a_abs, tmp, vl);
1966
+ float max_abs_a = __riscv_vfmv_f_s_f32m1_f32(v_a_max);
1967
+
1968
+ float scale_a = max_abs_a / ((1 << 7) - 1);
1969
+ float rep_scale_a = scale_a ? 1.0f / scale_a : 0.0f;
1970
+ scale_a_ptr[mi] = scale_a;
1971
+
1972
+ vfloat32m4_t v_a_scale = __riscv_vfmul_vf_f32m4(v_a, rep_scale_a, vl);
1973
+ vint16m2_t v_a_quant = __riscv_vfncvt_x_f_w_i16m2(v_a_scale, vl);
1974
+ vint8m1_t v_a_quant_i8 = __riscv_vncvt_x_x_w_i8m1(v_a_quant, vl);
1975
+
1976
+ vint16m1_t tmp_sum = __riscv_vmv_v_x_i16m1(0, vl);
1977
+ vint16m1_t v_a_sum = __riscv_vwredsum_vs_i8m1_i16m1(v_a_quant_i8, tmp_sum, vl);
1978
+ int16_t a_sum = __riscv_vmv_x_s_i16m1_i16(v_a_sum);
1979
+ a_sum_ptr[mi] = -a_sum;
1980
+
1981
+ __riscv_vse8_v_i8m1(quant_a_blk + mi * blk_len, v_a_quant_i8, vl);
1982
+ }
1983
+ }
1984
+ } else {
1985
+ quantize_a_nrow_i8_ref<4>(blk_len, a_ptr, count_k, quant_a_ptr);
1986
+ }
1987
+ }
1988
+
1989
+ void quantize_a_row_i8_hp(size_t blk_len, const float * a_ptr, size_t count_k, uint8_t * quant_a_ptr) {
1990
+ constexpr size_t k_subblk_len = 32;
1991
+ GGML_ASSERT(blk_len == 256);
1992
+
1993
+ constexpr size_t subblk_count = 256 / k_subblk_len;
1994
+ int64_t a_blk_stride = q8_hp_blk_size(blk_len, true, true);
1995
+ int64_t a_subblk_stride = q8_hp_blk_size(k_subblk_len, false, false);
1996
+ size_t vlenb = __riscv_vlenb();
1997
+ float scale_temp[subblk_count] = { 0.0f };
1998
+
1999
+ if (vlenb == 128) {
2000
+ for (size_t k = 0; k < count_k; k += blk_len, quant_a_ptr += a_blk_stride) {
2001
+ _Float16 * a_sum_ptr = reinterpret_cast<_Float16 *>(quant_a_ptr + a_subblk_stride * subblk_count);
2002
+ _Float16 * scale_avg_ptr = reinterpret_cast<_Float16 *>(quant_a_ptr + a_blk_stride - sizeof(_Float16));
2003
+ float scale_avg = 0.0f;
2004
+
2005
+ for (size_t kk = 0; kk < subblk_count; ++kk) {
2006
+ const float * a_src_ptr = a_ptr + k + kk * k_subblk_len;
2007
+
2008
+ size_t vl = __riscv_vsetvl_e32m1(k_subblk_len);
2009
+ vfloat32m1_t v_a = __riscv_vle32_v_f32m1(a_src_ptr, vl);
2010
+ vfloat32m1_t v_a_abs = __riscv_vfabs_v_f32m1(v_a, vl);
2011
+
2012
+ vfloat32m1_t tmp = __riscv_vfmv_v_f_f32m1(0.0f, vl);
2013
+ vfloat32m1_t v_a_max = __riscv_vfredmax_vs_f32m1_f32m1(v_a_abs, tmp, vl);
2014
+ float max_abs_a = __riscv_vfmv_f_s_f32m1_f32(v_a_max);
2015
+
2016
+ scale_temp[kk] = max_abs_a / ((1 << 7) - 1);
2017
+ scale_avg += scale_temp[kk];
2018
+ }
2019
+
2020
+ scale_avg /= subblk_count;
2021
+ const float scale_factor = scale_avg ? 1.0f / scale_avg : 0.0f;
2022
+ scale_avg_ptr[0] = static_cast<_Float16>(scale_avg);
2023
+
2024
+ for (size_t kk = 0; kk < subblk_count; ++kk) {
2025
+ uint8_t * a_subblk_base = quant_a_ptr + kk * a_subblk_stride;
2026
+ _Float16 * scale_a_ptr = reinterpret_cast<_Float16 *>(a_subblk_base);
2027
+ int8_t * quant_a_blk = reinterpret_cast<int8_t *>(a_subblk_base + sizeof(_Float16));
2028
+ const float * a_src_ptr = a_ptr + k + kk * k_subblk_len;
2029
+
2030
+ size_t vl = __riscv_vsetvl_e32m1(k_subblk_len);
2031
+ vfloat32m1_t v_a = __riscv_vle32_v_f32m1(a_src_ptr, vl);
2032
+ float rep_scale_a = scale_temp[kk] ? 1.0f / scale_temp[kk] : 0.0f;
2033
+ scale_a_ptr[0] = static_cast<_Float16>(scale_temp[kk] * scale_factor);
2034
+
2035
+ vfloat32m1_t v_a_scale = __riscv_vfmul_vf_f32m1(v_a, rep_scale_a, vl);
2036
+ vint16mf2_t v_a_quant = __riscv_vfncvt_x_f_w_i16mf2(v_a_scale, vl);
2037
+ vint8mf4_t v_a_quant_i8 = __riscv_vncvt_x_x_w_i8mf4(v_a_quant, vl);
2038
+
2039
+ vint16m1_t tmp_sum = __riscv_vmv_v_x_i16m1(0, vl);
2040
+ vint16m1_t v_a_sum = __riscv_vwredsum_vs_i8mf4_i16m1(v_a_quant_i8, tmp_sum, vl);
2041
+ int16_t a_sum = __riscv_vmv_x_s_i16m1_i16(v_a_sum);
2042
+ a_sum_ptr[kk] = static_cast<_Float16>(-a_sum) * static_cast<_Float16>(8.0f);
2043
+
2044
+ __riscv_vse8_v_i8mf4(quant_a_blk, v_a_quant_i8, vl);
2045
+ }
2046
+ }
2047
+ } else if (vlenb == 32) {
2048
+ for (size_t k = 0; k < count_k; k += blk_len, quant_a_ptr += a_blk_stride) {
2049
+ _Float16 * a_sum_ptr = reinterpret_cast<_Float16 *>(quant_a_ptr + a_subblk_stride * subblk_count);
2050
+ _Float16 * scale_avg_ptr = reinterpret_cast<_Float16 *>(quant_a_ptr + a_blk_stride - sizeof(_Float16));
2051
+ float scale_avg = 0.0f;
2052
+
2053
+ for (size_t kk = 0; kk < subblk_count; ++kk) {
2054
+ const float * a_src_ptr = a_ptr + k + kk * k_subblk_len;
2055
+
2056
+ size_t vl = __riscv_vsetvl_e32m4(k_subblk_len);
2057
+ vfloat32m4_t v_a = __riscv_vle32_v_f32m4(a_src_ptr, vl);
2058
+ vfloat32m4_t v_a_abs = __riscv_vfabs_v_f32m4(v_a, vl);
2059
+
2060
+ vfloat32m1_t tmp = __riscv_vfmv_v_f_f32m1(0.0f, vl);
2061
+ vfloat32m1_t v_a_max = __riscv_vfredmax_vs_f32m4_f32m1(v_a_abs, tmp, vl);
2062
+ float max_abs_a = __riscv_vfmv_f_s_f32m1_f32(v_a_max);
2063
+
2064
+ scale_temp[kk] = max_abs_a / ((1 << 7) - 1);
2065
+ scale_avg += scale_temp[kk];
2066
+ }
2067
+
2068
+ scale_avg /= subblk_count;
2069
+ const float scale_factor = scale_avg ? 1.0f / scale_avg : 0.0f;
2070
+ scale_avg_ptr[0] = static_cast<_Float16>(scale_avg);
2071
+
2072
+ for (size_t kk = 0; kk < subblk_count; ++kk) {
2073
+ uint8_t * a_subblk_base = quant_a_ptr + kk * a_subblk_stride;
2074
+ _Float16 * scale_a_ptr = reinterpret_cast<_Float16 *>(a_subblk_base);
2075
+ int8_t * quant_a_blk = reinterpret_cast<int8_t *>(a_subblk_base + sizeof(_Float16));
2076
+ const float * a_src_ptr = a_ptr + k + kk * k_subblk_len;
2077
+
2078
+ size_t vl = __riscv_vsetvl_e32m4(k_subblk_len);
2079
+ vfloat32m4_t v_a = __riscv_vle32_v_f32m4(a_src_ptr, vl);
2080
+ float rep_scale_a = scale_temp[kk] ? 1.0f / scale_temp[kk] : 0.0f;
2081
+ scale_a_ptr[0] = static_cast<_Float16>(scale_temp[kk] * scale_factor);
2082
+
2083
+ vfloat32m4_t v_a_scale = __riscv_vfmul_vf_f32m4(v_a, rep_scale_a, vl);
2084
+ vint16m2_t v_a_quant = __riscv_vfncvt_x_f_w_i16m2(v_a_scale, vl);
2085
+ vint8m1_t v_a_quant_i8 = __riscv_vncvt_x_x_w_i8m1(v_a_quant, vl);
2086
+
2087
+ vint16m1_t tmp_sum = __riscv_vmv_v_x_i16m1(0, vl);
2088
+ vint16m1_t v_a_sum = __riscv_vwredsum_vs_i8m1_i16m1(v_a_quant_i8, tmp_sum, vl);
2089
+ int16_t a_sum = __riscv_vmv_x_s_i16m1_i16(v_a_sum);
2090
+ a_sum_ptr[kk] = static_cast<_Float16>(-a_sum) * static_cast<_Float16>(8.0f);
2091
+
2092
+ __riscv_vse8_v_i8m1(quant_a_blk, v_a_quant_i8, vl);
2093
+ }
2094
+ }
2095
+ } else {
2096
+ quantize_a_nrow_i8_hp_ref<1>(blk_len, a_ptr, count_k, quant_a_ptr);
2097
+ }
2098
+ }
2099
+
2100
+ void quantize_a_4row_i8_hp(size_t blk_len, const float * a_ptr, size_t count_k, uint8_t * quant_a_ptr) {
2101
+ constexpr size_t k_subblk_len = 32;
2102
+ GGML_ASSERT(blk_len == 256);
2103
+
2104
+ constexpr size_t subblk_count = 256 / k_subblk_len;
2105
+ int64_t a_blk_stride = q8_hp_blk_size(blk_len, true, true);
2106
+ int64_t a_nrow_block_stride = a_blk_stride * 4;
2107
+ int64_t a_subblk_stride = q8_hp_blk_size(k_subblk_len, false, false) * 4;
2108
+ size_t vlenb = __riscv_vlenb();
2109
+ float scale_temp[subblk_count] = { 0.0f };
2110
+
2111
+ if (vlenb == 128) {
2112
+ for (size_t k = 0; k < count_k; k += blk_len, quant_a_ptr += a_nrow_block_stride) {
2113
+ _Float16 * a_sum_ptr = reinterpret_cast<_Float16 *>(quant_a_ptr + a_subblk_stride * subblk_count);
2114
+ _Float16 * scale_avg_ptr =
2115
+ reinterpret_cast<_Float16 *>(quant_a_ptr + a_nrow_block_stride - sizeof(_Float16) * 4);
2116
+ float scale_avg = 0.0f;
2117
+
2118
+ for (size_t kk = 0; kk < subblk_count; ++kk) {
2119
+ const float * a_src_ptr0 = a_ptr + 0 * count_k + k + kk * k_subblk_len;
2120
+ const float * a_src_ptr1 = a_ptr + 1 * count_k + k + kk * k_subblk_len;
2121
+ const float * a_src_ptr2 = a_ptr + 2 * count_k + k + kk * k_subblk_len;
2122
+ const float * a_src_ptr3 = a_ptr + 3 * count_k + k + kk * k_subblk_len;
2123
+
2124
+ size_t vl = __riscv_vsetvl_e32m1(k_subblk_len);
2125
+ vfloat32m1_t v_a0 = __riscv_vle32_v_f32m1(a_src_ptr0, vl);
2126
+ vfloat32m1_t v_a1 = __riscv_vle32_v_f32m1(a_src_ptr1, vl);
2127
+ vfloat32m1_t v_a2 = __riscv_vle32_v_f32m1(a_src_ptr2, vl);
2128
+ vfloat32m1_t v_a3 = __riscv_vle32_v_f32m1(a_src_ptr3, vl);
2129
+ vfloat32m1_t v_a0_abs = __riscv_vfabs_v_f32m1(v_a0, vl);
2130
+ vfloat32m1_t v_a1_abs = __riscv_vfabs_v_f32m1(v_a1, vl);
2131
+ vfloat32m1_t v_a2_abs = __riscv_vfabs_v_f32m1(v_a2, vl);
2132
+ vfloat32m1_t v_a3_abs = __riscv_vfabs_v_f32m1(v_a3, vl);
2133
+
2134
+ vfloat32m1_t v_max_abs = __riscv_vfmax_vv_f32m1(v_a0_abs, v_a1_abs, vl);
2135
+ v_max_abs = __riscv_vfmax_vv_f32m1(v_max_abs, v_a2_abs, vl);
2136
+ v_max_abs = __riscv_vfmax_vv_f32m1(v_max_abs, v_a3_abs, vl);
2137
+
2138
+ vfloat32m1_t tmp = __riscv_vfmv_v_f_f32m1(0.0f, vl);
2139
+ vfloat32m1_t v_a_max = __riscv_vfredmax_vs_f32m1_f32m1(v_max_abs, tmp, vl);
2140
+ float max_abs_a = __riscv_vfmv_f_s_f32m1_f32(v_a_max);
2141
+
2142
+ scale_temp[kk] = max_abs_a / ((1 << 7) - 1);
2143
+ scale_avg += scale_temp[kk];
2144
+ }
2145
+
2146
+ scale_avg /= subblk_count;
2147
+ const float scale_factor = scale_avg ? 1.0f / scale_avg : 0.0f;
2148
+ scale_avg_ptr[0] = static_cast<_Float16>(scale_avg);
2149
+
2150
+ for (size_t kk = 0; kk < subblk_count; ++kk) {
2151
+ uint8_t * a_subblk_base = quant_a_ptr + kk * a_subblk_stride;
2152
+ _Float16 * scale_a_ptr = reinterpret_cast<_Float16 *>(a_subblk_base);
2153
+ int8_t * quant_a_blk = reinterpret_cast<int8_t *>(a_subblk_base + sizeof(_Float16) * 4);
2154
+ const float * a_src_ptr0 = a_ptr + 0 * count_k + k + kk * k_subblk_len;
2155
+ const float * a_src_ptr1 = a_ptr + 1 * count_k + k + kk * k_subblk_len;
2156
+ const float * a_src_ptr2 = a_ptr + 2 * count_k + k + kk * k_subblk_len;
2157
+ const float * a_src_ptr3 = a_ptr + 3 * count_k + k + kk * k_subblk_len;
2158
+
2159
+ size_t vl = __riscv_vsetvl_e32m1(k_subblk_len);
2160
+ vfloat32m1_t v_a0 = __riscv_vle32_v_f32m1(a_src_ptr0, vl);
2161
+ vfloat32m1_t v_a1 = __riscv_vle32_v_f32m1(a_src_ptr1, vl);
2162
+ vfloat32m1_t v_a2 = __riscv_vle32_v_f32m1(a_src_ptr2, vl);
2163
+ vfloat32m1_t v_a3 = __riscv_vle32_v_f32m1(a_src_ptr3, vl);
2164
+
2165
+ float rep_scale_a = scale_temp[kk] ? 1.0f / scale_temp[kk] : 0.0f;
2166
+ scale_a_ptr[0] = static_cast<_Float16>(scale_temp[kk] * scale_factor);
2167
+
2168
+ vfloat32m1_t v_a0_scale = __riscv_vfmul_vf_f32m1(v_a0, rep_scale_a, vl);
2169
+ vfloat32m1_t v_a1_scale = __riscv_vfmul_vf_f32m1(v_a1, rep_scale_a, vl);
2170
+ vfloat32m1_t v_a2_scale = __riscv_vfmul_vf_f32m1(v_a2, rep_scale_a, vl);
2171
+ vfloat32m1_t v_a3_scale = __riscv_vfmul_vf_f32m1(v_a3, rep_scale_a, vl);
2172
+ vint16mf2_t v_a0_quant = __riscv_vfncvt_x_f_w_i16mf2(v_a0_scale, vl);
2173
+ vint16mf2_t v_a1_quant = __riscv_vfncvt_x_f_w_i16mf2(v_a1_scale, vl);
2174
+ vint16mf2_t v_a2_quant = __riscv_vfncvt_x_f_w_i16mf2(v_a2_scale, vl);
2175
+ vint16mf2_t v_a3_quant = __riscv_vfncvt_x_f_w_i16mf2(v_a3_scale, vl);
2176
+ vint8mf4_t v_a0_quant_i8 = __riscv_vncvt_x_x_w_i8mf4(v_a0_quant, vl);
2177
+ vint8mf4_t v_a1_quant_i8 = __riscv_vncvt_x_x_w_i8mf4(v_a1_quant, vl);
2178
+ vint8mf4_t v_a2_quant_i8 = __riscv_vncvt_x_x_w_i8mf4(v_a2_quant, vl);
2179
+ vint8mf4_t v_a3_quant_i8 = __riscv_vncvt_x_x_w_i8mf4(v_a3_quant, vl);
2180
+
2181
+ vint16m1_t tmp_sum0 = __riscv_vmv_v_x_i16m1(0, vl);
2182
+ vint16m1_t tmp_sum1 = __riscv_vmv_v_x_i16m1(0, vl);
2183
+ vint16m1_t tmp_sum2 = __riscv_vmv_v_x_i16m1(0, vl);
2184
+ vint16m1_t tmp_sum3 = __riscv_vmv_v_x_i16m1(0, vl);
2185
+ vint16m1_t v_a0_sum = __riscv_vwredsum_vs_i8mf4_i16m1(v_a0_quant_i8, tmp_sum0, vl);
2186
+ vint16m1_t v_a1_sum = __riscv_vwredsum_vs_i8mf4_i16m1(v_a1_quant_i8, tmp_sum1, vl);
2187
+ vint16m1_t v_a2_sum = __riscv_vwredsum_vs_i8mf4_i16m1(v_a2_quant_i8, tmp_sum2, vl);
2188
+ vint16m1_t v_a3_sum = __riscv_vwredsum_vs_i8mf4_i16m1(v_a3_quant_i8, tmp_sum3, vl);
2189
+
2190
+ a_sum_ptr[0 * subblk_count + kk] =
2191
+ static_cast<_Float16>(-__riscv_vmv_x_s_i16m1_i16(v_a0_sum)) * static_cast<_Float16>(8.0f);
2192
+ a_sum_ptr[1 * subblk_count + kk] =
2193
+ static_cast<_Float16>(-__riscv_vmv_x_s_i16m1_i16(v_a1_sum)) * static_cast<_Float16>(8.0f);
2194
+ a_sum_ptr[2 * subblk_count + kk] =
2195
+ static_cast<_Float16>(-__riscv_vmv_x_s_i16m1_i16(v_a2_sum)) * static_cast<_Float16>(8.0f);
2196
+ a_sum_ptr[3 * subblk_count + kk] =
2197
+ static_cast<_Float16>(-__riscv_vmv_x_s_i16m1_i16(v_a3_sum)) * static_cast<_Float16>(8.0f);
2198
+
2199
+ __riscv_vse8_v_i8mf4(quant_a_blk + 0 * k_subblk_len, v_a0_quant_i8, vl);
2200
+ __riscv_vse8_v_i8mf4(quant_a_blk + 1 * k_subblk_len, v_a1_quant_i8, vl);
2201
+ __riscv_vse8_v_i8mf4(quant_a_blk + 2 * k_subblk_len, v_a2_quant_i8, vl);
2202
+ __riscv_vse8_v_i8mf4(quant_a_blk + 3 * k_subblk_len, v_a3_quant_i8, vl);
2203
+ }
2204
+ }
2205
+ } else if (vlenb == 32) {
2206
+ for (size_t k = 0; k < count_k; k += blk_len, quant_a_ptr += a_nrow_block_stride) {
2207
+ _Float16 * a_sum_ptr = reinterpret_cast<_Float16 *>(quant_a_ptr + a_subblk_stride * subblk_count);
2208
+ _Float16 * scale_avg_ptr =
2209
+ reinterpret_cast<_Float16 *>(quant_a_ptr + a_nrow_block_stride - sizeof(_Float16) * 4);
2210
+ float scale_avg = 0.0f;
2211
+
2212
+ for (size_t kk = 0; kk < subblk_count; ++kk) {
2213
+ const float * a_src_ptr0 = a_ptr + 0 * count_k + k + kk * k_subblk_len;
2214
+ const float * a_src_ptr1 = a_ptr + 1 * count_k + k + kk * k_subblk_len;
2215
+ const float * a_src_ptr2 = a_ptr + 2 * count_k + k + kk * k_subblk_len;
2216
+ const float * a_src_ptr3 = a_ptr + 3 * count_k + k + kk * k_subblk_len;
2217
+
2218
+ size_t vl = __riscv_vsetvl_e32m4(k_subblk_len);
2219
+ vfloat32m4_t v_a0 = __riscv_vle32_v_f32m4(a_src_ptr0, vl);
2220
+ vfloat32m4_t v_a1 = __riscv_vle32_v_f32m4(a_src_ptr1, vl);
2221
+ vfloat32m4_t v_a2 = __riscv_vle32_v_f32m4(a_src_ptr2, vl);
2222
+ vfloat32m4_t v_a3 = __riscv_vle32_v_f32m4(a_src_ptr3, vl);
2223
+
2224
+ vfloat32m4_t v_a0_abs = __riscv_vfabs_v_f32m4(v_a0, vl);
2225
+ vfloat32m4_t v_a1_abs = __riscv_vfabs_v_f32m4(v_a1, vl);
2226
+ vfloat32m4_t v_a2_abs = __riscv_vfabs_v_f32m4(v_a2, vl);
2227
+ vfloat32m4_t v_a3_abs = __riscv_vfabs_v_f32m4(v_a3, vl);
2228
+
2229
+ vfloat32m4_t v_max_abs = __riscv_vfmax_vv_f32m4(v_a0_abs, v_a1_abs, vl);
2230
+ v_max_abs = __riscv_vfmax_vv_f32m4(v_max_abs, v_a2_abs, vl);
2231
+ v_max_abs = __riscv_vfmax_vv_f32m4(v_max_abs, v_a3_abs, vl);
2232
+
2233
+ vfloat32m1_t tmp = __riscv_vfmv_v_f_f32m1(0.0f, vl);
2234
+ vfloat32m1_t v_a_max = __riscv_vfredmax_vs_f32m4_f32m1(v_max_abs, tmp, vl);
2235
+ float max_abs_a = __riscv_vfmv_f_s_f32m1_f32(v_a_max);
2236
+
2237
+ scale_temp[kk] = max_abs_a / ((1 << 7) - 1);
2238
+ scale_avg += scale_temp[kk];
2239
+ }
2240
+
2241
+ scale_avg /= subblk_count;
2242
+ const float scale_factor = scale_avg ? 1.0f / scale_avg : 0.0f;
2243
+ scale_avg_ptr[0] = static_cast<_Float16>(scale_avg);
2244
+
2245
+ for (size_t kk = 0; kk < subblk_count; ++kk) {
2246
+ uint8_t * a_subblk_base = quant_a_ptr + kk * a_subblk_stride;
2247
+ _Float16 * scale_a_ptr = reinterpret_cast<_Float16 *>(a_subblk_base);
2248
+ int8_t * quant_a_blk = reinterpret_cast<int8_t *>(a_subblk_base + sizeof(_Float16) * 4);
2249
+ const float * a_src_ptr0 = a_ptr + 0 * count_k + k + kk * k_subblk_len;
2250
+ const float * a_src_ptr1 = a_ptr + 1 * count_k + k + kk * k_subblk_len;
2251
+ const float * a_src_ptr2 = a_ptr + 2 * count_k + k + kk * k_subblk_len;
2252
+ const float * a_src_ptr3 = a_ptr + 3 * count_k + k + kk * k_subblk_len;
2253
+
2254
+ size_t vl = __riscv_vsetvl_e32m4(k_subblk_len);
2255
+ vfloat32m4_t v_a0 = __riscv_vle32_v_f32m4(a_src_ptr0, vl);
2256
+ vfloat32m4_t v_a1 = __riscv_vle32_v_f32m4(a_src_ptr1, vl);
2257
+ vfloat32m4_t v_a2 = __riscv_vle32_v_f32m4(a_src_ptr2, vl);
2258
+ vfloat32m4_t v_a3 = __riscv_vle32_v_f32m4(a_src_ptr3, vl);
2259
+
2260
+ float rep_scale_a = scale_temp[kk] ? 1.0f / scale_temp[kk] : 0.0f;
2261
+ scale_a_ptr[0] = static_cast<_Float16>(scale_temp[kk] * scale_factor);
2262
+
2263
+ vfloat32m4_t v_a0_scale = __riscv_vfmul_vf_f32m4(v_a0, rep_scale_a, vl);
2264
+ vfloat32m4_t v_a1_scale = __riscv_vfmul_vf_f32m4(v_a1, rep_scale_a, vl);
2265
+ vfloat32m4_t v_a2_scale = __riscv_vfmul_vf_f32m4(v_a2, rep_scale_a, vl);
2266
+ vfloat32m4_t v_a3_scale = __riscv_vfmul_vf_f32m4(v_a3, rep_scale_a, vl);
2267
+ vint16m2_t v_a0_quant = __riscv_vfncvt_x_f_w_i16m2(v_a0_scale, vl);
2268
+ vint16m2_t v_a1_quant = __riscv_vfncvt_x_f_w_i16m2(v_a1_scale, vl);
2269
+ vint16m2_t v_a2_quant = __riscv_vfncvt_x_f_w_i16m2(v_a2_scale, vl);
2270
+ vint16m2_t v_a3_quant = __riscv_vfncvt_x_f_w_i16m2(v_a3_scale, vl);
2271
+ vint8m1_t v_a0_quant_i8 = __riscv_vncvt_x_x_w_i8m1(v_a0_quant, vl);
2272
+ vint8m1_t v_a1_quant_i8 = __riscv_vncvt_x_x_w_i8m1(v_a1_quant, vl);
2273
+ vint8m1_t v_a2_quant_i8 = __riscv_vncvt_x_x_w_i8m1(v_a2_quant, vl);
2274
+ vint8m1_t v_a3_quant_i8 = __riscv_vncvt_x_x_w_i8m1(v_a3_quant, vl);
2275
+
2276
+ vint16m1_t tmp_sum0 = __riscv_vmv_v_x_i16m1(0, vl);
2277
+ vint16m1_t tmp_sum1 = __riscv_vmv_v_x_i16m1(0, vl);
2278
+ vint16m1_t tmp_sum2 = __riscv_vmv_v_x_i16m1(0, vl);
2279
+ vint16m1_t tmp_sum3 = __riscv_vmv_v_x_i16m1(0, vl);
2280
+ vint16m1_t v_a0_sum = __riscv_vwredsum_vs_i8m1_i16m1(v_a0_quant_i8, tmp_sum0, vl);
2281
+ vint16m1_t v_a1_sum = __riscv_vwredsum_vs_i8m1_i16m1(v_a1_quant_i8, tmp_sum1, vl);
2282
+ vint16m1_t v_a2_sum = __riscv_vwredsum_vs_i8m1_i16m1(v_a2_quant_i8, tmp_sum2, vl);
2283
+ vint16m1_t v_a3_sum = __riscv_vwredsum_vs_i8m1_i16m1(v_a3_quant_i8, tmp_sum3, vl);
2284
+
2285
+ a_sum_ptr[0 * subblk_count + kk] =
2286
+ static_cast<_Float16>(-__riscv_vmv_x_s_i16m1_i16(v_a0_sum)) * static_cast<_Float16>(8.0f);
2287
+ a_sum_ptr[1 * subblk_count + kk] =
2288
+ static_cast<_Float16>(-__riscv_vmv_x_s_i16m1_i16(v_a1_sum)) * static_cast<_Float16>(8.0f);
2289
+ a_sum_ptr[2 * subblk_count + kk] =
2290
+ static_cast<_Float16>(-__riscv_vmv_x_s_i16m1_i16(v_a2_sum)) * static_cast<_Float16>(8.0f);
2291
+ a_sum_ptr[3 * subblk_count + kk] =
2292
+ static_cast<_Float16>(-__riscv_vmv_x_s_i16m1_i16(v_a3_sum)) * static_cast<_Float16>(8.0f);
2293
+
2294
+ __riscv_vse8_v_i8m1(quant_a_blk + 0 * k_subblk_len, v_a0_quant_i8, vl);
2295
+ __riscv_vse8_v_i8m1(quant_a_blk + 1 * k_subblk_len, v_a1_quant_i8, vl);
2296
+ __riscv_vse8_v_i8m1(quant_a_blk + 2 * k_subblk_len, v_a2_quant_i8, vl);
2297
+ __riscv_vse8_v_i8m1(quant_a_blk + 3 * k_subblk_len, v_a3_quant_i8, vl);
2298
+ }
2299
+ }
2300
+ } else {
2301
+ quantize_a_nrow_i8_hp_ref<4>(blk_len, a_ptr, count_k, quant_a_ptr);
2302
+ }
2303
+ }
2304
+
2305
+ void quantize_a_row_i8k(size_t blk_len, const float * a_ptr, size_t count_k, uint8_t * quant_a_ptr) {
2306
+ GGML_ASSERT(blk_len == 256);
2307
+ constexpr int64_t a_blk_stride = q8k_blk_size(256);
2308
+ constexpr int64_t a_sum_size = 256 / 16;
2309
+ size_t vlenb = __riscv_vlenb();
2310
+
2311
+ if (vlenb == 128) {
2312
+ // vlen = 1024 bits, can process 32 float32 elements with m1
2313
+ for (size_t k = 0; k < count_k; k += blk_len, quant_a_ptr += a_blk_stride) {
2314
+ float * scale_a_ptr = reinterpret_cast<float *>(quant_a_ptr);
2315
+ int16_t * a_sum_ptr = reinterpret_cast<int16_t *>(quant_a_ptr + sizeof(float));
2316
+ int8_t * quant_a_blk =
2317
+ reinterpret_cast<int8_t *>(quant_a_ptr + sizeof(float) + sizeof(int16_t) * a_sum_size);
2318
+
2319
+ // Find max absolute value across all 256 elements
2320
+ size_t vl = __riscv_vsetvl_e32m1(16);
2321
+ vfloat32m1_t v_max_abs = __riscv_vfmv_v_f_f32m1(0.0f, vl);
2322
+
2323
+ for (size_t bki = 0; bki < a_sum_size; bki++) {
2324
+ vfloat32m1_t v_a = __riscv_vle32_v_f32m1(a_ptr + k + bki * 16, vl);
2325
+ vfloat32m1_t v_a_abs = __riscv_vfabs_v_f32m1(v_a, vl);
2326
+ v_max_abs = __riscv_vfmax_vv_f32m1(v_a_abs, v_max_abs, vl);
2327
+ }
2328
+ vfloat32m1_t tmp = __riscv_vfmv_v_f_f32m1(0.0f, vl);
2329
+ vfloat32m1_t v_local_max = __riscv_vfredmax_vs_f32m1_f32m1(v_max_abs, tmp, vl);
2330
+ float max_abs_a = __riscv_vfmv_f_s_f32m1_f32(v_local_max);
2331
+
2332
+ float scale_a = max_abs_a / ((1 << 7) - 1);
2333
+ float rep_scale_a = scale_a ? 1.0f / scale_a : 0.0f;
2334
+ scale_a_ptr[0] = scale_a;
2335
+
2336
+ // Quantize and compute sums for each 16-element group
2337
+ for (size_t bki = 0; bki < a_sum_size; bki++) {
2338
+ vfloat32m1_t v_a = __riscv_vle32_v_f32m1(a_ptr + k + bki * 16, vl);
2339
+ vfloat32m1_t v_a_scale = __riscv_vfmul_vf_f32m1(v_a, rep_scale_a, vl);
2340
+ vint16mf2_t v_a_quant = __riscv_vfncvt_x_f_w_i16mf2(v_a_scale, vl);
2341
+ vint8mf4_t v_a_quant_i8 = __riscv_vncvt_x_x_w_i8mf4(v_a_quant, vl);
2342
+
2343
+ vint16m1_t tmp_sum = __riscv_vmv_v_x_i16m1(0, vl);
2344
+ vint16m1_t v_a_sum = __riscv_vwredsum_vs_i8mf4_i16m1(v_a_quant_i8, tmp_sum, vl);
2345
+ int16_t a_sum = __riscv_vmv_x_s_i16m1_i16(v_a_sum);
2346
+ a_sum_ptr[bki] = -a_sum;
2347
+
2348
+ __riscv_vse8_v_i8mf4(quant_a_blk + bki * 16, v_a_quant_i8, vl);
2349
+ }
2350
+ }
2351
+ } else if (vlenb == 32) {
2352
+ // vlen = 256 bits, can process 8 float32 elements with m1
2353
+ for (size_t k = 0; k < count_k; k += blk_len, quant_a_ptr += a_blk_stride) {
2354
+ float * scale_a_ptr = reinterpret_cast<float *>(quant_a_ptr);
2355
+ int16_t * a_sum_ptr = reinterpret_cast<int16_t *>(quant_a_ptr + sizeof(float));
2356
+ int8_t * quant_a_blk =
2357
+ reinterpret_cast<int8_t *>(quant_a_ptr + sizeof(float) + sizeof(int16_t) * a_sum_size);
2358
+
2359
+ // Find max absolute value across all 256 elements
2360
+ size_t vl = __riscv_vsetvl_e32m2(16);
2361
+ vfloat32m2_t v_max_abs = __riscv_vfmv_v_f_f32m2(0.0f, vl);
2362
+
2363
+ for (size_t bki = 0; bki < a_sum_size; bki++) {
2364
+ vfloat32m2_t v_a = __riscv_vle32_v_f32m2(a_ptr + k + bki * 16, vl);
2365
+ vfloat32m2_t v_a_abs = __riscv_vfabs_v_f32m2(v_a, vl);
2366
+ v_max_abs = __riscv_vfmax_vv_f32m2(v_a_abs, v_max_abs, vl);
2367
+ }
2368
+ vfloat32m1_t tmp = __riscv_vfmv_v_f_f32m1(0.0f, vl);
2369
+ vfloat32m1_t v_local_max = __riscv_vfredmax_vs_f32m2_f32m1(v_max_abs, tmp, vl);
2370
+ float max_abs_a = __riscv_vfmv_f_s_f32m1_f32(v_local_max);
2371
+
2372
+ float scale_a = max_abs_a / ((1 << 7) - 1);
2373
+ float rep_scale_a = scale_a ? 1.0f / scale_a : 0.0f;
2374
+ scale_a_ptr[0] = scale_a;
2375
+
2376
+ // Quantize and compute sums for each 16-element group
2377
+ for (size_t bki = 0; bki < a_sum_size; bki++) {
2378
+ vfloat32m2_t v_a = __riscv_vle32_v_f32m2(a_ptr + k + bki * 16, vl);
2379
+ vfloat32m2_t v_a_scale = __riscv_vfmul_vf_f32m2(v_a, rep_scale_a, vl);
2380
+ vint16m1_t v_a_quant = __riscv_vfncvt_x_f_w_i16m1(v_a_scale, vl);
2381
+ vint8mf2_t v_a_quant_i8 = __riscv_vncvt_x_x_w_i8mf2(v_a_quant, vl);
2382
+
2383
+ vint16m1_t tmp_sum = __riscv_vmv_v_x_i16m1(0, vl);
2384
+ vint16m1_t v_a_sum = __riscv_vwredsum_vs_i8mf2_i16m1(v_a_quant_i8, tmp_sum, vl);
2385
+ int16_t a_sum = __riscv_vmv_x_s_i16m1_i16(v_a_sum);
2386
+ a_sum_ptr[bki] = -a_sum;
2387
+
2388
+ __riscv_vse8_v_i8mf2(quant_a_blk + bki * 16, v_a_quant_i8, vl);
2389
+ }
2390
+ }
2391
+ } else {
2392
+ quantize_a_nrow_i8k_ref<1>(blk_len, a_ptr, count_k, quant_a_ptr);
2393
+ }
2394
+ }
2395
+
2396
+ void quantize_a_4row_i8k(size_t blk_len, const float * a_ptr, size_t count_k, uint8_t * quant_a_ptr) {
2397
+ GGML_ASSERT(blk_len == 256);
2398
+ constexpr int64_t a_blk_stride = q8k_blk_size(256);
2399
+ constexpr int64_t a_nrow_block_stride = a_blk_stride * 4;
2400
+ constexpr int64_t a_sum_size = 256 / 16;
2401
+ size_t vlenb = __riscv_vlenb();
2402
+
2403
+ if (vlenb == 128) {
2404
+ // vlen = 1024 bits
2405
+ for (size_t k = 0; k < count_k; k += blk_len, quant_a_ptr += a_nrow_block_stride) {
2406
+ float * scale_a_ptr = reinterpret_cast<float *>(quant_a_ptr);
2407
+ int16_t * a_sum_ptr = reinterpret_cast<int16_t *>(quant_a_ptr + sizeof(float) * 4);
2408
+ int8_t * quant_a_blk =
2409
+ reinterpret_cast<int8_t *>(quant_a_ptr + sizeof(float) * 4 + sizeof(int16_t) * a_sum_size * 4);
2410
+
2411
+ for (size_t mi = 0; mi < 4; mi++) {
2412
+ // Find max absolute value across all 256 elements for this row
2413
+ size_t vl = __riscv_vsetvl_e32m1(16);
2414
+ vfloat32m1_t v_max_abs = __riscv_vfmv_v_f_f32m1(0.0f, vl);
2415
+
2416
+ for (size_t bki = 0; bki < a_sum_size; bki++) {
2417
+ vfloat32m1_t v_a = __riscv_vle32_v_f32m1(a_ptr + mi * count_k + k + bki * 16, vl);
2418
+ vfloat32m1_t v_a_abs = __riscv_vfabs_v_f32m1(v_a, vl);
2419
+ v_max_abs = __riscv_vfmax_vv_f32m1(v_a_abs, v_max_abs, vl);
2420
+ }
2421
+ vfloat32m1_t tmp = __riscv_vfmv_v_f_f32m1(0.0f, vl);
2422
+ vfloat32m1_t v_local_max = __riscv_vfredmax_vs_f32m1_f32m1(v_max_abs, tmp, vl);
2423
+ float max_abs_a = __riscv_vfmv_f_s_f32m1_f32(v_local_max);
2424
+
2425
+ float scale_a = max_abs_a / ((1 << 7) - 1);
2426
+ float rep_scale_a = scale_a ? 1.0f / scale_a : 0.0f;
2427
+ scale_a_ptr[mi] = scale_a;
2428
+
2429
+ // Quantize and compute sums for each 16-element group
2430
+ for (size_t bki = 0; bki < a_sum_size; bki++) {
2431
+ vfloat32m1_t v_a = __riscv_vle32_v_f32m1(a_ptr + mi * count_k + k + bki * 16, vl);
2432
+ vfloat32m1_t v_a_scale = __riscv_vfmul_vf_f32m1(v_a, rep_scale_a, vl);
2433
+ vint16mf2_t v_a_quant = __riscv_vfncvt_x_f_w_i16mf2(v_a_scale, vl);
2434
+ vint8mf4_t v_a_quant_i8 = __riscv_vncvt_x_x_w_i8mf4(v_a_quant, vl);
2435
+
2436
+ vint16m1_t tmp_sum = __riscv_vmv_v_x_i16m1(0, vl);
2437
+ vint16m1_t v_a_sum = __riscv_vwredsum_vs_i8mf4_i16m1(v_a_quant_i8, tmp_sum, vl);
2438
+ int16_t a_sum = __riscv_vmv_x_s_i16m1_i16(v_a_sum);
2439
+ a_sum_ptr[mi * a_sum_size + bki] = -a_sum;
2440
+
2441
+ __riscv_vse8_v_i8mf4(quant_a_blk + mi * blk_len + bki * 16, v_a_quant_i8, vl);
2442
+ }
2443
+ }
2444
+ }
2445
+ } else if (vlenb == 32) {
2446
+ // vlen = 256 bits
2447
+ for (size_t k = 0; k < count_k; k += blk_len, quant_a_ptr += a_nrow_block_stride) {
2448
+ float * scale_a_ptr = reinterpret_cast<float *>(quant_a_ptr);
2449
+ int16_t * a_sum_ptr = reinterpret_cast<int16_t *>(quant_a_ptr + sizeof(float) * 4);
2450
+ int8_t * quant_a_blk =
2451
+ reinterpret_cast<int8_t *>(quant_a_ptr + sizeof(float) * 4 + sizeof(int16_t) * a_sum_size * 4);
2452
+
2453
+ for (size_t mi = 0; mi < 4; mi++) {
2454
+ // Find max absolute value across all 256 elements for this row
2455
+ size_t vl = __riscv_vsetvl_e32m2(16);
2456
+ vfloat32m2_t v_max_abs = __riscv_vfmv_v_f_f32m2(0.0f, vl);
2457
+
2458
+ for (size_t bki = 0; bki < a_sum_size; bki++) {
2459
+ vfloat32m2_t v_a = __riscv_vle32_v_f32m2(a_ptr + mi * count_k + k + bki * 16, vl);
2460
+ vfloat32m2_t v_a_abs = __riscv_vfabs_v_f32m2(v_a, vl);
2461
+ v_max_abs = __riscv_vfmax_vv_f32m2(v_a_abs, v_max_abs, vl);
2462
+ }
2463
+ vfloat32m1_t tmp = __riscv_vfmv_v_f_f32m1(0.0f, vl);
2464
+ vfloat32m1_t v_local_max = __riscv_vfredmax_vs_f32m2_f32m1(v_max_abs, tmp, vl);
2465
+ float max_abs_a = __riscv_vfmv_f_s_f32m1_f32(v_local_max);
2466
+
2467
+ float scale_a = max_abs_a / ((1 << 7) - 1);
2468
+ float rep_scale_a = scale_a ? 1.0f / scale_a : 0.0f;
2469
+ scale_a_ptr[mi] = scale_a;
2470
+
2471
+ // Quantize and compute sums for each 16-element group
2472
+ for (size_t bki = 0; bki < a_sum_size; bki++) {
2473
+ vfloat32m2_t v_a = __riscv_vle32_v_f32m2(a_ptr + mi * count_k + k + bki * 16, vl);
2474
+ vfloat32m2_t v_a_scale = __riscv_vfmul_vf_f32m2(v_a, rep_scale_a, vl);
2475
+ vint16m1_t v_a_quant = __riscv_vfncvt_x_f_w_i16m1(v_a_scale, vl);
2476
+ vint8mf2_t v_a_quant_i8 = __riscv_vncvt_x_x_w_i8mf2(v_a_quant, vl);
2477
+
2478
+ vint16m1_t tmp_sum = __riscv_vmv_v_x_i16m1(0, vl);
2479
+ vint16m1_t v_a_sum = __riscv_vwredsum_vs_i8mf2_i16m1(v_a_quant_i8, tmp_sum, vl);
2480
+ int16_t a_sum = __riscv_vmv_x_s_i16m1_i16(v_a_sum);
2481
+ a_sum_ptr[mi * a_sum_size + bki] = -a_sum;
2482
+
2483
+ __riscv_vse8_v_i8mf2(quant_a_blk + mi * blk_len + bki * 16, v_a_quant_i8, vl);
2484
+ }
2485
+ }
2486
+ }
2487
+ } else {
2488
+ quantize_a_nrow_i8k_ref<4>(blk_len, a_ptr, count_k, quant_a_ptr);
2489
+ }
2490
+ }
2491
+
2492
+ void forward_cpy_with_permute(ggml_compute_params * params, ggml_tensor * op) {
2493
+ const ggml_tensor * src0 = op->src[0];
2494
+ ggml_tensor * dst = op;
2495
+ const int ith = params->ith;
2496
+ const int nth = params->nth;
2497
+
2498
+ // [batch, m, n] -> [batch, n, m]
2499
+ int64_t batch = src0->ne[2] * src0->ne[3];
2500
+ int64_t m = src0->ne[1];
2501
+ int64_t n = src0->ne[0];
2502
+
2503
+ int64_t batch_stride = src0->nb[2];
2504
+ int64_t m_src_stride = src0->nb[0];
2505
+ int64_t n_src_stride = src0->nb[1];
2506
+ int64_t n_dst_stride = n_src_stride * m;
2507
+
2508
+ permute_transpose_impl(src0, dst, batch, m, n, batch_stride, m_src_stride, n_src_stride, n_dst_stride, ith, nth);
2509
+ }
2510
+
2511
+ void forward_cont_with_permute(ggml_compute_params * params, ggml_tensor * op) {
2512
+ const ggml_tensor * src0 = op->src[0];
2513
+ ggml_tensor * dst = op;
2514
+ const int ith = params->ith;
2515
+ const int nth = params->nth;
2516
+
2517
+ // [batch, m, n] -> [batch, n, m]
2518
+ int64_t batch = dst->ne[2] * dst->ne[3];
2519
+ int64_t n = dst->ne[1];
2520
+ int64_t m = dst->ne[0];
2521
+
2522
+ int64_t batch_stride = dst->nb[2];
2523
+ int64_t m_src_stride = src0->nb[0];
2524
+ int64_t n_src_stride = src0->nb[1];
2525
+ int64_t n_dst_stride = dst->nb[1];
2526
+
2527
+ permute_transpose_impl(src0, dst, batch, m, n, batch_stride, m_src_stride, n_src_stride, n_dst_stride, ith, nth);
2528
+ }
2529
+
2530
+ void forward_norm_f32(ggml_compute_params * params, ggml_tensor * op) {
2531
+ const ggml_tensor * src0 = op->src[0];
2532
+ ggml_tensor * dst = op;
2533
+ GGML_ASSERT(ggml_are_same_shape(src0, dst));
2534
+ GGML_ASSERT(src0->nb[0] == sizeof(float));
2535
+
2536
+ int ith = params->ith;
2537
+ int nth = params->nth;
2538
+
2539
+ GGML_TENSOR_UNARY_OP_LOCALS
2540
+
2541
+ float epsilon = *((float *) dst->op_params);
2542
+
2543
+ GGML_ASSERT(epsilon > 0.0f);
2544
+
2545
+ auto * input = (char *) src0->data;
2546
+ auto * output = (char *) dst->data;
2547
+
2548
+ const auto hidden_size = ne00;
2549
+ const auto task_count = ne01 * ne02 * ne03;
2550
+ const auto task_per_thread = (task_count + nth - 1) / nth;
2551
+
2552
+ const auto task_begin = ith * task_per_thread;
2553
+ const auto task_end = std::min((ith + 1) * task_per_thread, task_count);
2554
+
2555
+ for (auto task_idx = task_begin; task_idx < task_end; task_idx++) {
2556
+ int64_t i03 = task_idx / (ne02 * ne01);
2557
+ int64_t i02 = (task_idx - i03 * ne02 * ne01) / ne01;
2558
+ int64_t i01 = (task_idx - i03 * ne02 * ne01 - i02 * ne01);
2559
+
2560
+ auto * p_input = (float *) (input + i01 * nb01 + i02 * nb02 + i03 * nb03);
2561
+ auto * p_output = (float *) (output + i01 * nb1 + i02 * nb2 + i03 * nb3);
2562
+ auto * p_temp_output = p_output;
2563
+
2564
+ size_t gvl = __riscv_vsetvlmax_e32m4();
2565
+ vfloat32m4_t sum = __riscv_vfmv_v_f_f32m4(0.f, gvl);
2566
+ vfloat32m4_t sum_sq = __riscv_vfmv_v_f_f32m4(0.f, gvl);
2567
+ int64_t length = hidden_size;
2568
+ while (length > 0) {
2569
+ gvl = __riscv_vsetvl_e32m4(length);
2570
+ // load data
2571
+ vfloat32m4_t src_data = __riscv_vle32_v_f32m4(p_input, gvl);
2572
+
2573
+ sum = __riscv_vfadd_vv_f32m4(sum, src_data, gvl);
2574
+ sum_sq = __riscv_vfmacc_vv_f32m4(sum_sq, src_data, src_data, gvl);
2575
+
2576
+ __riscv_vse32_v_f32m4(p_temp_output, src_data, gvl);
2577
+
2578
+ p_input += gvl;
2579
+ p_temp_output += gvl;
2580
+ length -= gvl;
2581
+ }
2582
+
2583
+ gvl = __riscv_vsetvlmax_e32m1();
2584
+
2585
+ float mean = 0.f;
2586
+ vfloat32m1_t zero_v = __riscv_vfmv_v_f_f32m1(0.f, gvl);
2587
+ vfloat32m1_t mean_v =
2588
+ __riscv_vfadd_vv_f32m1(__riscv_vget_v_f32m4_f32m1(sum, 0), __riscv_vget_v_f32m4_f32m1(sum, 1), gvl);
2589
+ mean_v = __riscv_vfadd_vv_f32m1(mean_v, __riscv_vget_v_f32m4_f32m1(sum, 2), gvl);
2590
+ mean_v = __riscv_vfadd_vv_f32m1(mean_v, __riscv_vget_v_f32m4_f32m1(sum, 3), gvl);
2591
+ mean_v = __riscv_vfredusum_vs_f32m1_f32m1(mean_v, zero_v, gvl);
2592
+ mean = __riscv_vfmv_f_s_f32m1_f32(mean_v);
2593
+ mean /= hidden_size;
2594
+
2595
+ vfloat32m1_t mean_square_v =
2596
+ __riscv_vfadd_vv_f32m1(__riscv_vget_v_f32m4_f32m1(sum_sq, 0), __riscv_vget_v_f32m4_f32m1(sum_sq, 1), gvl);
2597
+ mean_square_v = __riscv_vfadd_vv_f32m1(mean_square_v, __riscv_vget_v_f32m4_f32m1(sum_sq, 2), gvl);
2598
+ mean_square_v = __riscv_vfadd_vv_f32m1(mean_square_v, __riscv_vget_v_f32m4_f32m1(sum_sq, 3), gvl);
2599
+ mean_square_v = __riscv_vfredusum_vs_f32m1_f32m1(mean_square_v, zero_v, gvl);
2600
+
2601
+ float mean_square = __riscv_vfmv_f_s_f32m1_f32(mean_square_v);
2602
+ mean_square /= hidden_size;
2603
+ mean_square = sqrt(mean_square - mean * mean + epsilon);
2604
+
2605
+ mean_square = 1.0f / mean_square;
2606
+ length = hidden_size;
2607
+ p_temp_output = p_output;
2608
+
2609
+ while (length > 0) {
2610
+ gvl = __riscv_vsetvl_e32m4(length);
2611
+ vfloat32m4_t src_data = __riscv_vle32_v_f32m4(p_temp_output, gvl);
2612
+ src_data = __riscv_vfsub_vf_f32m4(src_data, mean, gvl);
2613
+ src_data = __riscv_vfmul_vf_f32m4(src_data, mean_square, gvl);
2614
+ __riscv_vse32_v_f32m4(p_output, src_data, gvl);
2615
+ p_temp_output += gvl;
2616
+ p_output += gvl;
2617
+ length -= gvl;
2618
+ }
2619
+ }
2620
+ }
2621
+
2622
+ template <ggml_op op_type, typename T> void forward_binary(ggml_compute_params * params, ggml_tensor * op) {
2623
+ const ggml_tensor * src0 = op->src[0];
2624
+ const ggml_tensor * src1 = op->src[1];
2625
+ ggml_tensor * dst = op;
2626
+ GGML_ASSERT(ggml_can_repeat(src1, src0) && ggml_are_same_shape(src0, dst));
2627
+
2628
+ auto src0_rows = ggml_nrows(src0);
2629
+ auto src1_rows = ggml_nrows(src1);
2630
+
2631
+ int ith = params->ith;
2632
+ int nth = params->nth;
2633
+
2634
+ GGML_TENSOR_BINARY_OP_LOCALS
2635
+
2636
+ GGML_ASSERT(nb0 == sizeof(T));
2637
+ GGML_ASSERT(nb00 == sizeof(T));
2638
+
2639
+ const auto [ir0, ir1] = get_thread_range(params, src0);
2640
+
2641
+ auto compute_func_vv = [&](int64_t blk_len, int64_t r, T * src0_ptr, T * src1_ptr, T * dst_ptr) {
2642
+ int64_t idx = 0;
2643
+ if constexpr (op_type == GGML_OP_ADD) {
2644
+ if constexpr (std::is_same_v<T, float>) {
2645
+ for (size_t vl; blk_len > 0; blk_len -= vl, idx += vl) {
2646
+ vl = __riscv_vsetvl_e32m4(blk_len);
2647
+ vfloat32m4_t lhs = __riscv_vle32_v_f32m4(src0_ptr + idx + r, vl);
2648
+ vfloat32m4_t rhs = __riscv_vle32_v_f32m4(src1_ptr + idx, vl);
2649
+ vfloat32m4_t res = __riscv_vfadd_vv_f32m4(lhs, rhs, vl);
2650
+ __riscv_vse32_v_f32m4(dst_ptr + idx + r, res, vl);
2651
+ }
2652
+ } else if constexpr (std::is_same_v<T, _Float16>) {
2653
+ for (size_t vl; blk_len > 0; blk_len -= vl, idx += vl) {
2654
+ vl = __riscv_vsetvl_e16m4(blk_len);
2655
+ vfloat16m4_t lhs = __riscv_vle16_v_f16m4((src0_ptr + idx + r), vl);
2656
+ vfloat16m4_t rhs = __riscv_vle16_v_f16m4((src1_ptr + idx), vl);
2657
+ vfloat16m4_t res = __riscv_vfadd_vv_f16m4(lhs, rhs, vl);
2658
+ __riscv_vse16_v_f16m4((dst_ptr + idx + r), res, vl);
2659
+ }
2660
+ } else {
2661
+ GGML_ABORT("fatal error");
2662
+ }
2663
+ } else if constexpr (op_type == GGML_OP_SUB) {
2664
+ if constexpr (std::is_same_v<T, float>) {
2665
+ for (size_t vl; blk_len > 0; blk_len -= vl, idx += vl) {
2666
+ vl = __riscv_vsetvl_e32m4(blk_len);
2667
+ vfloat32m4_t lhs = __riscv_vle32_v_f32m4(src0_ptr + idx + r, vl);
2668
+ vfloat32m4_t rhs = __riscv_vle32_v_f32m4(src1_ptr + idx, vl);
2669
+ vfloat32m4_t res = __riscv_vfsub_vv_f32m4(lhs, rhs, vl);
2670
+ __riscv_vse32_v_f32m4(dst_ptr + idx + r, res, vl);
2671
+ }
2672
+ } else if constexpr (std::is_same_v<T, _Float16>) {
2673
+ for (size_t vl; blk_len > 0; blk_len -= vl, idx += vl) {
2674
+ vl = __riscv_vsetvl_e16m4(blk_len);
2675
+ vfloat16m4_t lhs = __riscv_vle16_v_f16m4((src0_ptr + idx + r), vl);
2676
+ vfloat16m4_t rhs = __riscv_vle16_v_f16m4((src1_ptr + idx), vl);
2677
+ vfloat16m4_t res = __riscv_vfsub_vv_f16m4(lhs, rhs, vl);
2678
+ __riscv_vse16_v_f16m4((dst_ptr + idx + r), res, vl);
2679
+ }
2680
+ } else {
2681
+ GGML_ABORT("fatal error");
2682
+ }
2683
+ } else if constexpr (op_type == GGML_OP_MUL) {
2684
+ if constexpr (std::is_same_v<T, float>) {
2685
+ for (size_t vl; blk_len > 0; blk_len -= vl, idx += vl) {
2686
+ vl = __riscv_vsetvl_e32m4(blk_len);
2687
+ vfloat32m4_t lhs = __riscv_vle32_v_f32m4(src0_ptr + idx + r, vl);
2688
+ vfloat32m4_t rhs = __riscv_vle32_v_f32m4(src1_ptr + idx, vl);
2689
+ vfloat32m4_t res = __riscv_vfmul_vv_f32m4(lhs, rhs, vl);
2690
+ __riscv_vse32_v_f32m4(dst_ptr + idx + r, res, vl);
2691
+ }
2692
+ } else if constexpr (std::is_same_v<T, _Float16>) {
2693
+ for (size_t vl; blk_len > 0; blk_len -= vl, idx += vl) {
2694
+ vl = __riscv_vsetvl_e16m4(blk_len);
2695
+ vfloat16m4_t lhs = __riscv_vle16_v_f16m4((src0_ptr + idx + r), vl);
2696
+ vfloat16m4_t rhs = __riscv_vle16_v_f16m4((src1_ptr + idx), vl);
2697
+ vfloat16m4_t res = __riscv_vfmul_vv_f16m4(lhs, rhs, vl);
2698
+ __riscv_vse16_v_f16m4((dst_ptr + idx + r), res, vl);
2699
+ }
2700
+ } else {
2701
+ GGML_ABORT("fatal error");
2702
+ }
2703
+ } else if constexpr (op_type == GGML_OP_DIV) {
2704
+ if constexpr (std::is_same_v<T, float>) {
2705
+ for (size_t vl; blk_len > 0; blk_len -= vl, idx += vl) {
2706
+ vl = __riscv_vsetvl_e32m4(blk_len);
2707
+ vfloat32m4_t lhs = __riscv_vle32_v_f32m4(src0_ptr + idx + r, vl);
2708
+ vfloat32m4_t rhs = __riscv_vle32_v_f32m4(src1_ptr + idx, vl);
2709
+ vfloat32m4_t res = __riscv_vfdiv_vv_f32m4(lhs, rhs, vl);
2710
+ __riscv_vse32_v_f32m4(dst_ptr + idx + r, res, vl);
2711
+ }
2712
+ } else if constexpr (std::is_same_v<T, _Float16>) {
2713
+ for (size_t vl; blk_len > 0; blk_len -= vl, idx += vl) {
2714
+ vl = __riscv_vsetvl_e16m4(blk_len);
2715
+ vfloat16m4_t lhs = __riscv_vle16_v_f16m4((src0_ptr + idx + r), vl);
2716
+ vfloat16m4_t rhs = __riscv_vle16_v_f16m4((src1_ptr + idx), vl);
2717
+ vfloat16m4_t res = __riscv_vfdiv_vv_f16m4(lhs, rhs, vl);
2718
+ __riscv_vse16_v_f16m4((dst_ptr + idx + r), res, vl);
2719
+ }
2720
+ } else {
2721
+ GGML_ABORT("fatal error");
2722
+ }
2723
+ } else {
2724
+ GGML_ABORT("fatal error");
2725
+ }
2726
+ };
2727
+
2728
+ if (src0_rows == src1_rows && src0_rows == 1 && ne00 == ne10) {
2729
+ int64_t task_per_thread = (ne00 + nth - 1) / nth;
2730
+ int64_t task_begin = ith * task_per_thread;
2731
+ int64_t task_end = std::min((ith + 1) * task_per_thread, ne00);
2732
+
2733
+ T * dst_ptr = ((T *) dst->data) + task_begin;
2734
+ T * src0_ptr = ((T *) src0->data) + task_begin;
2735
+ T * src1_ptr = ((T *) src1->data) + task_begin;
2736
+
2737
+ compute_func_vv(task_end - task_begin, 0, src0_ptr, src1_ptr, dst_ptr);
2738
+ } else if (ne10 > 1) {
2739
+ for (int64_t ir = ir0; ir < ir1; ++ir) {
2740
+ const int64_t i03 = ir / (ne02 * ne01);
2741
+ const int64_t i02 = (ir - i03 * ne02 * ne01) / ne01;
2742
+ const int64_t i01 = (ir - i03 * ne02 * ne01 - i02 * ne01);
2743
+
2744
+ const int64_t i13 = i03 % ne13;
2745
+ const int64_t i12 = i02 % ne12;
2746
+ const int64_t i11 = i01 % ne11;
2747
+
2748
+ T * dst_ptr = (T *) ((char *) dst->data + i03 * nb3 + i02 * nb2 + i01 * nb1);
2749
+ T * src0_ptr = (T *) ((char *) src0->data + i03 * nb03 + i02 * nb02 + i01 * nb01);
2750
+ T * src1_ptr = (T *) ((char *) src1->data + i13 * nb13 + i12 * nb12 + i11 * nb11);
2751
+
2752
+ // src1 is broadcastable across src0 and dst in i1, i2, i3
2753
+ for (int64_t r = 0; r < ne00; r += ne10) {
2754
+ compute_func_vv(ne10, r, src0_ptr, src1_ptr, dst_ptr);
2755
+ }
2756
+ }
2757
+ } else {
2758
+ for (int64_t ir = ir0; ir < ir1; ++ir) {
2759
+ const int64_t i03 = ir / (ne02 * ne01);
2760
+ const int64_t i02 = (ir - i03 * ne02 * ne01) / ne01;
2761
+ const int64_t i01 = (ir - i03 * ne02 * ne01 - i02 * ne01);
2762
+
2763
+ const int64_t i13 = i03 % ne13;
2764
+ const int64_t i12 = i02 % ne12;
2765
+ const int64_t i11 = i01 % ne11;
2766
+
2767
+ T * dst_ptr = (T *) ((char *) dst->data + i03 * nb3 + i02 * nb2 + i01 * nb1);
2768
+ T * src0_ptr = (T *) ((char *) src0->data + i03 * nb03 + i02 * nb02 + i01 * nb01);
2769
+ T * src1_ptr = (T *) ((char *) src1->data + i13 * nb13 + i12 * nb12 + i11 * nb11);
2770
+
2771
+ T rhs_scalar = src1_ptr[0];
2772
+ int64_t blk_len = ne00;
2773
+ int64_t r = 0;
2774
+
2775
+ for (size_t vl; blk_len > 0; blk_len -= vl, r += vl) {
2776
+ if constexpr (op_type == GGML_OP_ADD) {
2777
+ if constexpr (std::is_same_v<T, float>) {
2778
+ vl = __riscv_vsetvl_e32m4(blk_len);
2779
+ vfloat32m4_t lhs = __riscv_vle32_v_f32m4(src0_ptr + r, vl);
2780
+ vfloat32m4_t res = __riscv_vfadd_vf_f32m4(lhs, rhs_scalar, vl);
2781
+ __riscv_vse32_v_f32m4(dst_ptr + r, res, vl);
2782
+ } else if constexpr (std::is_same_v<T, _Float16>) {
2783
+ vl = __riscv_vsetvl_e16m4(blk_len);
2784
+ vfloat16m4_t lhs = __riscv_vle16_v_f16m4((src0_ptr + r), vl);
2785
+ vfloat16m4_t res = __riscv_vfadd_vf_f16m4(lhs, rhs_scalar, vl);
2786
+ __riscv_vse16_v_f16m4((dst_ptr + r), res, vl);
2787
+ } else {
2788
+ GGML_ABORT("fatal error");
2789
+ }
2790
+ } else if constexpr (op_type == GGML_OP_SUB) {
2791
+ if constexpr (std::is_same_v<T, float>) {
2792
+ vl = __riscv_vsetvl_e32m4(blk_len);
2793
+ vfloat32m4_t lhs = __riscv_vle32_v_f32m4(src0_ptr + r, vl);
2794
+ vfloat32m4_t res = __riscv_vfsub_vf_f32m4(lhs, rhs_scalar, vl);
2795
+ __riscv_vse32_v_f32m4(dst_ptr + r, res, vl);
2796
+ } else if constexpr (std::is_same_v<T, _Float16>) {
2797
+ vl = __riscv_vsetvl_e16m4(blk_len);
2798
+ vfloat16m4_t lhs = __riscv_vle16_v_f16m4((src0_ptr + r), vl);
2799
+ vfloat16m4_t res = __riscv_vfsub_vf_f16m4(lhs, rhs_scalar, vl);
2800
+ __riscv_vse16_v_f16m4((dst_ptr + r), res, vl);
2801
+ } else {
2802
+ GGML_ABORT("fatal error");
2803
+ }
2804
+ } else if constexpr (op_type == GGML_OP_MUL) {
2805
+ if constexpr (std::is_same_v<T, float>) {
2806
+ vl = __riscv_vsetvl_e32m4(blk_len);
2807
+ vfloat32m4_t lhs = __riscv_vle32_v_f32m4(src0_ptr + r, vl);
2808
+ vfloat32m4_t res = __riscv_vfmul_vf_f32m4(lhs, rhs_scalar, vl);
2809
+ __riscv_vse32_v_f32m4(dst_ptr + r, res, vl);
2810
+ } else if constexpr (std::is_same_v<T, _Float16>) {
2811
+ vl = __riscv_vsetvl_e16m4(blk_len);
2812
+ vfloat16m4_t lhs = __riscv_vle16_v_f16m4((src0_ptr + r), vl);
2813
+ vfloat16m4_t res = __riscv_vfmul_vf_f16m4(lhs, rhs_scalar, vl);
2814
+ __riscv_vse16_v_f16m4((dst_ptr + r), res, vl);
2815
+ } else {
2816
+ GGML_ABORT("fatal error");
2817
+ }
2818
+ } else if constexpr (op_type == GGML_OP_DIV) {
2819
+ if constexpr (std::is_same_v<T, float>) {
2820
+ vl = __riscv_vsetvl_e32m4(blk_len);
2821
+ vfloat32m4_t lhs = __riscv_vle32_v_f32m4(src0_ptr + r, vl);
2822
+ vfloat32m4_t res = __riscv_vfdiv_vf_f32m4(lhs, rhs_scalar, vl);
2823
+ __riscv_vse32_v_f32m4(dst_ptr + r, res, vl);
2824
+ } else if constexpr (std::is_same_v<T, _Float16>) {
2825
+ vl = __riscv_vsetvl_e16m4(blk_len);
2826
+ vfloat16m4_t lhs = __riscv_vle16_v_f16m4((src0_ptr + r), vl);
2827
+ vfloat16m4_t res = __riscv_vfdiv_vf_f16m4(lhs, rhs_scalar, vl);
2828
+ __riscv_vse16_v_f16m4((dst_ptr + r), res, vl);
2829
+ } else {
2830
+ GGML_ABORT("fatal error");
2831
+ }
2832
+ } else {
2833
+ GGML_ABORT("fatal error");
2834
+ }
2835
+ }
2836
+ }
2837
+ }
2838
+ }
2839
+
2840
+ template <typename T> void forward_sum_rows(const ggml_compute_params * params, ggml_tensor * op) {
2841
+ const ggml_tensor * src0 = op->src[0];
2842
+ ggml_tensor * dst = op;
2843
+
2844
+ const int ith = params->ith;
2845
+ const int nth = params->nth;
2846
+
2847
+ GGML_TENSOR_UNARY_OP_LOCALS
2848
+
2849
+ GGML_ASSERT(ne0 == 1);
2850
+ GGML_ASSERT(ne1 == ne01);
2851
+ GGML_ASSERT(ne2 == ne02);
2852
+ GGML_ASSERT(ne3 == ne03);
2853
+
2854
+ int64_t n_task = ne01 * ne02 * ne03;
2855
+ int64_t task_per_thread = (n_task + nth - 1) / nth;
2856
+ int64_t ir_start = ith * task_per_thread;
2857
+ int64_t ir_end = std::min(ir_start + task_per_thread, n_task);
2858
+
2859
+ for (int64_t ir = ir_start; ir < ir_end; ir++) {
2860
+ const int64_t i3 = ir / (ne02 * ne01);
2861
+ const int64_t i2 = (ir - i3 * ne02 * ne01) / ne01;
2862
+ const int64_t i1 = (ir - i3 * ne02 * ne01 - i2 * ne01);
2863
+
2864
+ T * src_row = (T *) ((char *) src0->data + i1 * nb01 + i2 * nb02 + i3 * nb03);
2865
+ T * dst_row = (T *) ((char *) op->data + i1 * nb1 + i2 * nb2 + i3 * nb3);
2866
+
2867
+ float row_sum = 0;
2868
+
2869
+ if constexpr (std::is_same_v<T, float>) {
2870
+ size_t gvl = __riscv_vsetvlmax_e32m4();
2871
+ vfloat32m4_t acc_vec = __riscv_vfmv_v_f_f32m4(0.0f, gvl);
2872
+ int64_t length = ne00;
2873
+ const float * p_data = src_row;
2874
+
2875
+ while (length > 0) {
2876
+ size_t vl = __riscv_vsetvl_e32m4(length);
2877
+ vfloat32m4_t vec = __riscv_vle32_v_f32m4(p_data, vl);
2878
+ acc_vec = __riscv_vfadd_vv_f32m4(acc_vec, vec, vl);
2879
+ p_data += vl;
2880
+ length -= vl;
2881
+ }
2882
+
2883
+ gvl = __riscv_vsetvlmax_e32m1();
2884
+ vfloat32m1_t zero_v = __riscv_vfmv_v_f_f32m1(0.0f, gvl);
2885
+ vfloat32m1_t sum_v = __riscv_vfadd_vv_f32m1(__riscv_vget_v_f32m4_f32m1(acc_vec, 0),
2886
+ __riscv_vget_v_f32m4_f32m1(acc_vec, 1), gvl);
2887
+ sum_v = __riscv_vfadd_vv_f32m1(sum_v, __riscv_vget_v_f32m4_f32m1(acc_vec, 2), gvl);
2888
+ sum_v = __riscv_vfadd_vv_f32m1(sum_v, __riscv_vget_v_f32m4_f32m1(acc_vec, 3), gvl);
2889
+ sum_v = __riscv_vfredusum_vs_f32m1_f32m1(sum_v, zero_v, gvl);
2890
+ row_sum = __riscv_vfmv_f_s_f32m1_f32(sum_v);
2891
+ } else if constexpr (std::is_same_v<T, _Float16>) {
2892
+ size_t gvl = __riscv_vsetvlmax_e16m2();
2893
+ vfloat32m4_t acc_vec = __riscv_vfmv_v_f_f32m4(0.0f, gvl);
2894
+ int64_t length = ne00;
2895
+ const _Float16 * p_data = src_row;
2896
+
2897
+ while (length > 0) {
2898
+ size_t vl = __riscv_vsetvl_e16m2(length);
2899
+ vfloat16m2_t vec_f16 = __riscv_vle16_v_f16m2(p_data, vl);
2900
+ vfloat32m4_t vec_f32 = __riscv_vfwcvt_f_f_v_f32m4(vec_f16, vl);
2901
+ acc_vec = __riscv_vfadd_vv_f32m4(acc_vec, vec_f32, vl);
2902
+ p_data += vl;
2903
+ length -= vl;
2904
+ }
2905
+
2906
+ gvl = __riscv_vsetvlmax_e32m1();
2907
+ vfloat32m1_t zero_v = __riscv_vfmv_v_f_f32m1(0.0f, gvl);
2908
+ vfloat32m1_t sum_v = __riscv_vfadd_vv_f32m1(__riscv_vget_v_f32m4_f32m1(acc_vec, 0),
2909
+ __riscv_vget_v_f32m4_f32m1(acc_vec, 1), gvl);
2910
+ sum_v = __riscv_vfadd_vv_f32m1(sum_v, __riscv_vget_v_f32m4_f32m1(acc_vec, 2), gvl);
2911
+ sum_v = __riscv_vfadd_vv_f32m1(sum_v, __riscv_vget_v_f32m4_f32m1(acc_vec, 3), gvl);
2912
+ sum_v = __riscv_vfredusum_vs_f32m1_f32m1(sum_v, zero_v, gvl);
2913
+ row_sum = __riscv_vfmv_f_s_f32m1_f32(sum_v);
2914
+ } else {
2915
+ GGML_ABORT("fatal error");
2916
+ }
2917
+
2918
+ dst_row[0] = row_sum;
2919
+ }
2920
+ }
2921
+
2922
+ template <typename T> void forward_repeat_nrows(ggml_compute_params * params, ggml_tensor * op) {
2923
+ const ggml_tensor * src0 = op->src[0];
2924
+ ggml_tensor * dst = op;
2925
+
2926
+ const int ith = params->ith;
2927
+ const int nth = params->nth;
2928
+
2929
+ int64_t nrows = ggml_nrows(src0);
2930
+ int64_t nrows_per_thread = (nrows + nth - 1) / nth;
2931
+ int64_t ir_start = ith * nrows_per_thread;
2932
+ int64_t ir_end = std::min(ir_start + nrows_per_thread, nrows);
2933
+
2934
+ if (src0->ne[0] == 1) {
2935
+ for (int64_t ir = ir_start; ir < ir_end; ir++) {
2936
+ T * src_row = (T *) ((char *) src0->data + ir * src0->nb[1]);
2937
+ T * dst_row = (T *) ((char *) dst->data + ir * dst->nb[1]);
2938
+
2939
+ T src_scalar = src_row[0];
2940
+
2941
+ int64_t length = dst->ne[0];
2942
+ int64_t idx = 0;
2943
+ size_t vl = 0;
2944
+
2945
+ while (length > 0) {
2946
+ if constexpr (std::is_same_v<T, int32_t>) {
2947
+ vl = __riscv_vsetvl_e32m4(length);
2948
+ vint32m4_t vec = __riscv_vmv_v_x_i32m4(src_scalar, vl);
2949
+ __riscv_vse32_v_i32m4(dst_row + idx, vec, vl);
2950
+ } else if constexpr (std::is_same_v<T, int16_t>) {
2951
+ vl = __riscv_vsetvl_e16m4(length);
2952
+ vint16m4_t vec = __riscv_vmv_v_x_i16m4(src_scalar, vl);
2953
+ __riscv_vse16_v_i16m4((dst_row + idx), vec, vl);
2954
+ } else {
2955
+ GGML_ABORT("fatal error");
2956
+ }
2957
+ idx += vl;
2958
+ length -= vl;
2959
+ }
2960
+ }
2961
+ } else if (src0->ne[0] == dst->ne[0]) {
2962
+ for (int64_t ir = ir_start; ir < ir_end; ir++) {
2963
+ T * src_row = (T *) ((char *) src0->data + ir * src0->nb[1]);
2964
+ T * dst_row = (T *) ((char *) dst->data + ir * dst->nb[1]);
2965
+
2966
+ int64_t length = dst->ne[0];
2967
+ int64_t idx = 0;
2968
+ size_t vl = 0;
2969
+
2970
+ while (length > 0) {
2971
+ if constexpr (std::is_same_v<T, int32_t>) {
2972
+ vl = __riscv_vsetvl_e32m4(length);
2973
+ vint32m4_t vec = __riscv_vle32_v_i32m4(src_row + idx, vl);
2974
+ __riscv_vse32_v_i32m4(dst_row + idx, vec, vl);
2975
+ } else if constexpr (std::is_same_v<T, int16_t>) {
2976
+ vl = __riscv_vsetvl_e16m4(length);
2977
+ vint16m4_t vec = __riscv_vle16_v_i16m4((src_row + idx), vl);
2978
+ __riscv_vse16_v_i16m4((dst_row + idx), vec, vl);
2979
+ } else {
2980
+ GGML_ABORT("fatal error");
2981
+ }
2982
+ idx += vl;
2983
+ length -= vl;
2984
+ }
2985
+ }
2986
+ } else {
2987
+ GGML_ABORT("fatal error");
2988
+ }
2989
+ }
2990
+
2991
+ template <typename T> void forward_repeat_dim1(ggml_compute_params * params, ggml_tensor * op) {
2992
+ const ggml_tensor * src0 = op->src[0];
2993
+ ggml_tensor * dst = op;
2994
+
2995
+ const int ith = params->ith;
2996
+ const int nth = params->nth;
2997
+
2998
+ const int64_t ne0 = dst->ne[0];
2999
+ const int64_t ne1 = dst->ne[1];
3000
+ const int64_t ne2 = dst->ne[2];
3001
+ const int64_t ne3 = dst->ne[3];
3002
+
3003
+ const int64_t total_batches = ne2 * ne3;
3004
+ const int64_t batches_per_thread = (total_batches + nth - 1) / nth;
3005
+ const int64_t batch_start = ith * batches_per_thread;
3006
+ const int64_t batch_end = std::min(batch_start + batches_per_thread, total_batches);
3007
+
3008
+ for (int64_t b = batch_start; b < batch_end; b++) {
3009
+ const int64_t i3 = b / ne2;
3010
+ const int64_t i2 = b % ne2;
3011
+
3012
+ T * src_base = (T *) ((char *) src0->data + i2 * src0->nb[2] + i3 * src0->nb[3]);
3013
+ T * dst_batch = (T *) ((char *) dst->data + i2 * dst->nb[2] + i3 * dst->nb[3]);
3014
+
3015
+ for (int64_t i1 = 0; i1 < ne1; i1++) {
3016
+ T * dst_ptr = (T *) ((char *) dst_batch + i1 * dst->nb[1]);
3017
+ int64_t length = ne0;
3018
+ int64_t idx = 0;
3019
+
3020
+ while (length > 0) {
3021
+ if constexpr (std::is_same_v<T, int32_t>) {
3022
+ size_t vl = __riscv_vsetvl_e32m4(length);
3023
+ vint32m4_t vec = __riscv_vle32_v_i32m4(src_base + idx, vl);
3024
+ __riscv_vse32_v_i32m4(dst_ptr + idx, vec, vl);
3025
+ idx += vl;
3026
+ length -= vl;
3027
+ } else if constexpr (std::is_same_v<T, int16_t>) {
3028
+ size_t vl = __riscv_vsetvl_e16m4(length);
3029
+ vint16m4_t vec = __riscv_vle16_v_i16m4((src_base + idx), vl);
3030
+ __riscv_vse16_v_i16m4((dst_ptr + idx), vec, vl);
3031
+ idx += vl;
3032
+ length -= vl;
3033
+ } else {
3034
+ GGML_ABORT("fatal error");
3035
+ }
3036
+ }
3037
+ }
3038
+ }
3039
+ }
3040
+
3041
+ template <typename T> void forward_get_rows(ggml_compute_params * params, ggml_tensor * op) {
3042
+ const ggml_tensor * src0 = op->src[0];
3043
+ const ggml_tensor * src1 = op->src[1];
3044
+ ggml_tensor * dst = op;
3045
+
3046
+ GGML_TENSOR_BINARY_OP_LOCALS
3047
+
3048
+ const int64_t nc = ne00;
3049
+ const int64_t nr = ggml_nelements(src1);
3050
+
3051
+ assert(ne0 == nc);
3052
+ assert(ne02 == ne11);
3053
+ assert(nb00 == sizeof(float));
3054
+ assert(ggml_nrows(op) == nr);
3055
+
3056
+ const int ith = params->ith;
3057
+ const int nth = params->nth;
3058
+
3059
+ int rows_nth = nth;
3060
+ int cols_nth = 1;
3061
+
3062
+ if (nr == 1) {
3063
+ rows_nth = 1;
3064
+ cols_nth = nth;
3065
+ }
3066
+
3067
+ // rows per thread
3068
+ const int dr = (nr + rows_nth - 1) / rows_nth;
3069
+ const int dc = (nc + cols_nth - 1) / cols_nth;
3070
+
3071
+ int rows_ith = ith % rows_nth;
3072
+ int cols_ith = ith % cols_nth;
3073
+
3074
+ // row range for this thread
3075
+ const int ir0 = dr * rows_ith;
3076
+ const int ir1 = MIN(ir0 + dr, nr);
3077
+
3078
+ const int cr0 = dc * cols_ith;
3079
+ const int cr1 = MIN(cr0 + dc, nc);
3080
+
3081
+ for (int64_t i = ir0; i < ir1; ++i) {
3082
+ const int64_t i12 = i / (ne11 * ne10);
3083
+ const int64_t i11 = (i - i12 * ne11 * ne10) / ne10;
3084
+ const int64_t i10 = (i - i12 * ne11 * ne10 - i11 * ne10);
3085
+ const int64_t i01 = *(int32_t *) ((char *) src1->data + i10 * nb10 + i11 * nb11 + i12 * nb12);
3086
+
3087
+ GGML_ASSERT(i01 >= 0 && i01 < ne01);
3088
+
3089
+ memcpy1d(((char *) dst->data + i10 * nb1 + i11 * nb2 + i12 * nb3) + cr0 * sizeof(T),
3090
+ ((char *) src0->data + i01 * nb01 + i11 * nb02 + i12 * nb03) + cr0 * sizeof(T),
3091
+ (cr1 - cr0) * sizeof(T));
3092
+ }
3093
+ }
3094
+
3095
+ template <typename T> void forward_concat(ggml_compute_params * params, ggml_tensor * op) {
3096
+ const ggml_tensor * src0 = op->src[0];
3097
+ const ggml_tensor * src1 = op->src[1];
3098
+ ggml_tensor * dst = op;
3099
+
3100
+ GGML_ASSERT(ggml_type_size(src0->type) == sizeof(float));
3101
+
3102
+ GGML_TENSOR_BINARY_OP_LOCALS
3103
+
3104
+ const int32_t dim = ggml_get_op_params_i32(dst, 0);
3105
+
3106
+ GGML_ASSERT(dim == 0 && nb0 == sizeof(float) && nb1 == sizeof(float) * (ne00 + ne10));
3107
+
3108
+ const int64_t nr = ggml_nrows(dst);
3109
+ const int64_t nc = ne0;
3110
+
3111
+ const int ith = params->ith;
3112
+ const int nth = params->nth;
3113
+
3114
+ int rows_nth = nth;
3115
+ int cols_nth = 1;
3116
+
3117
+ if (nr == 1) {
3118
+ rows_nth = 1;
3119
+ cols_nth = nth;
3120
+ }
3121
+
3122
+ const int dr = (nr + rows_nth - 1) / rows_nth;
3123
+ const int dc = (nc + cols_nth - 1) / cols_nth;
3124
+
3125
+ int rows_ith = ith % rows_nth;
3126
+ int cols_ith = ith % cols_nth;
3127
+
3128
+ // row range for this thread
3129
+ const int ir0 = dr * rows_ith;
3130
+ const int ir1 = MIN(ir0 + dr, nr);
3131
+
3132
+ const int cr0 = dc * cols_ith;
3133
+ const int cr1 = MIN(cr0 + dc, nc);
3134
+
3135
+ int64_t o[4] = { 0, 0, 0, 0 };
3136
+ o[dim] = src0->ne[dim];
3137
+ const float * x;
3138
+
3139
+ for (int64_t i = ir0; i < ir1; ++i) {
3140
+ const int64_t i3 = i / (ne02 * ne01);
3141
+ const int64_t i2 = (i - i3 * ne02 * ne01) / ne01;
3142
+ const int64_t i1 = (i - i3 * ne02 * ne01 - i2 * ne01);
3143
+
3144
+ for (int i0 = cr0; i0 < cr1; i0++) {
3145
+ if (i0 < ne00 && i1 < ne01 && i2 < ne02 && i3 < ne03) {
3146
+ x = (const float *) ((const char *) src0->data + (i0) *nb00 + (i1) *nb01 + (i2) *nb02 + (i3) *nb03);
3147
+ } else {
3148
+ x = (const float *) ((const char *) src1->data + (i0 - o[0]) * nb10 + (i1 - o[1]) * nb11 +
3149
+ (i2 - o[2]) * nb12 + (i3 - o[3]) * nb13);
3150
+ }
3151
+
3152
+ float * y = (float *) ((char *) dst->data + i0 * nb0 + i1 * nb1 + i2 * nb2 + i3 * nb3);
3153
+
3154
+ *y = *x;
3155
+ }
3156
+ }
3157
+ }
3158
+
3159
+ template void forward_binary<GGML_OP_ADD, float>(ggml_compute_params * params, ggml_tensor * op);
3160
+ template void forward_binary<GGML_OP_SUB, float>(ggml_compute_params * params, ggml_tensor * op);
3161
+ template void forward_binary<GGML_OP_MUL, float>(ggml_compute_params * params, ggml_tensor * op);
3162
+ template void forward_binary<GGML_OP_DIV, float>(ggml_compute_params * params, ggml_tensor * op);
3163
+ template void forward_binary<GGML_OP_ADD, _Float16>(ggml_compute_params * params, ggml_tensor * op);
3164
+ template void forward_binary<GGML_OP_SUB, _Float16>(ggml_compute_params * params, ggml_tensor * op);
3165
+ template void forward_binary<GGML_OP_MUL, _Float16>(ggml_compute_params * params, ggml_tensor * op);
3166
+ template void forward_binary<GGML_OP_DIV, _Float16>(ggml_compute_params * params, ggml_tensor * op);
3167
+ template void forward_sum_rows<float>(const ggml_compute_params * params, ggml_tensor * op);
3168
+ template void forward_sum_rows<_Float16>(const ggml_compute_params * params, ggml_tensor * op);
3169
+ template void forward_repeat_nrows<int32_t>(ggml_compute_params * params, ggml_tensor * op);
3170
+ template void forward_repeat_nrows<int16_t>(ggml_compute_params * params, ggml_tensor * op);
3171
+ template void forward_repeat_dim1<int32_t>(ggml_compute_params * params, ggml_tensor * op);
3172
+ template void forward_repeat_dim1<int16_t>(ggml_compute_params * params, ggml_tensor * op);
3173
+ template void forward_get_rows<int32_t>(ggml_compute_params * params, ggml_tensor * op);
3174
+ template void forward_get_rows<int16_t>(ggml_compute_params * params, ggml_tensor * op);
3175
+ template void forward_concat<int32_t>(ggml_compute_params * params, ggml_tensor * op);
3176
+ template void forward_concat<int16_t>(ggml_compute_params * params, ggml_tensor * op);
3177
+
3178
+ } // namespace spacemit_kernels::rvv