whispercpp 1.3.7 → 1.3.8

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (308) hide show
  1. checksums.yaml +4 -4
  2. data/README.md +5 -4
  3. data/ext/options.rb +1 -1
  4. data/ext/ruby_whisper.c +0 -1
  5. data/ext/ruby_whisper.h +7 -1
  6. data/ext/ruby_whisper_context.c +50 -1
  7. data/ext/ruby_whisper_log_settable.h +1 -2
  8. data/ext/ruby_whisper_params.c +9 -8
  9. data/ext/ruby_whisper_transcribe.cpp +0 -19
  10. data/ext/ruby_whisper_vad_context.c +30 -10
  11. data/ext/ruby_whisper_vad_context_detect.cpp +8 -9
  12. data/ext/ruby_whisper_vad_params.c +4 -4
  13. data/ext/ruby_whisper_vad_segment.c +2 -2
  14. data/ext/sources/CMakeLists.txt +2 -1
  15. data/ext/sources/cmake/parakeet.pc.in +2 -2
  16. data/ext/sources/cmake/whisper.pc.in +2 -2
  17. data/ext/sources/examples/cli/cli.cpp +9 -1
  18. data/ext/sources/examples/common-ggml.cpp +2 -0
  19. data/ext/sources/examples/vad-speech-segments/speech.cpp +3 -2
  20. data/ext/sources/ggml/CMakeLists.txt +3 -4
  21. data/ext/sources/ggml/include/ggml-cuda.h +0 -3
  22. data/ext/sources/ggml/include/ggml-sycl.h +8 -0
  23. data/ext/sources/ggml/include/ggml.h +3 -1
  24. data/ext/sources/ggml/src/CMakeLists.txt +8 -1
  25. data/ext/sources/ggml/src/ggml-backend-meta.cpp +7 -4
  26. data/ext/sources/ggml/src/ggml-common.h +13 -2
  27. data/ext/sources/ggml/src/ggml-cpu/CMakeLists.txt +1 -1
  28. data/ext/sources/ggml/src/ggml-cpu/amx/mmq.cpp +5 -6
  29. data/ext/sources/ggml/src/ggml-cpu/arch/arm/quants.c +78 -4
  30. data/ext/sources/ggml/src/ggml-cpu/arch/x86/quants.c +142 -4
  31. data/ext/sources/ggml/src/ggml-cpu/arch-fallback.h +7 -2
  32. data/ext/sources/ggml/src/ggml-cpu/ggml-cpu.c +14 -0
  33. data/ext/sources/ggml/src/ggml-cpu/llamafile/sgemm.cpp +26 -19
  34. data/ext/sources/ggml/src/ggml-cpu/ops.cpp +129 -46
  35. data/ext/sources/ggml/src/ggml-cpu/quants.c +51 -0
  36. data/ext/sources/ggml/src/ggml-cpu/quants.h +3 -0
  37. data/ext/sources/ggml/src/ggml-cpu/simd-gemm.h +1 -1
  38. data/ext/sources/ggml/src/ggml-cpu/simd-mappings.h +11 -0
  39. data/ext/sources/ggml/src/ggml-cpu/vec.cpp +2 -2
  40. data/ext/sources/ggml/src/ggml-cuda/binbcast.cu +90 -46
  41. data/ext/sources/ggml/src/ggml-cuda/col2im-1d.cu +81 -0
  42. data/ext/sources/ggml/src/ggml-cuda/col2im-1d.cuh +3 -0
  43. data/ext/sources/ggml/src/ggml-cuda/common.cuh +4 -0
  44. data/ext/sources/ggml/src/ggml-cuda/concat.cu +33 -21
  45. data/ext/sources/ggml/src/ggml-cuda/conv-transpose-1d.cu +14 -12
  46. data/ext/sources/ggml/src/ggml-cuda/convert.cu +86 -34
  47. data/ext/sources/ggml/src/ggml-cuda/cpy.cu +80 -29
  48. data/ext/sources/ggml/src/ggml-cuda/fattn-common.cuh +9 -5
  49. data/ext/sources/ggml/src/ggml-cuda/fattn-mma-f16.cuh +4 -0
  50. data/ext/sources/ggml/src/ggml-cuda/fattn-tile.cuh +9 -5
  51. data/ext/sources/ggml/src/ggml-cuda/fattn.cu +27 -21
  52. data/ext/sources/ggml/src/ggml-cuda/gated_delta_net.cu +40 -25
  53. data/ext/sources/ggml/src/ggml-cuda/gated_delta_net.cuh +10 -0
  54. data/ext/sources/ggml/src/ggml-cuda/getrows.cu +15 -12
  55. data/ext/sources/ggml/src/ggml-cuda/ggml-cuda.cu +718 -1248
  56. data/ext/sources/ggml/src/ggml-cuda/mmq.cu +7 -0
  57. data/ext/sources/ggml/src/ggml-cuda/mmvq.cu +77 -40
  58. data/ext/sources/ggml/src/ggml-cuda/out-prod.cu +55 -12
  59. data/ext/sources/ggml/src/ggml-cuda/set-rows.cu +64 -4
  60. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_16-ncols2_2.cu +1 -0
  61. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_32-ncols2_2.cu +1 -0
  62. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_2.cu +1 -0
  63. data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_8-ncols2_2.cu +1 -0
  64. data/ext/sources/ggml/src/ggml-cuda/topk-moe.cu +7 -1
  65. data/ext/sources/ggml/src/ggml-cuda/vendors/hip.h +1 -0
  66. data/ext/sources/ggml/src/ggml-cuda/vendors/musa.h +1 -0
  67. data/ext/sources/ggml/src/ggml-hexagon/CMakeLists.txt +0 -5
  68. data/ext/sources/ggml/src/ggml-hexagon/ggml-hexagon.cpp +1634 -1293
  69. data/ext/sources/ggml/src/ggml-hexagon/htp/CMakeLists.txt +11 -40
  70. data/ext/sources/ggml/src/ggml-hexagon/htp/cmake-toolchain.cmake +13 -15
  71. data/ext/sources/ggml/src/ggml-hexagon/htp/concat-ops.c +1 -1
  72. data/ext/sources/ggml/src/ggml-hexagon/htp/flash-attn-ops.c +1749 -399
  73. data/ext/sources/ggml/src/ggml-hexagon/htp/flash-attn-ops.h +303 -0
  74. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-common.h +80 -0
  75. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-dma.h +26 -23
  76. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-profile.h +64 -0
  77. data/ext/sources/ggml/src/ggml-hexagon/htp/hex-utils.h +1 -83
  78. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-fa-kernels.h +555 -0
  79. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-mm-kernels-tiled.h +1303 -0
  80. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-queue.c +9 -0
  81. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-queue.h +27 -4
  82. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-utils.h +59 -37
  83. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-ctx.h +11 -3
  84. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-ops.h +52 -12
  85. data/ext/sources/ggml/src/ggml-hexagon/htp/htp-vtcm.h +19 -0
  86. data/ext/sources/ggml/src/ggml-hexagon/htp/htp_iface.idl +2 -1
  87. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-base.h +14 -30
  88. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-exp.h +39 -0
  89. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-fa-kernels.h +232 -0
  90. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-flat.h +1511 -0
  91. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-tiled.h +1200 -0
  92. data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-sigmoid.h +39 -0
  93. data/ext/sources/ggml/src/ggml-hexagon/htp/main.c +127 -32
  94. data/ext/sources/ggml/src/ggml-hexagon/htp/matmul-ops.c +3023 -4425
  95. data/ext/sources/ggml/src/ggml-hexagon/htp/matmul-ops.h +650 -0
  96. data/ext/sources/ggml/src/ggml-hexagon/htp/rope-ops.c +48 -13
  97. data/ext/sources/ggml/src/ggml-hexagon/htp/ssm-conv.c +10 -9
  98. data/ext/sources/ggml/src/ggml-hexagon/htp/worker-pool.c +15 -3
  99. data/ext/sources/ggml/src/ggml-hexagon/htp/worker-pool.h +8 -0
  100. data/ext/sources/ggml/src/ggml-hexagon/htp-opnode.h +168 -50
  101. data/ext/sources/ggml/src/ggml-hexagon/libggml-htp.inf +0 -4
  102. data/ext/sources/ggml/src/ggml-hip/CMakeLists.txt +5 -0
  103. data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.cpp +69 -5
  104. data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.h +4 -1
  105. data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.m +27 -6
  106. data/ext/sources/ggml/src/ggml-metal/ggml-metal-impl.h +38 -0
  107. data/ext/sources/ggml/src/ggml-metal/ggml-metal-ops.cpp +132 -2
  108. data/ext/sources/ggml/src/ggml-metal/ggml-metal-ops.h +2 -0
  109. data/ext/sources/ggml/src/ggml-metal/ggml-metal.metal +345 -87
  110. data/ext/sources/ggml/src/ggml-opencl/CMakeLists.txt +13 -0
  111. data/ext/sources/ggml/src/ggml-opencl/fa_tune.h +92 -0
  112. data/ext/sources/ggml/src/ggml-opencl/ggml-opencl.cpp +4060 -357
  113. data/ext/sources/ggml/src/ggml-opencl/kernels/cvt.cl +198 -0
  114. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f16.cl +81 -41
  115. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32.cl +88 -39
  116. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_f16.cl +1995 -96
  117. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_q4_0.cl +1615 -0
  118. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_q8_0.cl +1486 -0
  119. data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_pre_f16.cl +156 -0
  120. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_mxfp4_f32_ns.cl +74 -6
  121. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_0_f32_ns.cl +74 -6
  122. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_1_f32_ns.cl +74 -6
  123. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_k_f32_ns.cl +71 -6
  124. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_0_f32_ns.cl +74 -6
  125. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_1_f32_ns.cl +74 -6
  126. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_k_f32_ns.cl +74 -6
  127. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q6_k_f32_ns.cl +74 -6
  128. data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q1_0_f32.cl +94 -0
  129. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q1_0_f32.cl +121 -0
  130. data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q8_0_f32.cl +1 -1
  131. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q1_0_f32_l4_lm.cl +156 -0
  132. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_f16_f32_l4.cl +1149 -0
  133. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q1_0_f32.cl +141 -0
  134. data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q1_0_f32_flat.cl +190 -0
  135. data/ext/sources/ggml/src/ggml-opencl/kernels/norm.cl +5 -2
  136. data/ext/sources/ggml/src/ggml-opencl/kernels/set_rows.cl +500 -0
  137. data/ext/sources/ggml/src/ggml-opencl/libdl.h +79 -0
  138. data/ext/sources/ggml/src/ggml-openvino/.clang-format +0 -5
  139. data/ext/sources/ggml/src/ggml-openvino/CMakeLists.txt +2 -4
  140. data/ext/sources/ggml/src/ggml-openvino/ggml-decoder.cpp +733 -130
  141. data/ext/sources/ggml/src/ggml-openvino/ggml-decoder.h +76 -23
  142. data/ext/sources/ggml/src/ggml-openvino/ggml-openvino-extra.cpp +57 -3
  143. data/ext/sources/ggml/src/ggml-openvino/ggml-openvino-extra.h +29 -8
  144. data/ext/sources/ggml/src/ggml-openvino/ggml-openvino.cpp +307 -59
  145. data/ext/sources/ggml/src/ggml-openvino/ggml-quants.cpp +66 -0
  146. data/ext/sources/ggml/src/ggml-openvino/ggml-quants.h +10 -4
  147. data/ext/sources/ggml/src/ggml-openvino/openvino/decoder.h +56 -16
  148. data/ext/sources/ggml/src/ggml-openvino/openvino/frontend.h +1 -1
  149. data/ext/sources/ggml/src/ggml-openvino/openvino/input_model.h +4 -4
  150. data/ext/sources/ggml/src/ggml-openvino/openvino/node_context.h +94 -37
  151. data/ext/sources/ggml/src/ggml-openvino/openvino/op/add_id.cpp +76 -0
  152. data/ext/sources/ggml/src/ggml-openvino/openvino/op/argsort.cpp +47 -0
  153. data/ext/sources/ggml/src/ggml-openvino/openvino/op/clamp.cpp +33 -0
  154. data/ext/sources/ggml/src/ggml-openvino/openvino/op/concat.cpp +48 -0
  155. data/ext/sources/ggml/src/ggml-openvino/openvino/op/cont.cpp +8 -16
  156. data/ext/sources/ggml/src/ggml-openvino/openvino/op/cpy.cpp +14 -1
  157. data/ext/sources/ggml/src/ggml-openvino/openvino/op/div.cpp +146 -0
  158. data/ext/sources/ggml/src/ggml-openvino/openvino/op/flash_attn_ext.cpp +108 -21
  159. data/ext/sources/ggml/src/ggml-openvino/openvino/op/gated_delta_net.cpp +282 -0
  160. data/ext/sources/ggml/src/ggml-openvino/openvino/op/gated_delta_net.hpp +65 -0
  161. data/ext/sources/ggml/src/ggml-openvino/openvino/op/get_rows.cpp +2 -9
  162. data/ext/sources/ggml/src/ggml-openvino/openvino/op/glu_geglu.cpp +21 -7
  163. data/ext/sources/ggml/src/ggml-openvino/openvino/op/glu_swiglu.cpp +41 -8
  164. data/ext/sources/ggml/src/ggml-openvino/openvino/op/im2col.cpp +120 -0
  165. data/ext/sources/ggml/src/ggml-openvino/openvino/op/l2_norm.cpp +44 -0
  166. data/ext/sources/ggml/src/ggml-openvino/openvino/op/mul_mat_id.cpp +226 -0
  167. data/ext/sources/ggml/src/ggml-openvino/openvino/op/mulmat.cpp +19 -9
  168. data/ext/sources/ggml/src/ggml-openvino/openvino/op/norm.cpp +58 -0
  169. data/ext/sources/ggml/src/ggml-openvino/openvino/op/pad.cpp +95 -0
  170. data/ext/sources/ggml/src/ggml-openvino/openvino/op/permute.cpp +58 -13
  171. data/ext/sources/ggml/src/ggml-openvino/openvino/op/repeat.cpp +74 -0
  172. data/ext/sources/ggml/src/ggml-openvino/openvino/op/reshape.cpp +13 -6
  173. data/ext/sources/ggml/src/ggml-openvino/openvino/op/rms_norm.cpp +1 -1
  174. data/ext/sources/ggml/src/ggml-openvino/openvino/op/rope.cpp +134 -38
  175. data/ext/sources/ggml/src/ggml-openvino/openvino/op/set_rows.cpp +3 -3
  176. data/ext/sources/ggml/src/ggml-openvino/openvino/op/softmax.cpp +126 -49
  177. data/ext/sources/ggml/src/ggml-openvino/openvino/op/ssm_conv.cpp +59 -0
  178. data/ext/sources/ggml/src/ggml-openvino/openvino/op/sum_rows.cpp +27 -0
  179. data/ext/sources/ggml/src/ggml-openvino/openvino/op/transpose.cpp +32 -1
  180. data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_silu.cpp +1 -1
  181. data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_softplus.cpp +38 -0
  182. data/ext/sources/ggml/src/ggml-openvino/openvino/op/view.cpp +90 -25
  183. data/ext/sources/ggml/src/ggml-openvino/openvino/op_table.cpp +41 -23
  184. data/ext/sources/ggml/src/ggml-openvino/openvino/op_table.h +18 -5
  185. data/ext/sources/ggml/src/ggml-openvino/openvino/pass/mark_decompression_convert_constant_folding.h +1 -1
  186. data/ext/sources/ggml/src/ggml-openvino/openvino/translate_session.cpp +43 -40
  187. data/ext/sources/ggml/src/ggml-openvino/openvino/translate_session.h +5 -4
  188. data/ext/sources/ggml/src/ggml-openvino/openvino/utils.cpp +548 -3
  189. data/ext/sources/ggml/src/ggml-openvino/openvino/utils.h +28 -26
  190. data/ext/sources/ggml/src/ggml-openvino/utils.cpp +383 -94
  191. data/ext/sources/ggml/src/ggml-openvino/utils.h +11 -8
  192. data/ext/sources/ggml/src/ggml-quants.c +76 -0
  193. data/ext/sources/ggml/src/ggml-quants.h +3 -0
  194. data/ext/sources/ggml/src/ggml-sycl/CMakeLists.txt +5 -5
  195. data/ext/sources/ggml/src/ggml-sycl/backend.hpp +2 -0
  196. data/ext/sources/ggml/src/ggml-sycl/binbcast.cpp +12 -0
  197. data/ext/sources/ggml/src/ggml-sycl/col2im-1d.cpp +102 -0
  198. data/ext/sources/ggml/src/ggml-sycl/col2im-1d.hpp +8 -0
  199. data/ext/sources/ggml/src/ggml-sycl/common.cpp +6 -8
  200. data/ext/sources/ggml/src/ggml-sycl/common.hpp +19 -2
  201. data/ext/sources/ggml/src/ggml-sycl/concat.cpp +21 -1
  202. data/ext/sources/ggml/src/ggml-sycl/conv2d-dw.cpp +158 -0
  203. data/ext/sources/ggml/src/ggml-sycl/conv2d-dw.hpp +10 -0
  204. data/ext/sources/ggml/src/ggml-sycl/conv2d-transpose.cpp +125 -0
  205. data/ext/sources/ggml/src/ggml-sycl/conv2d-transpose.hpp +10 -0
  206. data/ext/sources/ggml/src/ggml-sycl/conv2d.cpp +150 -0
  207. data/ext/sources/ggml/src/ggml-sycl/conv2d.hpp +10 -0
  208. data/ext/sources/ggml/src/ggml-sycl/conv3d.cpp +224 -0
  209. data/ext/sources/ggml/src/ggml-sycl/conv3d.hpp +8 -0
  210. data/ext/sources/ggml/src/ggml-sycl/convert.cpp +6 -0
  211. data/ext/sources/ggml/src/ggml-sycl/cpy.cpp +706 -0
  212. data/ext/sources/ggml/src/ggml-sycl/cpy.hpp +281 -0
  213. data/ext/sources/ggml/src/ggml-sycl/cross_entropy_loss.cpp +255 -0
  214. data/ext/sources/ggml/src/ggml-sycl/cross_entropy_loss.hpp +7 -0
  215. data/ext/sources/ggml/src/ggml-sycl/dequantize.hpp +15 -0
  216. data/ext/sources/ggml/src/ggml-sycl/dmmv.cpp +492 -319
  217. data/ext/sources/ggml/src/ggml-sycl/dpct/helper.hpp +15 -7
  218. data/ext/sources/ggml/src/ggml-sycl/element_wise.cpp +215 -115
  219. data/ext/sources/ggml/src/ggml-sycl/element_wise.hpp +2 -0
  220. data/ext/sources/ggml/src/ggml-sycl/ggml-sycl.cpp +1006 -336
  221. data/ext/sources/ggml/src/ggml-sycl/mmvq.cpp +252 -67
  222. data/ext/sources/ggml/src/ggml-sycl/mmvq.hpp +17 -0
  223. data/ext/sources/ggml/src/ggml-sycl/norm.cpp +103 -49
  224. data/ext/sources/ggml/src/ggml-sycl/outprod.cpp +45 -9
  225. data/ext/sources/ggml/src/ggml-sycl/pool.cpp +185 -0
  226. data/ext/sources/ggml/src/ggml-sycl/pool.hpp +22 -0
  227. data/ext/sources/ggml/src/ggml-sycl/presets.hpp +3 -1
  228. data/ext/sources/ggml/src/ggml-sycl/set_rows.cpp +10 -2
  229. data/ext/sources/ggml/src/ggml-sycl/softmax.cpp +9 -10
  230. data/ext/sources/ggml/src/ggml-sycl/vecdotq.hpp +35 -0
  231. data/ext/sources/ggml/src/ggml-vulkan/CMakeLists.txt +5 -0
  232. data/ext/sources/ggml/src/ggml-vulkan/ggml-vulkan.cpp +833 -215
  233. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/col2im_1d.comp +61 -0
  234. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/conv2d_mm.comp +1 -1
  235. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/conv3d_mm.comp +431 -0
  236. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/diag.comp +3 -3
  237. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp +1 -0
  238. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp +1 -0
  239. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/generic_unary_head.glsl +21 -19
  240. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/get_rows_back.comp +25 -0
  241. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/glu_head.glsl +23 -4
  242. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/glu_main.glsl +14 -18
  243. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/l2_norm.comp +4 -7
  244. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq.comp +21 -24
  245. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp +31 -23
  246. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp +6 -5
  247. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl +84 -67
  248. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/norm.comp +10 -10
  249. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/repeat_back.comp +3 -3
  250. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/roll.comp +3 -3
  251. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/tri.comp +3 -3
  252. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/unary.comp +168 -0
  253. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +121 -74
  254. data/ext/sources/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp +26 -19
  255. data/ext/sources/ggml/src/ggml-webgpu/ggml-webgpu.cpp +31 -36
  256. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/common_decls.tmpl +16 -2
  257. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_split.wgsl +7 -7
  258. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/get_rows.wgsl +21 -0
  259. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_decls.tmpl +439 -320
  260. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_id_vec.wgsl +2 -2
  261. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec.wgsl +45 -39
  262. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_acc.tmpl +586 -465
  263. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_q_acc.tmpl +63 -69
  264. data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/quantize_q8.wgsl +14 -9
  265. data/ext/sources/ggml/src/ggml.c +36 -14
  266. data/ext/sources/include/whisper.h +21 -0
  267. data/ext/sources/src/whisper.cpp +164 -14
  268. data/lib/whisper/log_settable.rb +5 -8
  269. data/lib/whisper/model/uri.rb +0 -7
  270. data/sig/whisper.rbs +6 -0
  271. data/test/test_vad.rb +9 -0
  272. data/test/test_vad_context.rb +2 -2
  273. data/whispercpp.gemspec +1 -1
  274. metadata +62 -37
  275. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-flash-attn-ops.c +0 -1878
  276. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-matmul-ops.c +0 -2066
  277. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-ops.c +0 -6
  278. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-ops.h +0 -88
  279. data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-profile.h +0 -34
  280. data/ext/sources/ggml/src/ggml-hexagon/htp/vtcm-utils.h +0 -16
  281. data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_gelu.cpp +0 -25
  282. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/abs.comp +0 -21
  283. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/ceil.comp +0 -22
  284. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/clamp.comp +0 -17
  285. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/cos.comp +0 -17
  286. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/elu.comp +0 -27
  287. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/exp.comp +0 -20
  288. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/floor.comp +0 -22
  289. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu.comp +0 -25
  290. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu_erf.comp +0 -39
  291. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu_quick.comp +0 -23
  292. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/hardsigmoid.comp +0 -22
  293. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/hardswish.comp +0 -22
  294. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/leaky_relu.comp +0 -22
  295. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/neg.comp +0 -20
  296. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/relu.comp +0 -21
  297. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/round.comp +0 -29
  298. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sgn.comp +0 -21
  299. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sigmoid.comp +0 -20
  300. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/silu.comp +0 -22
  301. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sin.comp +0 -17
  302. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/softplus.comp +0 -23
  303. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sqrt.comp +0 -17
  304. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/square.comp +0 -17
  305. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/step.comp +0 -22
  306. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/tanh.comp +0 -20
  307. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/trunc.comp +0 -22
  308. data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/xielu.comp +0 -35
@@ -3,228 +3,41 @@
3
3
  #pragma clang diagnostic ignored "-Wunused-but-set-variable"
4
4
 
5
5
  #include <assert.h>
6
+ #include <HAP_compute_res.h>
6
7
  #include <HAP_farf.h>
7
8
  #include <HAP_perf.h>
8
9
  #include <math.h>
10
+ #include <stdbool.h>
11
+ #include <stdatomic.h>
12
+ #include <stddef.h>
13
+ #include <stdint.h>
9
14
  #include <string.h>
10
15
 
11
16
  #include "hex-dma.h"
17
+ #include "hex-fastdiv.h"
18
+ #include "hex-profile.h"
19
+ #include "hmx-queue.h"
20
+ #include "hmx-utils.h"
12
21
  #include "hvx-utils.h"
13
22
  #include "hvx-dump.h"
23
+ #include "hvx-copy.h"
24
+ #include "hvx-reduce.h"
14
25
  #include "hvx-flash-attn.h"
26
+ #include "htp-vtcm.h"
27
+ #include "worker-pool.h"
15
28
 
16
29
  #define GGML_COMMON_DECL_C
17
30
  #include "ggml-common.h"
18
31
  #include "htp-ctx.h"
19
32
  #include "htp-ops.h"
20
- #include "htp-ops.h"
21
- #include "hmx-ops.h"
33
+
34
+ #include "flash-attn-ops.h"
35
+ #include "hvx-fa-kernels.h"
36
+ #include "hmx-fa-kernels.h"
22
37
 
23
38
  // Must be multiple of 32
24
39
  #define FLASH_ATTN_BLOCK_SIZE (32 * 2)
25
40
 
26
- #if __HVX_ARCH__ < 79
27
- #define HVX_OP_ADD_F32(a, b) Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_VsfVsf(a, b))
28
- #define HVX_OP_SUB_F32(a, b) Q6_Vsf_equals_Vqf32(Q6_Vqf32_vsub_VsfVsf(a, b))
29
- #define HVX_OP_MUL_F32(a, b) Q6_Vsf_equals_Vqf32(Q6_Vqf32_vmpy_VsfVsf(a, b))
30
- #else
31
- #define HVX_OP_ADD_F32(a, b) Q6_Vsf_vadd_VsfVsf(a, b)
32
- #define HVX_OP_SUB_F32(a, b) Q6_Vsf_vsub_VsfVsf(a, b)
33
- #define HVX_OP_MUL_F32(a, b) Q6_Vsf_vmpy_VsfVsf(a, b)
34
- #endif
35
-
36
- // This is a bit of a hack because the compiler is strugling to properly inline
37
- // the default hvx_vec_f32_to_f16 with output into the local array.
38
- static __attribute__((noinline)) void hvx_vec_f32_to_f16_a(void *ptr, HVX_Vector v0, HVX_Vector v1)
39
- {
40
- *(HVX_Vector *) ptr = hvx_vec_f32_to_f16(v0, v1);
41
- }
42
-
43
- // Dot product of two F16 vectors, accumulating to float
44
- static inline void hvx_dot_f16_f16_aa(float * restrict r, const void * restrict x, const void * restrict y, unsigned int n, float s) {
45
- const HVX_Vector * restrict vx = (const HVX_Vector * restrict) x; // fp16
46
- const HVX_Vector * restrict vy = (const HVX_Vector * restrict) y; // fp16
47
-
48
- uint32_t nvec = n / VLEN_FP16; // num full fp16 hvx vectors
49
- uint32_t nloe = n % VLEN_FP16; // leftover elements
50
-
51
- HVX_VectorPair rsum_p = Q6_W_vcombine_VV(Q6_V_vsplat_R(0), Q6_V_vsplat_R(0));
52
-
53
- uint32_t i = 0;
54
-
55
- #pragma unroll(4)
56
- for (i = 0; i < nvec; i++) {
57
- rsum_p = hvx_vec_mpyacc_f32_f16(rsum_p, vx[i], vy[i]);
58
- }
59
-
60
- if (nloe) {
61
- HVX_VectorPred bmask = Q6_Q_vsetq_R(nloe * 2);
62
- HVX_Vector y_hf = Q6_V_vand_QV(bmask, vy[i]);
63
- HVX_Vector x_hf = Q6_V_vand_QV(bmask, vx[i]);
64
-
65
- rsum_p = hvx_vec_mpyacc_f32_f16(rsum_p, x_hf, y_hf);
66
- }
67
-
68
- HVX_Vector rsum = HVX_OP_ADD_F32(Q6_V_lo_W(rsum_p), Q6_V_hi_W(rsum_p));
69
- rsum = HVX_OP_MUL_F32(hvx_vec_splat_f32(s), hvx_vec_reduce_sum_f32(rsum));
70
- hvx_vec_store_u(r, 4, rsum);
71
- }
72
-
73
- static inline HVX_Vector hvx_dot_f16_f16_aa_rx4(const void * restrict y,
74
- const uint8_t * restrict x,
75
- const size_t stride_x,
76
- const size_t nvec,
77
- const size_t nloe) {
78
- const HVX_Vector * restrict vx0 = (const HVX_Vector * restrict) x; // fp16
79
- const HVX_Vector * restrict vx1 = (const HVX_Vector * restrict) (x + stride_x); // fp16
80
- const HVX_Vector * restrict vx2 = (const HVX_Vector * restrict) (x + stride_x * 2); // fp16
81
- const HVX_Vector * restrict vx3 = (const HVX_Vector * restrict) (x + stride_x * 3); // fp16
82
- const HVX_Vector * restrict vy = (const HVX_Vector * restrict) y; // fp16
83
-
84
- HVX_VectorPair rsum0_p = Q6_W_vcombine_VV(Q6_V_vsplat_R(0), Q6_V_vsplat_R(0));
85
- HVX_VectorPair rsum1_p = Q6_W_vcombine_VV(Q6_V_vsplat_R(0), Q6_V_vsplat_R(0));
86
- HVX_VectorPair rsum2_p = Q6_W_vcombine_VV(Q6_V_vsplat_R(0), Q6_V_vsplat_R(0));
87
- HVX_VectorPair rsum3_p = Q6_W_vcombine_VV(Q6_V_vsplat_R(0), Q6_V_vsplat_R(0));
88
-
89
- uint32_t i = 0;
90
-
91
- for (i = 0; i < nvec; i++) {
92
- HVX_Vector y_hf = vy[i];
93
- HVX_Vector x0_hf = vx0[i];
94
- HVX_Vector x1_hf = vx1[i];
95
- HVX_Vector x2_hf = vx2[i];
96
- HVX_Vector x3_hf = vx3[i];
97
-
98
- rsum0_p = hvx_vec_mpyacc_f32_f16(rsum0_p, x0_hf, y_hf);
99
- rsum1_p = hvx_vec_mpyacc_f32_f16(rsum1_p, x1_hf, y_hf);
100
- rsum2_p = hvx_vec_mpyacc_f32_f16(rsum2_p, x2_hf, y_hf);
101
- rsum3_p = hvx_vec_mpyacc_f32_f16(rsum3_p, x3_hf, y_hf);
102
- }
103
-
104
- if (nloe) {
105
- // Load x (fp16) and zero-out unused elements
106
- HVX_VectorPred bmask = Q6_Q_vsetq_R(nloe * 2);
107
- HVX_Vector y_hf = Q6_V_vand_QV(bmask, vy[i]);
108
- HVX_Vector x0_hf = Q6_V_vand_QV(bmask, vx0[i]);
109
- HVX_Vector x1_hf = Q6_V_vand_QV(bmask, vx1[i]);
110
- HVX_Vector x2_hf = Q6_V_vand_QV(bmask, vx2[i]);
111
- HVX_Vector x3_hf = Q6_V_vand_QV(bmask, vx3[i]);
112
-
113
- rsum0_p = hvx_vec_mpyacc_f32_f16(rsum0_p, x0_hf, y_hf);
114
- rsum1_p = hvx_vec_mpyacc_f32_f16(rsum1_p, x1_hf, y_hf);
115
- rsum2_p = hvx_vec_mpyacc_f32_f16(rsum2_p, x2_hf, y_hf);
116
- rsum3_p = hvx_vec_mpyacc_f32_f16(rsum3_p, x3_hf, y_hf);
117
- }
118
-
119
- HVX_Vector rsum0 = HVX_OP_ADD_F32(Q6_V_lo_W(rsum0_p), Q6_V_hi_W(rsum0_p));
120
- HVX_Vector rsum1 = HVX_OP_ADD_F32(Q6_V_lo_W(rsum1_p), Q6_V_hi_W(rsum1_p));
121
- HVX_Vector rsum2 = HVX_OP_ADD_F32(Q6_V_lo_W(rsum2_p), Q6_V_hi_W(rsum2_p));
122
- HVX_Vector rsum3 = HVX_OP_ADD_F32(Q6_V_lo_W(rsum3_p), Q6_V_hi_W(rsum3_p));
123
-
124
- HVX_Vector_x4 rsum0123 = { .v = { rsum0, rsum1, rsum2, rsum3 } };
125
- return hvx_vec_reduce_sum_f32x4(rsum0123);
126
- }
127
-
128
- static inline HVX_Vector hvx_dot_f16_f16_aa_rx32(const void * restrict y,
129
- const uint8_t * restrict x,
130
- const size_t stride_x,
131
- const size_t n,
132
- float s) {
133
-
134
- const size_t nvec = n / VLEN_FP16; // num full fp16 hvx vectors
135
- const size_t nloe = n % VLEN_FP16; // leftover elements
136
-
137
- HVX_Vector sums = Q6_V_vzero();
138
- const size_t stride_x_4 = stride_x * 4;
139
- for (uint32_t j = 0; j < VLEN_FP32; j += 4) {
140
- HVX_Vector sums_x4 = hvx_dot_f16_f16_aa_rx4(y, x, stride_x, nvec, nloe);
141
- HVX_VectorPred pred = Q6_Q_vsetq_R(j * SIZEOF_FP32);
142
- sums = Q6_V_vmux_QVV(pred, sums, sums_x4);
143
- x += stride_x_4;
144
- }
145
-
146
- return HVX_OP_MUL_F32(hvx_vec_splat_f32(s), sums);
147
- }
148
-
149
- // MAD: y (F32) += x (F16) * s (F16)
150
- static inline void hvx_mad_f32_f16_aa(float * restrict y, const void * restrict x, const __fp16 * restrict s, int n) {
151
- const HVX_Vector * restrict vx0 = (const HVX_Vector *) x;
152
-
153
- HVX_VectorPair * restrict vy_p = (HVX_VectorPair *) y;
154
- HVX_Vector * restrict vy = (HVX_Vector *) y;
155
-
156
- uint32_t nvec = n / VLEN_FP16; // num full fp16 hvx vectors
157
- uint32_t nloe = n % VLEN_FP16; // leftover elements
158
-
159
- HVX_Vector S0 = hvx_vec_splat_f16(*s);
160
-
161
- uint32_t i = 0;
162
-
163
- #pragma unroll(2)
164
- for (i = 0; i < nvec; ++i) {
165
- vy_p[i] = hvx_vec_mpyacc_f32_f16(vy_p[i], Q6_Vh_vshuff_Vh(vx0[i]), S0);
166
- }
167
-
168
- if (nloe) {
169
- HVX_VectorPair xy_p = vy_p[i];
170
- xy_p = hvx_vec_mpyacc_f32_f16(xy_p, Q6_Vh_vshuff_Vh(vx0[i]), S0);
171
-
172
- HVX_Vector xy = Q6_V_lo_W(xy_p);
173
- i = 2 * i; // index for vy
174
-
175
- if (nloe >= VLEN_FP32) {
176
- vy[i] = xy;
177
- nloe -= VLEN_FP32; ++i; xy = Q6_V_hi_W(xy_p);
178
- }
179
-
180
- if (nloe) {
181
- hvx_vec_store_a(&vy[i], nloe * 4, xy);
182
- }
183
- }
184
- }
185
-
186
- // MAD: y (F32) += x0 (F16) * s0 (F16) + x1 (F16) * s1 (F16)
187
- static inline void hvx_mad_f32_f16_aa_rx2(float * restrict y, const void * restrict x0, const void * restrict x1,
188
- const __fp16 * restrict s0, const __fp16 * restrict s1, int n) {
189
- const HVX_Vector * restrict vx0 = (const HVX_Vector *) x0;
190
- const HVX_Vector * restrict vx1 = (const HVX_Vector *) x1;
191
-
192
- HVX_VectorPair * restrict vy_p = (HVX_VectorPair *) y;
193
- HVX_Vector * restrict vy = (HVX_Vector *) y;
194
-
195
- uint32_t nvec = n / VLEN_FP16; // num full fp16 hvx vectors
196
- uint32_t nloe = n % VLEN_FP16; // leftover elements
197
-
198
- HVX_Vector S0 = hvx_vec_splat_f16(*s0);
199
- HVX_Vector S1 = hvx_vec_splat_f16(*s1);
200
-
201
- uint32_t i = 0;
202
-
203
- #pragma unroll(2)
204
- for (i = 0; i < nvec; ++i) {
205
- vy_p[i] = hvx_vec_mpyacc_f32_f16(vy_p[i], Q6_Vh_vshuff_Vh(vx0[i]), S0);
206
- vy_p[i] = hvx_vec_mpyacc_f32_f16(vy_p[i], Q6_Vh_vshuff_Vh(vx1[i]), S1);
207
- }
208
-
209
- if (nloe) {
210
- HVX_VectorPair xy_p = vy_p[i];
211
- xy_p = hvx_vec_mpyacc_f32_f16(xy_p, Q6_Vh_vshuff_Vh(vx0[i]), S0);
212
- xy_p = hvx_vec_mpyacc_f32_f16(xy_p, Q6_Vh_vshuff_Vh(vx1[i]), S1);
213
-
214
- HVX_Vector xy = Q6_V_lo_W(xy_p);
215
- i = 2 * i; // index for vy
216
-
217
- if (nloe >= VLEN_FP32) {
218
- vy[i] = xy;
219
- nloe -= VLEN_FP32; ++i; xy = Q6_V_hi_W(xy_p);
220
- }
221
-
222
- if (nloe) {
223
- hvx_vec_store_a(&vy[i], nloe * 4, xy);
224
- }
225
- }
226
- }
227
-
228
41
  struct htp_fa_context {
229
42
  const struct htp_ops_context * octx;
230
43
 
@@ -241,12 +54,12 @@ struct htp_fa_context {
241
54
 
242
55
  float scale;
243
56
  float max_bias;
244
- float logit_softcap;
57
+ __fp16 logit_softcap;
245
58
 
246
59
  uint32_t n_head_log2;
247
60
  float m0;
248
61
  float m1;
249
- float slopes[512];
62
+ __fp16 slopes[512];
250
63
 
251
64
  uint32_t n_blocks;
252
65
 
@@ -263,28 +76,80 @@ struct htp_fa_context {
263
76
 
264
77
  bool is_q_fp32;
265
78
 
266
- uint64_t t_start;
267
- };
268
-
269
- static inline void hvx_scale_vec_f32_aa(uint8_t * restrict dst, const uint8_t * restrict src, const int n, HVX_Vector vs) {
270
- assert((size_t) dst % 128 == 0);
271
- assert((size_t) src % 128 == 0);
79
+ size_t size_q_block;
80
+ size_t size_vkq_acc;
272
81
 
273
- const HVX_Vector * restrict vsrc = (const HVX_Vector * restrict) src;
274
- HVX_Vector * restrict vdst = (HVX_Vector * restrict) dst;
82
+ uint8_t * spad_q;
83
+ uint8_t * spad_k;
84
+ uint8_t * spad_v;
85
+ uint8_t * spad_m;
86
+ uint8_t * spad_a;
275
87
 
276
- const uint32_t nvec = n / VLEN_FP32;
277
- const uint32_t nloe = n % VLEN_FP32;
88
+ uint64_t t_start;
89
+ };
278
90
 
279
- uint32_t i = 0;
280
- #pragma unroll(4)
281
- for (; i < nvec; ++i) {
282
- vdst[i] = HVX_OP_MUL_F32(vsrc[i], vs);
283
- }
284
- if (nloe) {
285
- hvx_vec_store_a(&vdst[i], nloe * sizeof(float), HVX_OP_MUL_F32(vsrc[i], vs));
286
- }
287
- }
91
+ struct hmx_fa_context {
92
+ const struct htp_ops_context * octx;
93
+ const struct htp_tensor * sinks; // attention sinks (src[4]), NULL if absent
94
+ bool pipeline; // true when n_kv_blocks >= FA_MIN_KV_BLOCKS && n_threads >= 2
95
+ uint32_t n_threads;
96
+
97
+ // Op parameters
98
+ __fp16 scale;
99
+ float max_bias;
100
+ __fp16 logit_softcap;
101
+ uint32_t n_head_log2;
102
+ float m0, m1;
103
+
104
+ // Dimensions
105
+ uint32_t DK, DV;
106
+ uint32_t n_kv; // kv_len
107
+ uint32_t n_kv_heads; // number of KV heads
108
+ uint32_t n_heads; // number of Q heads
109
+ uint32_t G; // GQA factor = n_heads / n_kv_heads
110
+ struct fastdiv_values div_G;
111
+ struct fastdiv_values src3_div2;
112
+ struct fastdiv_values src3_div3;
113
+ uint32_t n_kv_blocks;
114
+ uint32_t neq1; // Q token count
115
+
116
+ // Types
117
+ bool is_q_fp32;
118
+ bool is_dst_fp32;
119
+
120
+ // Dynamic block sizes
121
+ uint32_t Br; // Q tokens per block (before GQA expansion)
122
+ uint32_t Bc;
123
+ uint32_t g_br; // hex_align_up(G * Br, 32) - actual tile row dim
124
+
125
+ // VTCM buffers (allocated by vtcm_seq_alloc)
126
+ __fp16 * vtcm_q_tiles; // Q tile format [g_br, D]
127
+ __fp16 * vtcm_o_tiles[2]; // O ping-pong [g_br, D]
128
+ __fp16 * vtcm_k_fp16[2]; // K DMA double-buffer [Bc, D]
129
+ __fp16 * vtcm_v_fp16[2]; // V DMA double-buffer [Bc, D]
130
+ __fp16 * vtcm_k_tiles; // K tiles (transposed)
131
+ __fp16 * vtcm_v_tiles[2]; // V tiles (column-major, double-buffered)
132
+ __fp16 * vtcm_s_tiles; // S = QK^T [g_br, Bc]
133
+ __fp16 * vtcm_p_tiles; // P = softmax(S) [g_br, Bc]
134
+ __fp16 * vtcm_d_tiles; // Diagonal rescale [g_br, g_br]
135
+ HVX_Vector * vtcm_m_vec; // Row max [g_br]
136
+ HVX_Vector * vtcm_l_vec; // Row sum [g_br]
137
+ HVX_Vector * vtcm_s_rowmax; // Softmax intermediate [g_br]
138
+ HVX_Vector * vtcm_p_rowsum; // Softmax intermediate [g_br]
139
+ HVX_Vector * vtcm_row_bufs; // Per-thread softmax row scratch [n_threads][2][Bc/64]
140
+ uint8_t * vtcm_hmx_scales_id; // HMX output scales (identity)
141
+ uint8_t * vtcm_hmx_scales_qk; // HMX output scales (qk_scale)
142
+ __fp16 * vtcm_mask_buf; // VTCM mask buffer [Br * m_line], DMA'd per KV block
143
+ __fp16 * vtcm_slopes; // ALiBi slopes [g_br]
144
+ size_t row_buf_stride; // HVX vectors per row buffer (Bc/64)
145
+ size_t mask_buf_row_stride; // elements (__fp16) per row in mask buffer
146
+ size_t q_tile_bytes;
147
+ size_t o_tile_bytes;
148
+ size_t col_vec_bytes;
149
+ size_t d_tile_bytes;
150
+ bool mask_broadcast; // true when mask->ne[2] == 1 (head-independent, single 2D DMA)
151
+ dma_cache m_cache;
152
+ };
288
153
 
289
154
  static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void * data) {
290
155
  struct htp_fa_context * factx = (struct htp_fa_context *) data;
@@ -339,6 +204,8 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
339
204
 
340
205
  if (ir0 >= ir1) return;
341
206
 
207
+ struct htp_thread_trace * tr = octx->ctx ? &octx->ctx->trace[ith] : NULL;
208
+
342
209
  dma_queue * dma = octx->ctx->dma[ith];
343
210
 
344
211
  const uint32_t DK = nek0;
@@ -349,16 +216,14 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
349
216
  const size_t size_v_row = DV * sizeof(__fp16);
350
217
 
351
218
  // Scratchpad buffers for Q, K, V, Mask, and VKQ32 accumulator
352
- uint8_t * spad_q = octx->src0_spad.data + octx->src0_spad.size_per_thread * ith;
353
- uint8_t * spad_k = octx->src1_spad.data + octx->src1_spad.size_per_thread * ith;
354
- uint8_t * spad_v = octx->src2_spad.data + octx->src2_spad.size_per_thread * ith;
355
- uint8_t * spad_m = octx->src3_spad.data + octx->src3_spad.size_per_thread * ith;
356
- uint8_t * spad_a = octx->dst_spad.data + octx->dst_spad.size_per_thread * ith;
357
-
358
- const HVX_Vector logit_cap = hvx_vec_splat_f32(factx->logit_softcap);
219
+ uint8_t * spad_q = factx->spad_q + factx->size_q_block * ith;
220
+ uint8_t * spad_k = factx->spad_k + factx->size_k_block * 2 * ith;
221
+ uint8_t * spad_v = factx->spad_v + factx->size_v_block * 2 * ith;
222
+ uint8_t * spad_m = factx->spad_m + (mask ? factx->size_m_block * HVX_FA_DMA_CACHE_SIZE : 0) * ith;
223
+ uint8_t * spad_a = factx->spad_a + factx->size_vkq_acc * ith;
359
224
 
360
225
  dma_cache m_cache;
361
- dma_cache_init(&m_cache, spad_m, factx->size_m_block, DMA_CACHE_MAX_SIZE);
226
+ dma_cache_init(&m_cache, spad_m, factx->size_m_block, HVX_FA_DMA_CACHE_SIZE);
362
227
 
363
228
  for (uint32_t ir = ir0; ir < ir1; ++ir) {
364
229
  const uint32_t iq3 = fastdiv(ir, &factx->src0_div21);
@@ -375,9 +240,6 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
375
240
  const uint8_t * q_row_ptr = (const uint8_t *) q->data + (iq1*nbq1 + iq2*nbq2 + iq3*nbq3);
376
241
  dma_queue_push(dma, dma_make_ptr(spad_q, q_row_ptr), factx->size_q_row_padded, nbq1, size_q_row, 1);
377
242
 
378
- // FARF(HIGH, "fa %u: prefetch Q: ir %u iq1 %u iq2 %u iq3 %u q_row_ptr %p size %u : usec %u", ith, ir, iq1, iq2, iq3, q_row_ptr, size_q_row,
379
- // (unsigned)HAP_perf_qtimer_count_to_us(HAP_perf_get_qtimer_count() - factx->t_start));
380
-
381
243
  const __fp16 * mp_base = NULL;
382
244
  if (mask) {
383
245
  const uint32_t im2 = fastmodulo(iq2, mask->ne[2], &factx->src3_div2);
@@ -406,18 +268,13 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
406
268
  // Mask is 1D contiguous for this row
407
269
  dma_cache_push(dma, &m_cache, m_src, current_block_size * 2, current_block_size * 2, current_block_size * 2, 1);
408
270
  }
409
-
410
- // FARF(HIGH, "fa %u: prefetch KVM: ir %u ib %u iq1 %u iq2 %u iq3 %u : size_k_row %u size_v_row %u bs %u: usec %u",
411
- // ith, ir, ib, iq1, iq2, iq3,
412
- // size_k_row, size_v_row, current_block_size,
413
- // (unsigned)HAP_perf_qtimer_count_to_us(HAP_perf_get_qtimer_count() - factx->t_start));
414
271
  }
415
272
 
416
273
  const uint32_t h = iq2; // head index
417
- const float slope = factx->slopes[h];
274
+ const __fp16 slope = factx->slopes[h];
418
275
 
419
276
  HVX_Vector S_vec = hvx_vec_splat_f32(0.0f);
420
- HVX_Vector M_vec = hvx_vec_splat_f32(-INFINITY);
277
+ HVX_Vector M_vec = hvx_vec_splat_f32(HTP_FA_M_INITIAL_VAL);
421
278
 
422
279
  // Clear accumulator
423
280
  hvx_splat_f32_a(spad_a, 0, DV);
@@ -429,6 +286,7 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
429
286
  }
430
287
 
431
288
  const HVX_Vector slope_vec = hvx_vec_splat_f16(slope);
289
+ const HVX_Vector v_neg_inf = Q6_Vh_vsplat_R(0xfbff);
432
290
  for (uint32_t ib = 0; ib < factx->n_blocks; ++ib) {
433
291
  const uint32_t ic_start = ib * FLASH_ATTN_BLOCK_SIZE;
434
292
  const uint32_t current_block_size = MIN(FLASH_ATTN_BLOCK_SIZE, nek1 - ic_start);
@@ -438,113 +296,101 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
438
296
  uint8_t * v_base = dma_queue_pop(dma).dst; // V
439
297
  __fp16 * m_base = mask ? dma_queue_pop(dma).dst : NULL; // M
440
298
 
441
- // FARF(HIGH, "fa %u: process: ir %u ib %u : iq1 %u iq2 %u iq3 %u q_ptr_vtcm %p : usec %u",
442
- // ith, ir, ib, iq1, iq2, iq3, q_ptr_vtcm,
443
- // (unsigned)HAP_perf_qtimer_count_to_us(HAP_perf_get_qtimer_count() - factx->t_start));
299
+ htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_FA_QK, ir);
444
300
 
445
301
  // Inner loop processing the block from VTCM
446
- uint32_t ic = 0;
447
-
448
- // Process in sub-blocks of 32 (VLEN_FP32)
449
- HVX_Vector sb_scores[FLASH_ATTN_BLOCK_SIZE / VLEN_FP32];
450
- HVX_Vector v_max = hvx_vec_splat_f32(-INFINITY);
451
- for (uint32_t iv = 0; ic < current_block_size; ic += VLEN_FP32, ++iv) {
452
- // 1. Compute scores
453
- HVX_Vector scores = hvx_dot_f16_f16_aa_rx32(q_ptr_vtcm, k_base + ic * factx->size_k_row_padded, factx->size_k_row_padded, DK, factx->scale);
454
-
455
- // 2. Softcap
456
- if (factx->logit_softcap != 0.0f) {
457
- scores = hvx_vec_tanh_f32(scores);
458
- scores = HVX_OP_MUL_F32(scores, logit_cap);
459
- }
302
+ // 1. Compute scores (64 elements FP16)
303
+ HVX_Vector scores_f16 = Q6_V_vzero();
304
+ if (current_block_size > 0) {
305
+ HVX_Vector scores0 = hvx_dot_f16_f16_aa_rx32(q_ptr_vtcm, k_base, factx->size_k_row_padded, DK, factx->scale);
306
+ HVX_Vector scores1 = (current_block_size > 32) ? hvx_dot_f16_f16_aa_rx32(q_ptr_vtcm, k_base + 32 * factx->size_k_row_padded, factx->size_k_row_padded, DK, factx->scale) : Q6_V_vzero();
307
+ scores_f16 = hvx_vec_f32_to_f16(scores0, scores1);
308
+ }
460
309
 
461
- // 3. Mask
462
- if (mask) {
463
- const __fp16 * mp = m_base + ic;
464
- HVX_Vector m_vals_f16 = *(const HVX_UVector *) mp;
465
-
466
- // Multiplying -INFINITY (0xFC00) by a slope in VhfVhf instructions can incorrectly produce NaN on v79.
467
- // Clamp -INFINITY to the max negative fp16 finite value (-65504.0f).
468
- HVX_Vector vinf = Q6_Vh_vsplat_R(0xFC00);
469
- HVX_Vector vmin = Q6_Vh_vsplat_R(0xFBFF);
470
- HVX_VectorPred is_inf = Q6_Q_vcmp_eq_VhVh(m_vals_f16, vinf);
471
- m_vals_f16 = Q6_V_vmux_QVV(is_inf, vmin, m_vals_f16);
472
-
473
- #if __HVX_ARCH__ >= 79
474
- HVX_VectorPair m_vals_f32_pair = Q6_Wsf_vmpy_VhfVhf(Q6_Vh_vshuff_Vh(m_vals_f16), slope_vec);
475
- HVX_Vector add_val = Q6_V_lo_W(m_vals_f32_pair);
476
- scores = Q6_Vsf_vadd_VsfVsf(add_val, scores);
477
- #else
478
- HVX_VectorPair m_vals_f32_pair = Q6_Wqf32_vmpy_VhfVhf(Q6_Vh_vshuff_Vh(m_vals_f16), slope_vec);
479
- HVX_Vector add_val = Q6_V_lo_W(m_vals_f32_pair);
480
- scores = Q6_Vsf_equals_Vqf32(Q6_Vqf32_vadd_Vqf32Vsf(add_val, scores));
481
- #endif
482
- }
310
+ // 2. Softcap (in FP16)
311
+ if (factx->logit_softcap != 0.0f) {
312
+ const HVX_Vector v_cap = hvx_vec_splat_f16(factx->logit_softcap);
313
+ scores_f16 = hvx_vec_tanh_f16(scores_f16);
314
+ scores_f16 = hvx_vec_mul_f16_f16(scores_f16, v_cap);
315
+ }
483
316
 
484
- // Mask out invalid lanes for leftover handling
485
- uint32_t valid_lanes = current_block_size - ic;
486
- if (valid_lanes < VLEN_FP32) {
487
- HVX_VectorPred valid_pred = Q6_Q_vsetq_R(valid_lanes * 4); // 4 bytes per fp32 lane
488
- scores = Q6_V_vmux_QVV(valid_pred, scores, hvx_vec_splat_f32(-INFINITY));
489
- }
317
+ HVX_VectorPred q_tail_keep = Q6_Q_vsetq2_R(current_block_size * sizeof(__fp16));
490
318
 
491
- sb_scores[iv] = scores;
492
- v_max = hvx_vec_reduce_max2_f32(scores, v_max); // All lanes have block max
319
+ // 3. Mask (in FP16)
320
+ if (mask) {
321
+ HVX_Vector m_vals_f16 = *(const HVX_UVector *) m_base;
322
+ HVX_Vector vinf = Q6_Vh_vsplat_R(0xFC00);
323
+ HVX_Vector vmin = Q6_Vh_vsplat_R(0xFBFF);
324
+ HVX_VectorPred is_inf = Q6_Q_vcmp_eq_VhVh(m_vals_f16, vinf);
325
+ m_vals_f16 = Q6_V_vmux_QVV(is_inf, vmin, m_vals_f16);
326
+
327
+ HVX_Vector m_scaled = hvx_vec_mul_f16_f16(m_vals_f16, slope_vec);
328
+ scores_f16 = Q6_V_vmux_QVV(q_tail_keep, hvx_vec_add_f16_f16(scores_f16, m_scaled), v_neg_inf);
329
+ } else {
330
+ scores_f16 = Q6_V_vmux_QVV(q_tail_keep, scores_f16, v_neg_inf);
493
331
  }
494
332
 
333
+ // Compute block max in FP16
334
+ HVX_Vector v_max_f16 = hvx_vec_reduce_max_f16(scores_f16);
335
+ HVX_Vector v_max = Q6_V_lo_W(hvx_vec_f16_to_f32(v_max_f16)); // splat block max in FP32
336
+ htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_FA_QK, ir);
337
+
338
+ htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_FA_SFM, ir);
495
339
  {
340
+ const HVX_Vector v_log2e = hvx_vec_splat_f16(EXP_LOG2E_F);
341
+
496
342
  // 4. Online Softmax Update
497
343
  HVX_Vector M_new_vec = Q6_Vsf_vmax_VsfVsf(v_max, M_vec);
498
344
  HVX_Vector diff_vec = HVX_OP_SUB_F32(M_vec, M_new_vec);
499
- HVX_Vector ms_vec = hvx_vec_exp_f32(diff_vec);
345
+
346
+ HVX_Vector diff_f16 = hvx_vec_f32_to_f16(diff_vec, diff_vec);
347
+ HVX_Vector diff_base2 = hvx_vec_mul_f16_f16(diff_f16, v_log2e);
348
+ HVX_Vector ms_f16 = hvx_vec_exp2_f16(diff_base2);
349
+ HVX_Vector ms_vec = Q6_V_lo_W(hvx_vec_f16_to_f32(ms_f16));
350
+
500
351
  M_vec = M_new_vec;
501
352
 
502
353
  hvx_scale_vec_f32_aa((uint8_t *) VKQ32, (const uint8_t *) VKQ32, DV, ms_vec);
503
354
 
504
- HVX_Vector p_sum_vec = hvx_vec_splat_f32(0.0f);
505
- for (uint32_t ic2 = 0, iv = 0; ic2 < current_block_size; ic2 += VLEN_FP32, ++iv) {
506
- HVX_Vector scores = sb_scores[iv];
507
- HVX_Vector scores_shifted = HVX_OP_SUB_F32(scores, M_vec);
508
- HVX_Vector P = hvx_vec_exp_f32(scores_shifted);
355
+ // Compute P = exp2((S - M) * log2(e)) in FP16
356
+ HVX_Vector v_m_vec_f16 = hvx_vec_f32_to_f16(M_vec, M_vec);
357
+ HVX_Vector v_s_minus_m = Q6_Vqf16_vsub_VhfVhf(scores_f16, v_m_vec_f16);
509
358
 
510
- p_sum_vec = HVX_OP_ADD_F32(p_sum_vec, P);
359
+ HVX_Vector v_s_minus_m_base2 = hvx_vec_mul_f16_f16(Q6_Vhf_equals_Vqf16(v_s_minus_m), v_log2e);
511
360
 
512
- // 5. Accumulate V
513
- __fp16 __attribute__((aligned(VLEN))) p_arr[VLEN_FP16];
514
- hvx_vec_f32_to_f16_a(p_arr, P, hvx_vec_splat_f32(0));
361
+ HVX_Vector P = hvx_vec_exp2_f16(v_s_minus_m_base2);
362
+ P = Q6_V_vmux_QVV(q_tail_keep, P, Q6_V_vzero());
515
363
 
516
- float __attribute__((aligned(128))) P_arr[VLEN_FP32];
517
- hvx_vec_store_a(P_arr, 128, P);
364
+ // Convert P to FP32 to update the running sum S_vec
365
+ HVX_VectorPair P_pair = hvx_vec_f16_to_f32(P);
366
+ HVX_Vector P0 = Q6_V_lo_W(P_pair);
367
+ HVX_Vector P1 = Q6_V_hi_W(P_pair);
368
+ HVX_Vector p_sum_vec = hvx_vec_reduce_sum_f32(HVX_OP_ADD_F32(P0, P1));
518
369
 
519
- for (uint32_t j = 0; j < VLEN_FP32; j += 2) {
520
- const uint32_t cur_ic = ic2 + j;
521
- if (cur_ic >= current_block_size) {
522
- break;
523
- }
370
+ S_vec = HVX_OP_ADD_F32(HVX_OP_MUL_F32(S_vec, ms_vec), p_sum_vec);
524
371
 
525
- if (cur_ic + 1 == current_block_size) {
526
- // Odd leftover, process single row
527
- if (P_arr[j] != 0.0f) {
528
- const uint8_t * v_ptr = v_base + cur_ic * factx->size_v_row_padded;
529
- hvx_mad_f32_f16_aa(VKQ32, v_ptr, (p_arr + j), DV);
530
- }
531
- break;
532
- }
372
+ // 5. Accumulate V (F16 * F16 -> F32 accumulator)
373
+ __fp16 __attribute__((aligned(128))) p_arr[VLEN_FP16];
374
+ hvx_vec_store_a(p_arr, 128, P);
533
375
 
534
- // Avoid NaN * 0.0 = NaN for uninitialized V cache rows.
535
- // Check the f32 values to safely avoid strict aliasing violations.
536
- if (P_arr[j] == 0.0f && P_arr[j + 1] == 0.0f) {
537
- continue;
376
+ for (uint32_t j = 0; j < current_block_size; j += 2) {
377
+ if (j + 1 == current_block_size) {
378
+ if (p_arr[j] != 0.0f) {
379
+ const uint8_t * v_ptr = v_base + j * factx->size_v_row_padded;
380
+ hvx_mad_f32_f16_aa(VKQ32, v_ptr, (p_arr + j), DV);
538
381
  }
382
+ break;
383
+ }
539
384
 
540
- const uint8_t * v_ptr = v_base + cur_ic * factx->size_v_row_padded;
541
- hvx_mad_f32_f16_aa_rx2(VKQ32, v_ptr, v_ptr + factx->size_v_row_padded, (p_arr + j), (p_arr + j + 1), DV);
385
+ if (p_arr[j] == 0.0f && p_arr[j + 1] == 0.0f) {
386
+ continue;
542
387
  }
543
- }
544
388
 
545
- p_sum_vec = hvx_vec_reduce_sum_f32(p_sum_vec);
546
- S_vec = HVX_OP_ADD_F32(HVX_OP_MUL_F32(S_vec, ms_vec), p_sum_vec);
389
+ const uint8_t * v_ptr = v_base + j * factx->size_v_row_padded;
390
+ hvx_mad_f32_f16_aa_rx2(VKQ32, v_ptr, v_ptr + factx->size_v_row_padded, (p_arr + j), (p_arr + j + 1), DV);
391
+ }
547
392
  }
393
+ htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_FA_SFM, ir);
548
394
 
549
395
  // Issue DMA for next+1 block (if exists)
550
396
  if (ib + 2 < factx->n_blocks) {
@@ -565,14 +411,10 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
565
411
  const uint8_t * m_src = (const uint8_t *) (mp_base + next_ic_start);
566
412
  dma_cache_push(dma, &m_cache, m_src, next_block_size * 2, next_block_size * 2, next_block_size * 2, 1);
567
413
  }
568
-
569
- // FARF(HIGH, "fa %u: prefetch KVM: ir %u ib %u : iq1 %u iq2 %u iq3 %u : size_k_row %u size_v_row %u bs %u: usec %u",
570
- // ith, ir, next_ib, iq1, iq2, iq3,
571
- // size_k_row, size_v_row, next_block_size,
572
- // (unsigned)HAP_perf_qtimer_count_to_us(HAP_perf_get_qtimer_count() - factx->t_start));
573
414
  }
574
415
  }
575
416
 
417
+ htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_O_PROC, ir);
576
418
  // sinks
577
419
  float M = hvx_vec_get_f32(M_vec);
578
420
  float S = hvx_vec_get_f32(S_vec);
@@ -601,9 +443,9 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
601
443
 
602
444
  // Store result
603
445
  // dst indices
604
- const int i1 = iq1;
605
- const int i2 = iq2;
606
- const int i3 = iq3;
446
+ const uint32_t i1 = iq1;
447
+ const uint32_t i2 = iq2;
448
+ const uint32_t i3 = iq3;
607
449
 
608
450
  // dst is permuted: [DV, n_heads, n_tokens, n_seq]
609
451
  // head stride is nb[1], token stride is nb[2], batch stride is nb[3]
@@ -614,9 +456,1544 @@ static void flash_attn_ext_f16_thread(unsigned int nth, unsigned int ith, void *
614
456
  } else if (dst->type == HTP_TYPE_F16) {
615
457
  hvx_copy_f16_f32_ua(dst_ptr, (uint8_t *) VKQ32, DV);
616
458
  }
459
+ htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_O_PROC, ir);
617
460
  }
618
461
  }
619
462
 
463
+ // ============================================================================
464
+ // HMX Phase args and thread logic
465
+ // ============================================================================
466
+
467
+ typedef struct {
468
+ struct hmx_fa_context * factx;
469
+ uint32_t kv_rows;
470
+ size_t src_stride;
471
+ void * curr_k;
472
+ uint32_t kv_start;
473
+ uint32_t rows_per_t;
474
+ } fa_k_int_args_t;
475
+
476
+ static void fa_k_interleave_thread(unsigned int n, unsigned int i, void * data) {
477
+ fa_k_int_args_t * args = (fa_k_int_args_t *) data;
478
+ struct hmx_fa_context * factx = args->factx;
479
+
480
+ const uint32_t total_rows = args->kv_rows;
481
+ const uint32_t rows_per_t = args->rows_per_t;
482
+ const uint32_t start = i * rows_per_t;
483
+ const uint32_t end = (uint32_t) hex_smin(start + rows_per_t, total_rows);
484
+
485
+ if (start >= total_rows) {
486
+ return;
487
+ }
488
+
489
+ struct htp_thread_trace * tr = factx->octx->ctx ? &factx->octx->ctx->trace[i] : NULL;
490
+ htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_FA_K_PREP, (uint16_t) (args->kv_start + start));
491
+ hmx_interleave_rows_to_tiles(factx->vtcm_k_tiles, (const __fp16 *) args->curr_k, total_rows, factx->DK,
492
+ args->src_stride, start, end);
493
+ htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_FA_K_PREP, (uint16_t) (args->kv_start + start));
494
+ }
495
+
496
+ static void fa_phase_k_interleave(struct hmx_fa_context * factx, uint32_t kv_rows, size_t src_stride, void * curr_k, uint32_t kv_start) {
497
+ worker_pool_context_t wp = factx->octx->ctx->worker_pool;
498
+ uint32_t n = 1;
499
+ if (factx->n_threads > 1 && kv_rows >= factx->n_threads * 2) {
500
+ n = factx->n_threads;
501
+ }
502
+ uint32_t rows_per_t = hex_align_up(hmx_ceil_div(kv_rows, n), 2);
503
+ fa_k_int_args_t args = { factx, kv_rows, src_stride, curr_k, kv_start, rows_per_t };
504
+ if (n > 1) {
505
+ worker_pool_run_func(wp, fa_k_interleave_thread, &args, n);
506
+ } else {
507
+ fa_k_interleave_thread(1, 0, &args);
508
+ }
509
+ }
510
+
511
+ typedef struct {
512
+ struct hmx_fa_context * factx;
513
+ uint32_t kv_rows;
514
+ size_t src_stride;
515
+ void * v_src;
516
+ void * v_tiles_dst;
517
+ size_t n_col_tiles;
518
+ uint32_t kv_start;
519
+ uint32_t rows_per_t;
520
+ } fa_v_int_args_t;
521
+
522
+ static void fa_v_interleave_thread(unsigned int n, unsigned int i, void * data) {
523
+ fa_v_int_args_t * args = (fa_v_int_args_t *) data;
524
+ struct hmx_fa_context * factx = args->factx;
525
+
526
+ const uint32_t total_rows = args->kv_rows;
527
+ const uint32_t rows_per_t = args->rows_per_t;
528
+ const uint32_t start = i * rows_per_t;
529
+ const uint32_t end = (uint32_t) hex_smin(start + rows_per_t, total_rows);
530
+
531
+ if (start >= total_rows) {
532
+ return;
533
+ }
534
+
535
+ __fp16 * v_tiles_dst = (__fp16 *) args->v_tiles_dst;
536
+
537
+ struct htp_thread_trace * tr = factx->octx->ctx ? &factx->octx->ctx->trace[i] : NULL;
538
+ htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_FA_V_PREP, (uint16_t) (args->kv_start + start));
539
+ hmx_interleave_cols_to_tiles(v_tiles_dst, (const __fp16 *) args->v_src, total_rows, factx->DV,
540
+ args->src_stride, (uint32_t) args->n_col_tiles, start, end);
541
+ htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_FA_V_PREP, (uint16_t) (args->kv_start + start));
542
+ }
543
+
544
+ static void fa_phase_v_interleave(struct hmx_fa_context * factx,
545
+ uint32_t kv_rows,
546
+ size_t src_stride,
547
+ void * v_src,
548
+ void * v_tiles_dst,
549
+ size_t n_col_tiles,
550
+ uint32_t kv_start) {
551
+ worker_pool_context_t wp = factx->octx->ctx->worker_pool;
552
+ uint32_t n = 1;
553
+ if (factx->n_threads > 1 && kv_rows >= factx->n_threads * 2) {
554
+ n = factx->n_threads;
555
+ }
556
+ uint32_t rows_per_t = hex_align_up(hmx_ceil_div(kv_rows, n), 2);
557
+ fa_v_int_args_t args = { factx, kv_rows, src_stride, v_src, v_tiles_dst, n_col_tiles, kv_start, rows_per_t };
558
+ if (n > 1) {
559
+ worker_pool_run_func(wp, fa_v_interleave_thread, &args, n);
560
+ } else {
561
+ fa_v_interleave_thread(1, 0, &args);
562
+ }
563
+ }
564
+
565
+ typedef struct {
566
+ struct hmx_fa_context * factx;
567
+ const struct htp_tensor * q;
568
+ uint32_t q_start;
569
+ uint32_t kv_head;
570
+ uint32_t ib3;
571
+ size_t n_rows_g;
572
+ size_t rows_per_t;
573
+ size_t n_rows_q;
574
+ bool q_transposed;
575
+ atomic_uint barrier;
576
+ } fa_q_load_args_t;
577
+
578
+ static void fa_q_load_thread(unsigned int n, unsigned int i, void * data) {
579
+ fa_q_load_args_t * args = (fa_q_load_args_t *) data;
580
+ struct hmx_fa_context * factx = args->factx;
581
+
582
+ const size_t n_rows_g = args->n_rows_g;
583
+ const size_t G = factx->G;
584
+ const size_t DK = factx->DK;
585
+
586
+ // Partition the padded Q rows (g_br) across threads.
587
+ // Keep start/end even so r and r+1 are always in the same thread's range.
588
+ const size_t rows_per_t = args->rows_per_t;
589
+ const size_t start = (size_t) i * rows_per_t;
590
+ const size_t end = hex_smin(start + rows_per_t, factx->g_br);
591
+
592
+ struct htp_thread_trace * tr = factx->octx->ctx ? &factx->octx->ctx->trace[i] : NULL;
593
+ htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_FA_Q_PREP, (uint16_t) (args->q_start * G + start));
594
+
595
+ // Parallel initialization of per-block state
596
+ {
597
+ const uint32_t g_br = factx->g_br;
598
+ const uint32_t DV = factx->DV;
599
+
600
+ const size_t col_vec_bytes = factx->col_vec_bytes;
601
+ const size_t d_tile_bytes = factx->d_tile_bytes;
602
+
603
+ // Initialize vtcm_l_vec & vtcm_m_vec
604
+ const size_t l_bytes_per_t = hex_align_up(col_vec_bytes / n, 128);
605
+ const size_t l_start = i * l_bytes_per_t;
606
+ const size_t l_end = hex_smin(l_start + l_bytes_per_t, col_vec_bytes);
607
+
608
+ const size_t m_bytes_per_t = hex_align_up(col_vec_bytes / n, 128);
609
+ const size_t m_start = i * m_bytes_per_t;
610
+ const size_t m_end = hex_smin(m_start + m_bytes_per_t, col_vec_bytes);
611
+
612
+ if (factx->sinks) {
613
+ const float * sinks_data = (const float *) (uintptr_t) factx->sinks->data;
614
+ float * m_vec = (float *) factx->vtcm_m_vec;
615
+ const size_t r_start = l_start / sizeof(float);
616
+ const size_t r_end = l_end / sizeof(float);
617
+ const float scale_factor = EXP_LOG2E_F;
618
+
619
+ const HVX_Vector v_scale = hvx_vec_splat_f32(scale_factor);
620
+
621
+ for (size_t r = r_start; r < r_end; r += 32) {
622
+ HVX_VectorAlias local_m;
623
+ for (size_t j = 0; j < 32; ++j) {
624
+ size_t curr_r = r + j;
625
+ if (curr_r < n_rows_g) {
626
+ const size_t h_idx = fastmodulo(curr_r, G, &factx->div_G);
627
+ const size_t head = args->kv_head * G + h_idx;
628
+ local_m.fp32[j] = sinks_data[head];
629
+ } else {
630
+ local_m.fp32[j] = HTP_FA_M_INITIAL_VAL;
631
+ }
632
+ }
633
+ HVX_Vector v_scaled = HVX_OP_MUL_F32(local_m.v, v_scale);
634
+ *(HVX_Vector *) (m_vec + r) = v_scaled;
635
+ }
636
+ if (l_start < col_vec_bytes) {
637
+ hvx_splat_u8_a((char *) factx->vtcm_l_vec + l_start, 0, l_end - l_start);
638
+ }
639
+ } else {
640
+ if (l_start < col_vec_bytes) {
641
+ hvx_splat_u8_a((char *) factx->vtcm_l_vec + l_start, 0, l_end - l_start);
642
+ }
643
+ if (m_start < col_vec_bytes) {
644
+ hvx_splat_f32_a((char *) factx->vtcm_m_vec + m_start, HTP_FA_M_INITIAL_VAL, (m_end - m_start) / sizeof(float));
645
+ }
646
+ }
647
+
648
+ // Initialize vtcm_d_tiles to 0
649
+ const size_t d_bytes_per_t = hex_align_up(d_tile_bytes / n, 128);
650
+ const size_t d_start = i * d_bytes_per_t;
651
+ const size_t d_end = hex_smin(d_start + d_bytes_per_t, d_tile_bytes);
652
+ if (d_start < d_tile_bytes) {
653
+ hvx_splat_u8_a((char *) factx->vtcm_d_tiles + d_start, 0, d_end - d_start);
654
+ }
655
+ }
656
+
657
+ if (start < factx->g_br) {
658
+ const struct htp_tensor * q = args->q;
659
+ const uint32_t q_start = args->q_start;
660
+ const uint32_t kv_head = args->kv_head;
661
+ const uint32_t ib3 = args->ib3;
662
+
663
+ assert(factx->DK == factx->DV);
664
+
665
+ const size_t o_tile_bytes = factx->o_tile_bytes;
666
+ const bool use_q_dma = (2 * o_tile_bytes >= factx->g_br * DK * (factx->is_q_fp32 ? 4 : 2));
667
+
668
+ __fp16 * q_tiles = factx->vtcm_q_tiles;
669
+ if (use_q_dma) {
670
+ const size_t g_rows_end = hex_smin(end, n_rows_g);
671
+ const uint32_t d_limit = factx->is_q_fp32 ? DK / 32 : DK / 64;
672
+
673
+ uint8_t * q_flat = (uint8_t *) factx->vtcm_o_tiles[0];
674
+ if (factx->is_q_fp32) {
675
+ switch (d_limit) {
676
+ case 2: hmx_fa_q_prep_fp32_d2(q_tiles, q_flat, start, end, g_rows_end, DK, G, args->n_rows_q, &factx->div_G, args->q_transposed); break;
677
+ case 4: hmx_fa_q_prep_fp32_d4(q_tiles, q_flat, start, end, g_rows_end, DK, G, args->n_rows_q, &factx->div_G, args->q_transposed); break;
678
+ default: hmx_fa_q_prep_fp32( q_tiles, q_flat, start, end, g_rows_end, DK, G, args->n_rows_q, &factx->div_G, d_limit, args->q_transposed); break;
679
+ }
680
+ } else {
681
+ switch (d_limit) {
682
+ case 1: hmx_fa_q_prep_fp16_d1(q_tiles, q_flat, start, end, g_rows_end, DK, G, args->n_rows_q, &factx->div_G, args->q_transposed); break;
683
+ case 2: hmx_fa_q_prep_fp16_d2(q_tiles, q_flat, start, end, g_rows_end, DK, G, args->n_rows_q, &factx->div_G, args->q_transposed); break;
684
+ default: hmx_fa_q_prep_fp16( q_tiles, q_flat, start, end, g_rows_end, DK, G, args->n_rows_q, &factx->div_G, d_limit, args->q_transposed); break;
685
+ }
686
+ }
687
+ } else {
688
+ // Fallback: direct-from-DDR/L2 path
689
+ hmx_fa_q_prep_fallback(q_tiles, q->data, q->nb[1], q->nb[2], q->nb[3],
690
+ q_start, kv_head, ib3, start, end, n_rows_g, G, DK, factx->is_q_fp32, &factx->div_G);
691
+ }
692
+ }
693
+
694
+ // Synchronize threads before zeroing out vtcm_o_tiles[0] to prevent race condition
695
+ if (n > 1) {
696
+ atomic_fetch_sub(&args->barrier, 1);
697
+ while (atomic_load(&args->barrier) > 0) {
698
+ // spin wait
699
+ }
700
+ }
701
+
702
+ // Zero out vtcm_o_tiles[0] as it was used as temp_q_vtcm
703
+ {
704
+ const uint32_t g_br = factx->g_br;
705
+ const uint32_t DV = factx->DV;
706
+ const size_t o_tile_bytes = factx->o_tile_bytes;
707
+ const size_t o_bytes_per_t = hex_align_up(o_tile_bytes / n, 128);
708
+ const size_t o_start = i * o_bytes_per_t;
709
+ const size_t o_end = hex_smin(o_start + o_bytes_per_t, o_tile_bytes);
710
+ if (o_start < o_tile_bytes) {
711
+ hvx_splat_u8_a((char *) factx->vtcm_o_tiles[0] + o_start, 0, o_end - o_start);
712
+ }
713
+ }
714
+ htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_FA_Q_PREP, (uint16_t) (args->q_start * G + start));
715
+ }
716
+
717
+ static void fa_phase_q_load(struct hmx_fa_context * factx,
718
+ const struct htp_tensor * q,
719
+ uint32_t q_start,
720
+ uint32_t kv_head,
721
+ uint32_t ib3,
722
+ size_t n_rows_g) {
723
+ worker_pool_context_t wp = factx->octx->ctx->worker_pool;
724
+ uint32_t n = 1;
725
+ if (factx->n_threads > 1 && n_rows_g >= (size_t) (factx->n_threads * 2)) {
726
+ n = factx->n_threads;
727
+ }
728
+ size_t rows_per_t = hex_align_up(hmx_ceil_div(factx->g_br, n), 2);
729
+ const uint32_t n_rows_q = hex_smin(factx->Br, factx->neq1 - q_start);
730
+ fa_q_load_args_t args;
731
+ args.factx = factx;
732
+ args.q = q;
733
+ args.q_start = q_start;
734
+ args.kv_head = kv_head;
735
+ args.ib3 = ib3;
736
+ args.n_rows_g = n_rows_g;
737
+ args.rows_per_t = rows_per_t;
738
+ args.n_rows_q = n_rows_q;
739
+ args.q_transposed = q->nb[1] < q->nb[2];
740
+ atomic_init(&args.barrier, n);
741
+ if (n > 1) {
742
+ worker_pool_run_func(wp, fa_q_load_thread, &args, n);
743
+ } else {
744
+ fa_q_load_thread(1, 0, &args);
745
+ }
746
+ }
747
+
748
+ typedef struct {
749
+ struct hmx_fa_context * factx;
750
+ const struct htp_tensor * dst;
751
+ const __fp16 * o_tile_src;
752
+ uint32_t q_start;
753
+ uint32_t kv_head;
754
+ uint32_t ib3;
755
+ size_t n_rows_g;
756
+ size_t rows_per_t;
757
+ } fa_o_store_args_t;
758
+
759
+ static void fa_o_store_thread_f32(unsigned int n, unsigned int i, void * data) {
760
+ fa_o_store_args_t * args = (fa_o_store_args_t *) data;
761
+ struct hmx_fa_context * factx = args->factx;
762
+
763
+ const size_t n_rows_g = args->n_rows_g;
764
+ const size_t G = factx->G;
765
+ const size_t DV = factx->DV;
766
+
767
+ const size_t rows_per_t = args->rows_per_t;
768
+ const size_t start = (size_t) i * rows_per_t;
769
+ const size_t end = hex_smin(start + rows_per_t, n_rows_g);
770
+
771
+ if (start >= n_rows_g) {
772
+ return;
773
+ }
774
+
775
+ struct htp_thread_trace * tr = factx->octx->ctx ? &factx->octx->ctx->trace[i] : NULL;
776
+ htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_O_PROC, (uint16_t) (args->q_start * G + start));
777
+
778
+ const struct htp_tensor * dst = args->dst;
779
+ const __fp16 * o_tile_src = args->o_tile_src;
780
+ const uint32_t q_start = args->q_start;
781
+ const uint32_t kv_head = args->kv_head;
782
+ const uint32_t ib3 = args->ib3;
783
+
784
+ for (size_t r = start; r < end; ++r) {
785
+ const size_t q_idx = fastdiv(r, &factx->div_G);
786
+ const size_t h_idx = fastmodulo(r, G, &factx->div_G);
787
+
788
+ float * out = (float *) ((uint8_t *) dst->data + (kv_head * G + h_idx) * dst->nb[1] +
789
+ (q_start + q_idx) * dst->nb[2] + ib3 * dst->nb[3]);
790
+
791
+ size_t r0 = r / HMX_FP16_TILE_N_ROWS;
792
+ size_t r1 = r % HMX_FP16_TILE_N_ROWS;
793
+ const __fp16 * tile_row_base = o_tile_src + r0 * HMX_FP16_TILE_N_ROWS * DV;
794
+
795
+ for (uint32_t d = 0; d < DV / 32; ++d) {
796
+ const HVX_Vector * in_tile = (const HVX_Vector *) (tile_row_base + d * HMX_FP16_TILE_N_ELMS);
797
+ HVX_VectorPair vp = hvx_vec_f16_to_f32_shuff(in_tile[r1 / 2]);
798
+ if (r1 % 2 == 0) {
799
+ *(HVX_UVector *) (out + d * 32) = Q6_V_lo_W(vp);
800
+ } else {
801
+ *(HVX_UVector *) (out + d * 32) = Q6_V_hi_W(vp);
802
+ }
803
+ }
804
+ }
805
+ htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_O_PROC, (uint16_t) (args->q_start * G + start));
806
+ }
807
+
808
+ static void fa_o_store_thread_f16(unsigned int n, unsigned int i, void * data) {
809
+ fa_o_store_args_t * args = (fa_o_store_args_t *) data;
810
+ struct hmx_fa_context * factx = args->factx;
811
+
812
+ const size_t n_rows_g = args->n_rows_g;
813
+ const size_t rows_per_t = args->rows_per_t;
814
+ const size_t G = factx->G;
815
+ const size_t DV = factx->DV;
816
+ const size_t start = (size_t) i * rows_per_t;
817
+ const size_t end = hex_smin(start + rows_per_t, n_rows_g);
818
+
819
+ if (start >= n_rows_g) {
820
+ return;
821
+ }
822
+
823
+ struct htp_thread_trace * tr = factx->octx->ctx ? &factx->octx->ctx->trace[i] : NULL;
824
+ htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_O_PROC, (uint16_t) (args->q_start * G + start));
825
+
826
+ const struct htp_tensor * dst = args->dst;
827
+ const __fp16 * o_tile_src = args->o_tile_src;
828
+ const uint32_t q_start = args->q_start;
829
+ const uint32_t kv_head = args->kv_head;
830
+ const uint32_t ib3 = args->ib3;
831
+
832
+ for (size_t r = start; r < end; ++r) {
833
+ const size_t q_idx = fastdiv(r, &factx->div_G);
834
+ const size_t h_idx = fastmodulo(r, G, &factx->div_G);
835
+
836
+ __fp16 * out = (__fp16 *) ((uint8_t *) dst->data + (kv_head * G + h_idx) * dst->nb[1] +
837
+ (q_start + q_idx) * dst->nb[2] + ib3 * dst->nb[3]);
838
+
839
+ size_t r0 = r / HMX_FP16_TILE_N_ROWS;
840
+ size_t r1 = r % HMX_FP16_TILE_N_ROWS;
841
+ const __fp16 * tile_row_base = o_tile_src + r0 * HMX_FP16_TILE_N_ROWS * DV;
842
+
843
+ for (uint32_t d = 0; d < DV / 64; ++d) {
844
+ const __fp16 * in_dtile = tile_row_base + d * HMX_FP16_TILE_N_ELMS * 2;
845
+ const HVX_Vector * pv_in0 = ((const HVX_Vector *) in_dtile) + r1 / 2;
846
+ const HVX_Vector * pv_in1 = pv_in0 + 16;
847
+ HVX_VectorPair vp = Q6_W_vdeal_VVR(*pv_in1, *pv_in0, -2);
848
+ if (r1 % 2 == 0) {
849
+ *(HVX_UVector *) (out + d * 64) = Q6_V_lo_W(vp);
850
+ } else {
851
+ *(HVX_UVector *) (out + d * 64) = Q6_V_hi_W(vp);
852
+ }
853
+ }
854
+ }
855
+ htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_O_PROC, (uint16_t) (args->q_start * G + start));
856
+ }
857
+
858
+ static void fa_phase_o_store(struct hmx_fa_context * factx,
859
+ const struct htp_tensor * dst,
860
+ const __fp16 * o_tile_src,
861
+ uint32_t q_start,
862
+ uint32_t kv_head,
863
+ uint32_t ib3,
864
+ size_t n_rows_g) {
865
+ worker_pool_context_t wp = factx->octx->ctx->worker_pool;
866
+ uint32_t n = 1;
867
+ if (factx->n_threads > 1 && n_rows_g >= (size_t) (factx->n_threads * 2)) {
868
+ n = factx->n_threads;
869
+ }
870
+ size_t rows_per_t = hmx_ceil_div(n_rows_g, n);
871
+ fa_o_store_args_t args = { factx, dst, o_tile_src, q_start, kv_head, ib3, n_rows_g, rows_per_t };
872
+ worker_callback_t store_fn = factx->is_dst_fp32 ? fa_o_store_thread_f32 : fa_o_store_thread_f16;
873
+ if (n > 1) {
874
+ worker_pool_run_func(wp, store_fn, &args, n);
875
+ } else {
876
+ store_fn(1, 0, &args);
877
+ }
878
+ }
879
+
880
+ typedef struct {
881
+ struct hmx_fa_context * factx;
882
+ size_t kv_rows;
883
+ size_t n_rows_g;
884
+ size_t n_col_tiles;
885
+ size_t n_tiles_per_bc;
886
+ size_t n_row_tiles;
887
+ size_t n_row_tiles_g_br;
888
+ uint32_t Bc;
889
+ uint32_t G;
890
+ uint32_t kv_head;
891
+ uint32_t kv_start;
892
+ uint32_t q_start;
893
+ uint32_t ib3;
894
+ bool has_alibi; // true when max_bias != 0 (need slope * mask + add)
895
+ __fp16 * slopes;
896
+ const struct htp_tensor * mask;
897
+ const __fp16 * mask_vtcm; // VTCM mask buffer base (NULL = DDR fallback)
898
+ size_t mask_vtcm_row_stride; // elements (__fp16) per row in VTCM mask buffer
899
+ struct fastdiv_values thread_div;
900
+ } fa_softmax_args_t;
901
+
902
+ static inline void fa_softmax_impl(
903
+ unsigned int n, unsigned int i, void * data,
904
+ const bool has_mask,
905
+ const bool mask_broadcast,
906
+ const bool is_g1,
907
+ const bool has_alibi,
908
+ const bool has_softcap
909
+ ) {
910
+ fa_softmax_args_t * args = (fa_softmax_args_t *) data;
911
+ struct hmx_fa_context * factx = args->factx;
912
+
913
+ const size_t n_rows_g = args->n_rows_g;
914
+ const size_t kv_rows = args->kv_rows;
915
+ const size_t Bc = args->Bc;
916
+ const size_t G = args->G;
917
+ const size_t n_tiles_per_bc = args->n_tiles_per_bc;
918
+ const size_t n_row_vec_cnt = hmx_ceil_div(n_rows_g, 64);
919
+ const uint32_t im3 = has_mask ? fastmodulo(args->ib3, args->mask->ne[3], &factx->src3_div3) : 0;
920
+
921
+ size_t vec_start = 0;
922
+ size_t vec_end = n_row_vec_cnt;
923
+ if (n > 1) {
924
+ const size_t vecs_per_t = fastdiv(n_row_vec_cnt + n - 1, &args->thread_div);
925
+ vec_start = i * vecs_per_t;
926
+ vec_end = hex_smin(vec_start + vecs_per_t, n_row_vec_cnt);
927
+ }
928
+
929
+ if (vec_start >= n_row_vec_cnt) {
930
+ return;
931
+ }
932
+
933
+ struct htp_thread_trace * tr = factx->octx->ctx ? &factx->octx->ctx->trace[i] : NULL;
934
+ htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_FA_SFM, (uint16_t) (args->q_start * G + vec_start * 64));
935
+
936
+ // Per-thread row scratch: thread i uses bufs at offset i * 2 * stride
937
+ const size_t row_buf_stride = factx->row_buf_stride;
938
+ HVX_Vector * my_row_buf0 = factx->vtcm_row_bufs + i * 2 * row_buf_stride;
939
+ HVX_Vector * my_row_buf1 = my_row_buf0 + row_buf_stride;
940
+
941
+ const HVX_Vector v_neg_inf = Q6_Vh_vsplat_R(0xfbff);
942
+
943
+ for (size_t r_vec_idx = vec_start; r_vec_idx < vec_end; ++r_vec_idx) {
944
+ HVX_Vector rowmax_acc_v = v_neg_inf;
945
+ HVX_Vector rowsum_acc_v = Q6_V_vzero();
946
+ HVX_Vector m_prev_v0 = factx->vtcm_m_vec[r_vec_idx * 2 + 0];
947
+ HVX_Vector m_prev_v1 = factx->vtcm_m_vec[r_vec_idx * 2 + 1];
948
+
949
+ HVX_Vector v_slopes = Q6_V_vzero();
950
+ if (has_alibi) {
951
+ v_slopes = hvx_vmem(args->slopes + r_vec_idx * 64);
952
+ }
953
+
954
+ for (uint32_t r_vec_off = 0; r_vec_off < 64; r_vec_off += 2) {
955
+ uint32_t r = r_vec_idx * 64 + r_vec_off;
956
+ if (r >= hex_align_up(n_rows_g, 2)) {
957
+ break;
958
+ }
959
+
960
+ uint32_t r0 = r / HMX_FP16_TILE_N_ROWS;
961
+ uint32_t r1 = r % HMX_FP16_TILE_N_ROWS;
962
+
963
+ const __fp16 * s_ld_base = factx->vtcm_s_tiles + r0 * HMX_FP16_TILE_N_ROWS * Bc;
964
+ __fp16 * p_st_base = factx->vtcm_p_tiles + r0 * HMX_FP16_TILE_N_ROWS * Bc;
965
+
966
+ // Decode 2 rows from S tiles into per-thread row buffers
967
+ if (has_softcap) {
968
+ const HVX_Vector v_cap = hvx_vec_splat_f16(factx->logit_softcap);
969
+ for (size_t c = 0; c < kv_rows; c += 64) {
970
+ size_t ci = c / 64;
971
+ const __fp16 * in_dtile = s_ld_base + ci * HMX_FP16_TILE_N_ELMS * 2;
972
+ const HVX_Vector * pv_s_in0 = ((const HVX_Vector *) in_dtile) + r1 / 2;
973
+ const HVX_Vector * pv_s_in1 = pv_s_in0 + 16;
974
+
975
+ HVX_VectorPair vp_s_drow = Q6_W_vdeal_VVR(*pv_s_in1, *pv_s_in0, -2);
976
+ HVX_Vector v_s_row0 = Q6_V_lo_W(vp_s_drow);
977
+ HVX_Vector v_s_row1 = Q6_V_hi_W(vp_s_drow);
978
+
979
+ HVX_Vector t0 = hvx_vec_tanh_f16(v_s_row0);
980
+ my_row_buf0[ci] = hvx_vec_mul_f16_f16(t0, v_cap);
981
+
982
+ HVX_Vector t1 = hvx_vec_tanh_f16(v_s_row1);
983
+ my_row_buf1[ci] = hvx_vec_mul_f16_f16(t1, v_cap);
984
+ }
985
+ } else {
986
+ for (size_t c = 0; c < kv_rows; c += 64) {
987
+ size_t ci = c / 64;
988
+ const __fp16 * in_dtile = s_ld_base + ci * HMX_FP16_TILE_N_ELMS * 2;
989
+ const HVX_Vector * pv_s_in0 = ((const HVX_Vector *) in_dtile) + r1 / 2;
990
+ const HVX_Vector * pv_s_in1 = pv_s_in0 + 16;
991
+
992
+ HVX_VectorPair vp_s_drow = Q6_W_vdeal_VVR(*pv_s_in1, *pv_s_in0, -2);
993
+ my_row_buf0[ci] = Q6_V_lo_W(vp_s_drow);
994
+ my_row_buf1[ci] = Q6_V_hi_W(vp_s_drow);
995
+ }
996
+ }
997
+
998
+ // Apply mask & compute rowmax(S)
999
+ HVX_Vector v_slope0 = Q6_V_vzero();
1000
+ HVX_Vector v_slope1 = Q6_V_vzero();
1001
+ if (has_alibi) {
1002
+ v_slope0 = hvx_vec_repl_f16(Q6_V_vror_VR(v_slopes, r_vec_off * 2));
1003
+ v_slope1 = (r + 1 < n_rows_g) ? hvx_vec_repl_f16(Q6_V_vror_VR(v_slopes, (r_vec_off + 1) * 2)) : Q6_V_vzero();
1004
+ }
1005
+
1006
+ const HVX_Vector v_threshold = Q6_Vh_vsplat_R(0xcc00); // fp16 -16.0
1007
+
1008
+ HVX_Vector v_s_rowmax0 = v_neg_inf;
1009
+ HVX_Vector v_s_rowmax1 = v_neg_inf;
1010
+ for (size_t c = 0; c < kv_rows; c += 64) {
1011
+ size_t ci = c / 64;
1012
+ const size_t ne = hex_smin(kv_rows - c, 64);
1013
+ HVX_VectorPred q_tail_keep = Q6_Q_vsetq2_R(ne * sizeof(__fp16));
1014
+
1015
+ if (has_mask) {
1016
+ HVX_Vector v_mask0, v_mask1;
1017
+
1018
+ if (mask_broadcast) {
1019
+ if (is_g1) {
1020
+ const size_t qi0 = r + 0;
1021
+ v_mask0 = *(const HVX_Vector *) (args->mask_vtcm + qi0 * args->mask_vtcm_row_stride + c);
1022
+ v_mask1 = v_neg_inf;
1023
+ if (r + 1 < n_rows_g) {
1024
+ const size_t qi1 = r + 1;
1025
+ v_mask1 = *(const HVX_Vector *) (args->mask_vtcm + qi1 * args->mask_vtcm_row_stride + c);
1026
+ }
1027
+ } else {
1028
+ const size_t qi0 = fastdiv(r + 0, &factx->div_G);
1029
+ v_mask0 = *(const HVX_Vector *) (args->mask_vtcm + qi0 * args->mask_vtcm_row_stride + c);
1030
+ v_mask1 = v_neg_inf;
1031
+ if (r + 1 < n_rows_g) {
1032
+ const size_t qi1 = fastdiv(r + 1, &factx->div_G);
1033
+ if (qi1 == qi0) {
1034
+ v_mask1 = v_mask0;
1035
+ } else {
1036
+ v_mask1 = *(const HVX_Vector *) (args->mask_vtcm + qi1 * args->mask_vtcm_row_stride + c);
1037
+ }
1038
+ }
1039
+ }
1040
+ } else {
1041
+ // Head-dependent mask: pre-interleaved per row r.
1042
+ const size_t r0 = r + 0;
1043
+ v_mask0 = *(const HVX_Vector *) (args->mask_vtcm + r0 * args->mask_vtcm_row_stride + c);
1044
+ v_mask1 = v_neg_inf;
1045
+ if (r + 1 < n_rows_g) {
1046
+ const size_t r1 = r + 1;
1047
+ v_mask1 = *(const HVX_Vector *) (args->mask_vtcm + r1 * args->mask_vtcm_row_stride + c);
1048
+ }
1049
+ }
1050
+
1051
+ // Threshold: mask values below -16.0 are treated as -inf (causal mask).
1052
+ HVX_VectorPred q_keep0 = Q6_Q_and_QQ(Q6_Q_vcmp_gt_VhfVhf(v_mask0, v_threshold), q_tail_keep);
1053
+ HVX_VectorPred q_keep1 = Q6_Q_and_QQ(Q6_Q_vcmp_gt_VhfVhf(v_mask1, v_threshold), q_tail_keep);
1054
+
1055
+ // Scale mask values by log2(e) for base-2 calculations
1056
+ const HVX_Vector v_log2e = hvx_vec_splat_f16(EXP_LOG2E_F);
1057
+ HVX_Vector v_mask0_scaled = hvx_vec_mul_f16_f16(v_mask0, v_log2e);
1058
+ HVX_Vector v_mask1_scaled = hvx_vec_mul_f16_f16(v_mask1, v_log2e);
1059
+
1060
+ if (has_alibi) {
1061
+ HVX_Vector v_sm0 = hvx_vec_mul_f16_f16(v_mask0_scaled, v_slope0);
1062
+ HVX_Vector v_sm1 = hvx_vec_mul_f16_f16(v_mask1_scaled, v_slope1);
1063
+ my_row_buf0[ci] = Q6_V_vmux_QVV(q_keep0, hvx_vec_add_f16_f16(my_row_buf0[ci], v_sm0), v_neg_inf);
1064
+ my_row_buf1[ci] = Q6_V_vmux_QVV(q_keep1, hvx_vec_add_f16_f16(my_row_buf1[ci], v_sm1), v_neg_inf);
1065
+ } else {
1066
+ my_row_buf0[ci] = Q6_V_vmux_QVV(q_keep0, hvx_vec_add_f16_f16(my_row_buf0[ci], v_mask0_scaled), v_neg_inf);
1067
+ my_row_buf1[ci] = Q6_V_vmux_QVV(q_keep1, hvx_vec_add_f16_f16(my_row_buf1[ci], v_mask1_scaled), v_neg_inf);
1068
+ }
1069
+ } else {
1070
+ if (ne < 64) {
1071
+ my_row_buf0[ci] = Q6_V_vmux_QVV(q_tail_keep, my_row_buf0[ci], v_neg_inf);
1072
+ my_row_buf1[ci] = Q6_V_vmux_QVV(q_tail_keep, my_row_buf1[ci], v_neg_inf);
1073
+ }
1074
+ }
1075
+
1076
+ v_s_rowmax0 = Q6_Vhf_vmax_VhfVhf(v_s_rowmax0, my_row_buf0[ci]);
1077
+ v_s_rowmax1 = Q6_Vhf_vmax_VhfVhf(v_s_rowmax1, my_row_buf1[ci]);
1078
+ }
1079
+
1080
+ v_s_rowmax0 = hvx_vec_reduce_max_f16(v_s_rowmax0);
1081
+ v_s_rowmax1 = hvx_vec_reduce_max_f16(v_s_rowmax1);
1082
+
1083
+ // Splat m_prev[r], m_prev[r+1] from the float per-row accumulators and convert to fp16 vectors
1084
+ HVX_Vector v_m_prev0, v_m_prev1;
1085
+ if (r_vec_off < 32) {
1086
+ HVX_Vector v0 = hvx_vec_repl_f32(Q6_V_vror_VR(m_prev_v0, r_vec_off * 4));
1087
+ v_m_prev0 = hvx_vec_f32_to_f16(v0, v0);
1088
+ if (r + 1 < n_rows_g) {
1089
+ HVX_Vector v1 = hvx_vec_repl_f32(Q6_V_vror_VR(m_prev_v0, (r_vec_off + 1) * 4));
1090
+ v_m_prev1 = hvx_vec_f32_to_f16(v1, v1);
1091
+ } else {
1092
+ v_m_prev1 = Q6_V_vzero();
1093
+ }
1094
+ } else {
1095
+ HVX_Vector v0 = hvx_vec_repl_f32(Q6_V_vror_VR(m_prev_v1, (r_vec_off - 32) * 4));
1096
+ v_m_prev0 = hvx_vec_f32_to_f16(v0, v0);
1097
+ if (r + 1 < n_rows_g) {
1098
+ HVX_Vector v1 = hvx_vec_repl_f32(Q6_V_vror_VR(m_prev_v1, (r_vec_off + 1 - 32) * 4));
1099
+ v_m_prev1 = hvx_vec_f32_to_f16(v1, v1);
1100
+ } else {
1101
+ v_m_prev1 = Q6_V_vzero();
1102
+ }
1103
+ }
1104
+
1105
+ HVX_Vector v_dup_m0 = Q6_Vhf_vmax_VhfVhf(v_m_prev0, v_s_rowmax0);
1106
+ HVX_Vector v_dup_m1 = Q6_Vhf_vmax_VhfVhf(v_m_prev1, v_s_rowmax1);
1107
+
1108
+ // Insert row r, r+1 rowmax into rowmax_acc_v
1109
+ {
1110
+ HVX_VectorPred p_start = Q6_Q_vsetq_R(r_vec_off * 2);
1111
+ HVX_VectorPred p_mid = Q6_Q_vsetq_R((r_vec_off + 1) * 2);
1112
+ HVX_VectorPred p_end = Q6_Q_vsetq2_R((r_vec_off + 2) * 2);
1113
+ HVX_VectorPred p_lane0 = Q6_Q_and_QQn(p_mid, p_start);
1114
+ HVX_VectorPred p_lane1 = Q6_Q_and_QQn(p_end, p_mid);
1115
+ rowmax_acc_v = Q6_V_vmux_QVV(p_lane0, v_dup_m0, rowmax_acc_v);
1116
+ rowmax_acc_v = Q6_V_vmux_QVV(p_lane1, v_dup_m1, rowmax_acc_v);
1117
+ }
1118
+
1119
+ // Compute P = exp(S - m_new)
1120
+ const HVX_Vector v_zero = Q6_V_vzero();
1121
+ HVX_Vector v_p_rowsum0 = v_zero;
1122
+ HVX_Vector v_p_rowsum1 = v_zero;
1123
+
1124
+ for (size_t c = 0; c < kv_rows; c += 64) {
1125
+ size_t ci = c / 64;
1126
+ HVX_Vector v_s_minus_m0 = Q6_Vqf16_vsub_VhfVhf(my_row_buf0[ci], v_dup_m0);
1127
+ HVX_Vector v_s_minus_m1 = Q6_Vqf16_vsub_VhfVhf(my_row_buf1[ci], v_dup_m1);
1128
+
1129
+ HVX_Vector v_p_row0_hf = hvx_vec_exp2_f16(Q6_Vhf_equals_Vqf16(v_s_minus_m0));
1130
+ HVX_Vector v_p_row1_hf = hvx_vec_exp2_f16(Q6_Vhf_equals_Vqf16(v_s_minus_m1));
1131
+ __fp16 * out_dtile = p_st_base + ci * HMX_FP16_TILE_N_ELMS * 2;
1132
+ HVX_Vector * pv_p_out0 = ((HVX_Vector *) out_dtile) + r1 / 2;
1133
+ HVX_Vector * pv_p_out1 = pv_p_out0 + 16;
1134
+
1135
+ HVX_VectorPair vp_p_dual = Q6_W_vshuff_VVR(v_p_row1_hf, v_p_row0_hf, -2);
1136
+ *pv_p_out0 = Q6_V_lo_W(vp_p_dual);
1137
+ *pv_p_out1 = Q6_V_hi_W(vp_p_dual);
1138
+
1139
+ HVX_VectorPair vp_p0 = hvx_vec_f16_to_f32_shuff(v_p_row0_hf);
1140
+ HVX_VectorPair vp_p1 = hvx_vec_f16_to_f32_shuff(v_p_row1_hf);
1141
+
1142
+ v_p_rowsum0 = Q6_Vqf32_vadd_Vqf32Vqf32(v_p_rowsum0, Q6_Vqf32_vadd_VsfVsf(Q6_V_lo_W(vp_p0), Q6_V_hi_W(vp_p0)));
1143
+ v_p_rowsum1 = Q6_Vqf32_vadd_Vqf32Vqf32(v_p_rowsum1, Q6_Vqf32_vadd_VsfVsf(Q6_V_lo_W(vp_p1), Q6_V_hi_W(vp_p1)));
1144
+ }
1145
+
1146
+ HVX_Vector rowsum0_sf = hvx_vec_reduce_sum_f32(Q6_Vsf_equals_Vqf32(v_p_rowsum0));
1147
+ HVX_Vector rowsum1_sf = hvx_vec_reduce_sum_f32(Q6_Vsf_equals_Vqf32(v_p_rowsum1));
1148
+ {
1149
+ HVX_Vector rv0_v = hvx_vec_f32_to_f16(rowsum0_sf, rowsum0_sf);
1150
+ HVX_Vector rv1_v = hvx_vec_f32_to_f16(rowsum1_sf, rowsum1_sf);
1151
+
1152
+ HVX_VectorPred p_start = Q6_Q_vsetq_R(r_vec_off * 2);
1153
+ HVX_VectorPred p_mid = Q6_Q_vsetq_R((r_vec_off + 1) * 2);
1154
+ HVX_VectorPred p_end = Q6_Q_vsetq2_R((r_vec_off + 2) * 2);
1155
+ HVX_VectorPred p_lane0 = Q6_Q_and_QQn(p_mid, p_start);
1156
+ HVX_VectorPred p_lane1 = Q6_Q_and_QQn(p_end, p_mid);
1157
+ rowsum_acc_v = Q6_V_vmux_QVV(p_lane0, rv0_v, rowsum_acc_v);
1158
+ rowsum_acc_v = Q6_V_vmux_QVV(p_lane1, rv1_v, rowsum_acc_v);
1159
+ }
1160
+ }
1161
+
1162
+ // Inline fa_ml_update_and_build_d for this vector (lock-free and in parallel)
1163
+ HVX_VectorPair rowmax_acc_pair = hvx_vec_f16_to_f32(rowmax_acc_v);
1164
+ HVX_Vector v_rowmax_acc_f32_0 = Q6_V_lo_W(rowmax_acc_pair);
1165
+ HVX_Vector v_rowmax_acc_f32_1 = Q6_V_hi_W(rowmax_acc_pair);
1166
+
1167
+ HVX_Vector v_m_curr0 = Q6_Vsf_vmax_VsfVsf(m_prev_v0, v_rowmax_acc_f32_0);
1168
+ HVX_Vector v_m_curr1 = Q6_Vsf_vmax_VsfVsf(m_prev_v1, v_rowmax_acc_f32_1);
1169
+
1170
+ HVX_Vector v_m_diff0 = HVX_OP_SUB_F32(m_prev_v0, v_m_curr0);
1171
+ HVX_Vector v_m_diff1 = HVX_OP_SUB_F32(m_prev_v1, v_m_curr1);
1172
+
1173
+ HVX_Vector v_m_diff_f16 = hvx_vec_f32_to_f16(v_m_diff0, v_m_diff1);
1174
+ HVX_Vector exp_m_diff_f16 = hvx_vec_exp2_f16(v_m_diff_f16);
1175
+
1176
+ HVX_VectorPair exp_m_diff_pair = hvx_vec_f16_to_f32(exp_m_diff_f16);
1177
+ HVX_Vector exp_m_diff0 = Q6_V_lo_W(exp_m_diff_pair);
1178
+ HVX_Vector exp_m_diff1 = Q6_V_hi_W(exp_m_diff_pair);
1179
+
1180
+ HVX_VectorPair rowsum_acc_pair = hvx_vec_f16_to_f32(rowsum_acc_v);
1181
+ HVX_Vector v_rowsum_acc_f32_0 = Q6_V_lo_W(rowsum_acc_pair);
1182
+ HVX_Vector v_rowsum_acc_f32_1 = Q6_V_hi_W(rowsum_acc_pair);
1183
+
1184
+ HVX_Vector v_l_curr0;
1185
+ HVX_Vector v_l_curr1;
1186
+ if (args->kv_start == 0 && factx->sinks != NULL) {
1187
+ // First KV block with sinks: m_prev holds the seeded sink value (not -inf),
1188
+ // so exp_m_diff = exp2(sink - m_curr) is the sink's contribution to the
1189
+ // denominator. l_prev is 0 here, so add exp_m_diff directly instead of
1190
+ // multiplying the (uninitialized) l_prev term.
1191
+ v_l_curr0 = HVX_OP_ADD_F32(exp_m_diff0, v_rowsum_acc_f32_0);
1192
+ v_l_curr1 = HVX_OP_ADD_F32(exp_m_diff1, v_rowsum_acc_f32_1);
1193
+ } else {
1194
+ HVX_Vector l_prev_v0 = factx->vtcm_l_vec[r_vec_idx * 2 + 0];
1195
+ HVX_Vector l_prev_v1 = factx->vtcm_l_vec[r_vec_idx * 2 + 1];
1196
+ v_l_curr0 = HVX_OP_ADD_F32(HVX_OP_MUL_F32(l_prev_v0, exp_m_diff0), v_rowsum_acc_f32_0);
1197
+ v_l_curr1 = HVX_OP_ADD_F32(HVX_OP_MUL_F32(l_prev_v1, exp_m_diff1), v_rowsum_acc_f32_1);
1198
+ }
1199
+
1200
+ factx->vtcm_m_vec[r_vec_idx * 2 + 0] = v_m_curr0;
1201
+ factx->vtcm_m_vec[r_vec_idx * 2 + 1] = v_m_curr1;
1202
+ factx->vtcm_l_vec[r_vec_idx * 2 + 0] = v_l_curr0;
1203
+ factx->vtcm_l_vec[r_vec_idx * 2 + 1] = v_l_curr1;
1204
+
1205
+ // Build diagonal tile D = diag(exp(m_diff))
1206
+ const HVX_Vector v_offsets = *(const HVX_Vector *) d_tile_scatter_offsets;
1207
+ const HVX_VectorPred q_32_mask = Q6_Q_vsetq_R(32 * sizeof(__fp16));
1208
+ HVX_Vector v_exp_m_diff = exp_m_diff_f16;
1209
+
1210
+ size_t t0 = r_vec_idx * 2;
1211
+ if (t0 < args->n_row_tiles) {
1212
+ const HVX_Vector v_content = v_exp_m_diff;
1213
+ __fp16 * out_base = factx->vtcm_d_tiles + t0 * (args->n_row_tiles_g_br + 1) * HMX_FP16_TILE_N_ELMS;
1214
+ Q6_vscatter_QRMVhV(q_32_mask, (size_t) out_base, HMX_FP16_TILE_SIZE - 1, v_offsets, v_content);
1215
+ }
1216
+
1217
+ size_t t1 = r_vec_idx * 2 + 1;
1218
+ if (t1 < args->n_row_tiles) {
1219
+ const HVX_Vector v_content = Q6_V_vror_VR(v_exp_m_diff, 64);
1220
+ __fp16 * out_base = factx->vtcm_d_tiles + t1 * (args->n_row_tiles_g_br + 1) * HMX_FP16_TILE_N_ELMS;
1221
+ Q6_vscatter_QRMVhV(q_32_mask, (size_t) out_base, HMX_FP16_TILE_SIZE - 1, v_offsets, v_content);
1222
+ }
1223
+ }
1224
+ htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_FA_SFM, (uint16_t) (args->q_start * G + vec_start * 64));
1225
+ }
1226
+
1227
+ static void fa_softmax_thread_nomask(unsigned int n, unsigned int i, void * data) {
1228
+ fa_softmax_impl(n, i, data,
1229
+ /*has_mask=*/false,
1230
+ /*mask_broadcast=*/false,
1231
+ /*is_g1=*/false,
1232
+ /*has_alibi=*/false,
1233
+ /*has_softcap=*/false);
1234
+ }
1235
+
1236
+ static void fa_softmax_thread_mask_broadcast_g1(unsigned int n, unsigned int i, void * data) {
1237
+ fa_softmax_impl(n, i, data,
1238
+ /*has_mask=*/true,
1239
+ /*mask_broadcast=*/true,
1240
+ /*is_g1=*/true,
1241
+ /*has_alibi=*/false,
1242
+ /*has_softcap=*/false);
1243
+ }
1244
+
1245
+ static void fa_softmax_thread_mask_broadcast_gn(unsigned int n, unsigned int i, void * data) {
1246
+ fa_softmax_impl(n, i, data,
1247
+ /*has_mask=*/true,
1248
+ /*mask_broadcast=*/true,
1249
+ /*is_g1=*/false,
1250
+ /*has_alibi=*/false,
1251
+ /*has_softcap=*/false);
1252
+ }
1253
+
1254
+ static void fa_softmax_thread(unsigned int n, unsigned int i, void * data) {
1255
+ fa_softmax_args_t * args = (fa_softmax_args_t *) data;
1256
+ struct hmx_fa_context * factx = args->factx;
1257
+
1258
+ const bool has_mask = (args->mask != NULL);
1259
+ const bool mask_broadcast = factx->mask_broadcast;
1260
+ const bool is_g1 = (args->G == 1);
1261
+ const bool has_alibi = args->has_alibi;
1262
+ const bool has_softcap = (factx->logit_softcap != 0.0f);
1263
+
1264
+ fa_softmax_impl(n, i, data, has_mask, mask_broadcast, is_g1, has_alibi, has_softcap);
1265
+ }
1266
+
1267
+ static __attribute__((noinline)) void fa_build_d_diag_inv_l(struct hmx_fa_context * factx,
1268
+ size_t n_row_tiles,
1269
+ size_t n_row_tiles_g_br) {
1270
+ const HVX_Vector v_offsets = *(const HVX_Vector *) d_tile_scatter_offsets;
1271
+ const HVX_VectorPred q_32_mask = Q6_Q_vsetq_R(32 * sizeof(__fp16));
1272
+ const HVX_Vector one = hvx_vec_splat_f32(1.0f);
1273
+
1274
+ HVX_Vector v_content = Q6_V_vzero();
1275
+ for (size_t i = 0; i < n_row_tiles; ++i) {
1276
+ if ((i % 2) == 0) {
1277
+ HVX_Vector inv_lo = HVX_OP_MUL_F32(one, hvx_vec_inverse_f32(factx->vtcm_l_vec[i]));
1278
+ HVX_Vector inv_hi = (i + 1 < n_row_tiles) ? HVX_OP_MUL_F32(one, hvx_vec_inverse_f32(factx->vtcm_l_vec[i + 1])) : Q6_V_vzero();
1279
+ v_content = hvx_vec_f32_to_f16(inv_lo, inv_hi);
1280
+ } else {
1281
+ v_content = Q6_V_vror_VR(v_content, 64);
1282
+ }
1283
+
1284
+ __fp16 * out_base = factx->vtcm_d_tiles + i * (n_row_tiles_g_br + 1) * HMX_FP16_TILE_N_ELMS;
1285
+ Q6_vscatter_QRMVhV(q_32_mask, (size_t) out_base, HMX_FP16_TILE_SIZE - 1, v_offsets, v_content);
1286
+ }
1287
+ }
1288
+
1289
+ static void fa_phase_softmax_and_build_d(struct hmx_fa_context * factx,
1290
+ fa_softmax_args_t * sargs,
1291
+ size_t n_row_tiles,
1292
+ size_t n_row_tiles_g_br) {
1293
+ worker_pool_context_t wp = factx->octx->ctx->worker_pool;
1294
+ const size_t n_row_vec_cnt = hmx_ceil_div(sargs->n_rows_g, 64);
1295
+
1296
+ worker_callback_t softmax_fn = fa_softmax_thread;
1297
+ if (sargs->mask == NULL && factx->logit_softcap == 0.0f && !sargs->has_alibi) {
1298
+ softmax_fn = fa_softmax_thread_nomask;
1299
+ } else if (sargs->mask != NULL && factx->mask_broadcast && factx->logit_softcap == 0.0f && !sargs->has_alibi) {
1300
+ if (sargs->G == 1) {
1301
+ softmax_fn = fa_softmax_thread_mask_broadcast_g1;
1302
+ } else {
1303
+ softmax_fn = fa_softmax_thread_mask_broadcast_gn;
1304
+ }
1305
+ }
1306
+
1307
+ if (factx->n_threads > 1 && n_row_vec_cnt >= 2) {
1308
+ uint32_t n_use = (uint32_t) hex_smin((size_t) factx->n_threads, n_row_vec_cnt);
1309
+ sargs->thread_div = init_fastdiv_values(n_use);
1310
+ worker_pool_run_func(wp, softmax_fn, sargs, n_use);
1311
+ } else {
1312
+ softmax_fn(1, 0, sargs);
1313
+ }
1314
+ }
1315
+
1316
+ // ============================================================================
1317
+ // HMX job structs and worker functions
1318
+ // ============================================================================
1319
+
1320
+ typedef struct {
1321
+ const __fp16 * q_tiles;
1322
+ const __fp16 * k_tiles;
1323
+ __fp16 * s_tiles;
1324
+ size_t n_row_tiles;
1325
+ size_t n_col_tiles;
1326
+ size_t n_dot_tiles; // DK / 32
1327
+ size_t n_tiles_per_bc;
1328
+ uint8_t * hmx_scales;
1329
+ } hmx_fa_qk_job_t;
1330
+
1331
+ static void hmx_fa_qk_dot_worker(void * data) {
1332
+ hmx_fa_qk_job_t * job = (hmx_fa_qk_job_t *) data;
1333
+ const size_t n_row_tiles = job->n_row_tiles;
1334
+ const size_t n_col_tiles = job->n_col_tiles;
1335
+ const size_t n_dot_tiles = job->n_dot_tiles;
1336
+ const size_t n_tiles_per_bc = job->n_tiles_per_bc;
1337
+ const __fp16 * restrict q_tiles = job->q_tiles;
1338
+ const __fp16 * restrict k_tiles = job->k_tiles;
1339
+ __fp16 * restrict s_tiles = job->s_tiles;
1340
+ __builtin_assume(n_row_tiles > 0);
1341
+ __builtin_assume(n_col_tiles > 0);
1342
+ __builtin_assume(n_dot_tiles > 0);
1343
+
1344
+ asm volatile(HMX_SET_BIAS("%0") :: "r"((unsigned int)job->hmx_scales));
1345
+ const size_t dot_stride = n_dot_tiles * HMX_FP16_TILE_N_ELMS;
1346
+ for (size_t r = 0; r < n_row_tiles; ++r) {
1347
+ const __fp16 * row_tiles = q_tiles + r * dot_stride;
1348
+ const __fp16 * col_tiles = k_tiles;
1349
+ __fp16 * out_tile = s_tiles + r * n_tiles_per_bc * HMX_FP16_TILE_N_ELMS;
1350
+
1351
+ for (size_t c = 0; c < n_col_tiles; ++c) {
1352
+ hmx_fa_qk_dot_tile(row_tiles, col_tiles, out_tile, n_dot_tiles);
1353
+ col_tiles += dot_stride;
1354
+ out_tile += HMX_FP16_TILE_N_ELMS;
1355
+ }
1356
+ }
1357
+ }
1358
+
1359
+ typedef struct {
1360
+ __fp16 * o_curr;
1361
+ const __fp16 * o_prev;
1362
+ const __fp16 * p_tiles;
1363
+ const __fp16 * v_tiles;
1364
+ const __fp16 * d_tiles;
1365
+ uint8_t * hmx_scales;
1366
+ size_t n_row_tiles;
1367
+ size_t n_col_tiles;
1368
+ size_t n_row_tiles_g_br;
1369
+ size_t n_tiles_per_bc;
1370
+ size_t DV;
1371
+ } hmx_fa_o_update_job_t;
1372
+
1373
+ static void hmx_fa_o_update_worker(void * data) {
1374
+ hmx_fa_o_update_job_t * job = (hmx_fa_o_update_job_t *) data;
1375
+ const size_t n_row_tiles = job->n_row_tiles;
1376
+ const size_t n_col_tiles = job->n_col_tiles;
1377
+ const size_t n_row_tiles_g_br = job->n_row_tiles_g_br;
1378
+ const size_t n_tiles_per_bc = job->n_tiles_per_bc;
1379
+ const size_t DV_tiles = job->DV / 32;
1380
+ const __fp16 * restrict d_tiles = job->d_tiles;
1381
+ const __fp16 * restrict p_tiles = job->p_tiles;
1382
+ const __fp16 * restrict v_tiles = job->v_tiles;
1383
+ const __fp16 * restrict o_prev = job->o_prev;
1384
+ __fp16 * restrict o_curr = job->o_curr;
1385
+ __builtin_assume(n_row_tiles > 0);
1386
+ __builtin_assume(n_col_tiles > 0);
1387
+ __builtin_assume(DV_tiles > 0);
1388
+
1389
+ asm volatile(HMX_SET_BIAS("%0") :: "r"((unsigned int)job->hmx_scales));
1390
+ const size_t o_stride = n_row_tiles_g_br * HMX_FP16_TILE_N_ELMS;
1391
+ const size_t v_stride = n_tiles_per_bc * HMX_FP16_TILE_N_ELMS;
1392
+ for (size_t r = 0; r < n_row_tiles; ++r) {
1393
+ const __fp16 * d_diag = d_tiles + r * (n_row_tiles_g_br + 1) * HMX_FP16_TILE_N_ELMS;
1394
+ const __fp16 * p_tile_in = p_tiles + (r * n_tiles_per_bc) * HMX_FP16_TILE_N_ELMS;
1395
+ const __fp16 * o_rc = o_prev + r * HMX_FP16_TILE_N_ELMS;
1396
+ const __fp16 * v_tile_in = v_tiles;
1397
+ __fp16 * o_tile_out = o_curr + r * HMX_FP16_TILE_N_ELMS;
1398
+
1399
+ for (size_t c = 0; c < DV_tiles; ++c) {
1400
+ hmx_fa_o_update_tile(d_diag, o_rc, p_tile_in, v_tile_in, o_tile_out, n_col_tiles);
1401
+ o_rc += o_stride;
1402
+ v_tile_in += v_stride;
1403
+ o_tile_out += o_stride;
1404
+ }
1405
+ }
1406
+ }
1407
+
1408
+ typedef struct {
1409
+ __fp16 * o_curr; // output (row-major tile layout)
1410
+ const __fp16 * o_prev; // input (column-major tile layout)
1411
+ const __fp16 * d_tiles; // diag(1/l) tiles
1412
+ uint8_t * hmx_scales;
1413
+ size_t n_row_tiles;
1414
+ size_t n_row_tiles_g_br;
1415
+ size_t DV;
1416
+ } hmx_fa_o_norm_job_t;
1417
+
1418
+ static void hmx_fa_o_norm_worker(void * data) {
1419
+ hmx_fa_o_norm_job_t * job = (hmx_fa_o_norm_job_t *) data;
1420
+ const size_t n_row_tiles = job->n_row_tiles;
1421
+ const size_t n_row_tiles_g_br = job->n_row_tiles_g_br;
1422
+ const size_t DV_tiles = job->DV / 32;
1423
+ const __fp16 * restrict d_tiles = job->d_tiles;
1424
+ const __fp16 * restrict o_prev = job->o_prev;
1425
+ __fp16 * restrict o_curr = job->o_curr;
1426
+ __builtin_assume(n_row_tiles > 0);
1427
+ __builtin_assume(DV_tiles > 0);
1428
+
1429
+ asm volatile(HMX_SET_BIAS("%0") :: "r"((unsigned int)job->hmx_scales));
1430
+ const size_t o_stride = n_row_tiles_g_br * HMX_FP16_TILE_N_ELMS;
1431
+ for (size_t r = 0; r < n_row_tiles; ++r) {
1432
+ const __fp16 * d_diag = d_tiles + r * (n_row_tiles_g_br + 1) * HMX_FP16_TILE_N_ELMS;
1433
+ const __fp16 * o_rc = o_prev + r * HMX_FP16_TILE_N_ELMS;
1434
+ __fp16 * o_out = o_curr + r * DV_tiles * HMX_FP16_TILE_N_ELMS;
1435
+
1436
+ for (size_t c = 0; c < DV_tiles; ++c) {
1437
+ hmx_fa_o_norm_tile(d_diag, o_rc, o_out);
1438
+ o_rc += o_stride;
1439
+ o_out += HMX_FP16_TILE_N_ELMS;
1440
+ }
1441
+ }
1442
+ }
1443
+
1444
+ // Populate per-GQA-row ALiBi slopes for a given KV head.
1445
+ static __attribute__((noinline)) void fa_compute_slopes(
1446
+ const struct hmx_fa_context * factx,
1447
+ uint32_t kv_head,
1448
+ size_t n_rows_g) {
1449
+ __fp16 * slopes = factx->vtcm_slopes;
1450
+ if (factx->max_bias == 0.0f) {
1451
+ hvx_splat_f16_a(slopes, 1.0f, n_rows_g);
1452
+ return;
1453
+ }
1454
+
1455
+ const uint32_t G = factx->G;
1456
+ const uint32_t n_head_log2 = factx->n_head_log2;
1457
+ const float m0 = factx->m0;
1458
+ const float m1 = factx->m1;
1459
+
1460
+ __fp16 temp_slopes[512] __attribute__((aligned(128)));
1461
+ if (G <= 32) {
1462
+ // Fast path: Compute G unique slope values in vector registers
1463
+ HVX_Vector v_val = hvx_alibi_slopes(kv_head, G, n_head_log2, m0, m1);
1464
+
1465
+ __fp16 temp_slopes_aligned[64] __attribute__((aligned(128)));
1466
+ hvx_vmem(temp_slopes_aligned) = hvx_vec_f32_to_f16(v_val, Q6_V_vzero());
1467
+
1468
+ for (uint32_t i = 0; i < G; ++i) {
1469
+ temp_slopes[i] = temp_slopes_aligned[i];
1470
+ }
1471
+ } else {
1472
+ // Fallback path: G > 32 (rare configurations)
1473
+ for (uint32_t i = 0; i < G; ++i) {
1474
+ temp_slopes[i] = (__fp16)alibi_slope(kv_head * G + i, n_head_log2, m0, m1);
1475
+ }
1476
+ }
1477
+
1478
+ // Allocate stack buffer to avoid scalar writes to VTCM (which generates L2 misses)
1479
+ __fp16 local_slopes[n_rows_g] __attribute__((aligned(128)));
1480
+ for (size_t r = 0; r < n_rows_g; ++r) {
1481
+ local_slopes[r] = temp_slopes[fastmodulo(r, G, &factx->div_G)];
1482
+ }
1483
+
1484
+ // Copy to VTCM slopes using HVX block copy (both are aligned to 128 bytes)
1485
+ hvx_copy_f16_aa((uint8_t *)slopes, (const uint8_t *)local_slopes, n_rows_g);
1486
+ }
1487
+
1488
+ static void fa_push_mask_dma_gqa(
1489
+ dma_queue * dma,
1490
+ const struct htp_tensor * mask,
1491
+ uint32_t q_start,
1492
+ uint32_t im3,
1493
+ uint32_t kv_start,
1494
+ uint32_t kv_head,
1495
+ uint32_t G,
1496
+ uint32_t m_line_bytes,
1497
+ uint32_t kv_rows,
1498
+ uint32_t n_rows_q,
1499
+ struct hmx_fa_context * factx
1500
+ ) {
1501
+ for (uint32_t g = 0; g < G; ++g) {
1502
+ const uint32_t h_idx = kv_head * G + g;
1503
+ const uint32_t im2 = fastmodulo(h_idx, mask->ne[2], &factx->src3_div2);
1504
+ const uint8_t * ms_src = (const uint8_t *) mask->data + q_start * mask->nb[1] +
1505
+ im2 * mask->nb[2] + im3 * mask->nb[3] + kv_start * sizeof(__fp16);
1506
+ uint8_t * ms_dst = (uint8_t *) factx->vtcm_mask_buf + g * m_line_bytes;
1507
+ dma_queue_push(dma, dma_make_ptr(ms_dst, ms_src), G * m_line_bytes, mask->nb[1], kv_rows * sizeof(__fp16), n_rows_q);
1508
+ }
1509
+ }
1510
+
1511
+ static void fa_pop_mask_dma_gqa(dma_queue * dma, uint32_t G) {
1512
+ for (uint32_t g = 0; g < G; ++g) {
1513
+ dma_queue_pop(dma);
1514
+ }
1515
+ }
1516
+
1517
+ // ============================================================================
1518
+ // Core HMX flash attention algorithm (GQA-merged)
1519
+ // ============================================================================
1520
+
1521
+ int hmx_flash_attn_ext(struct htp_ops_context * octx) {
1522
+ struct htp_thread_trace * tr_hvx = octx->ctx ? &octx->ctx->trace[0] : NULL;
1523
+ struct htp_thread_trace * tr_hmx = octx->ctx ? &octx->ctx->trace[HTP_MAX_NTHREADS] : NULL;
1524
+ const struct htp_tensor * q = octx->src[0];
1525
+ const struct htp_tensor * k = octx->src[1];
1526
+ const struct htp_tensor * v = octx->src[2];
1527
+ const struct htp_tensor * mask = (octx->src[3] && octx->src[3]->data) ? octx->src[3] : NULL;
1528
+ const struct htp_tensor * dst = octx->dst;
1529
+
1530
+ struct htp_context * const ctx = octx->ctx;
1531
+
1532
+ if (!ctx->hmx_enabled) {
1533
+ return HTP_STATUS_NO_SUPPORT;
1534
+ }
1535
+
1536
+ // Dimensions
1537
+ const uint32_t neq0 = q->ne[0]; // head_dim (DK)
1538
+ const uint32_t neq1 = q->ne[1]; // n_tokens
1539
+ const uint32_t neq2 = q->ne[2]; // n_heads
1540
+ const uint32_t neq3 = q->ne[3]; // n_seqs
1541
+
1542
+ const uint32_t nek0 = k->ne[0]; // head_dim
1543
+ const uint32_t nek1 = k->ne[1]; // kv_len
1544
+
1545
+ const uint32_t nev0 = v->ne[0]; // head_dim (DV)
1546
+
1547
+ const uint32_t DK = neq0;
1548
+ const uint32_t DV = nev0;
1549
+
1550
+ // HMX requires head_dim to be multiple of 32
1551
+ if (DK % 32 != 0 || DV % 32 != 0) {
1552
+ return HTP_STATUS_NO_SUPPORT;
1553
+ }
1554
+
1555
+ const struct htp_fa_kernel_params * kparams = (const struct htp_fa_kernel_params *) octx->kernel_params;
1556
+ const uint32_t n_kv_heads = k->ne[2];
1557
+
1558
+ // ======== Build context ========
1559
+ struct hmx_fa_context factx;
1560
+ memset(&factx, 0, sizeof(factx));
1561
+ factx.octx = octx;
1562
+ factx.sinks = octx->src[4]; // NULL if this op has no attention sinks
1563
+ factx.n_threads = kparams->n_threads;
1564
+ factx.DK = DK;
1565
+ factx.DV = DV;
1566
+ factx.n_kv = nek1;
1567
+ factx.n_kv_heads = n_kv_heads;
1568
+ factx.n_heads = neq2;
1569
+ factx.G = kparams->G;
1570
+ factx.div_G = kparams->u.hmx.div_G;
1571
+ factx.neq1 = neq1;
1572
+ factx.Br = kparams->Br;
1573
+ factx.Bc = kparams->Bc;
1574
+ factx.g_br = kparams->u.hmx.g_br;
1575
+ factx.n_kv_blocks = kparams->n_kv_blocks;
1576
+ factx.is_q_fp32 = (kparams->is_q_fp32 != 0);
1577
+ factx.is_dst_fp32 = (kparams->is_dst_fp32 != 0);
1578
+ factx.pipeline = (kparams->u.hmx.pipeline != 0);
1579
+ factx.mask_broadcast = (kparams->u.hmx.mask_broadcast != 0);
1580
+ if (mask) {
1581
+ factx.src3_div2 = kparams->src3_div2;
1582
+ factx.src3_div3 = kparams->src3_div3;
1583
+ }
1584
+
1585
+ if (kparams->logit_softcap == 0.0f) {
1586
+ factx.scale = (__fp16) (kparams->scale * EXP_LOG2E_F); // log2(e)
1587
+ } else {
1588
+ factx.scale = (__fp16) kparams->scale;
1589
+ }
1590
+ factx.max_bias = kparams->max_bias;
1591
+ factx.logit_softcap = (__fp16) (kparams->logit_softcap * EXP_LOG2E_F);
1592
+
1593
+ factx.n_head_log2 = kparams->n_head_log2;
1594
+ factx.m0 = kparams->m0;
1595
+ factx.m1 = kparams->m1;
1596
+
1597
+ const uint32_t Br = factx.Br;
1598
+ const uint32_t Bc = factx.Bc;
1599
+ const uint32_t g_br = factx.g_br;
1600
+ const bool pipeline = factx.pipeline;
1601
+ const uint32_t n_threads = factx.n_threads;
1602
+ const uint32_t G = factx.G;
1603
+
1604
+ // ======== VTCM allocation (GQA-aware) ========
1605
+ // K/V row sizes drive the DMA descriptors (not the VTCM layout) and are used
1606
+ // throughout the KV loop below.
1607
+ const size_t size_k_row = DK * sizeof(__fp16);
1608
+ const size_t size_v_row = DV * sizeof(__fp16);
1609
+ const size_t size_k_row_padded = hex_round_up(size_k_row, 128);
1610
+ const size_t size_v_row_padded = hex_round_up(size_v_row, 128);
1611
+
1612
+ // Build the VTCM layout once (shared with the host estimator) and place every
1613
+ // scratch buffer at its computed offset.
1614
+ struct hmx_fa_vtcm_layout L;
1615
+ hmx_fa_vtcm_layout_build(&L, G, DK, DV, Br, Bc, n_threads, pipeline);
1616
+
1617
+ if (L.total_bytes > ctx->vtcm_size) {
1618
+ return HTP_STATUS_VTCM_TOO_SMALL;
1619
+ }
1620
+
1621
+ uint8_t * const base = ctx->vtcm_base;
1622
+
1623
+ factx.vtcm_q_tiles = VTCM_LAYOUT_PTR(__fp16, base, L.off_q_tiles);
1624
+ factx.vtcm_o_tiles[0] = VTCM_LAYOUT_PTR(__fp16, base, L.off_o_tiles[0]);
1625
+ factx.vtcm_o_tiles[1] = VTCM_LAYOUT_PTR(__fp16, base, L.off_o_tiles[1]);
1626
+ factx.vtcm_k_fp16[0] = VTCM_LAYOUT_PTR(__fp16, base, L.off_k_fp16[0]);
1627
+ factx.vtcm_k_fp16[1] = VTCM_LAYOUT_PTR(__fp16, base, L.off_k_fp16[1]);
1628
+ factx.vtcm_v_fp16[0] = VTCM_LAYOUT_PTR(__fp16, base, L.off_v_fp16[0]);
1629
+ factx.vtcm_v_fp16[1] = VTCM_LAYOUT_PTR(__fp16, base, L.off_v_fp16[1]);
1630
+ factx.vtcm_k_tiles = VTCM_LAYOUT_PTR(__fp16, base, L.off_k_tiles);
1631
+ factx.vtcm_v_tiles[0] = VTCM_LAYOUT_PTR(__fp16, base, L.off_v_tiles[0]);
1632
+ factx.vtcm_v_tiles[1] = VTCM_LAYOUT_PTR_OPTIONAL(__fp16, base, L.off_v_tiles[1], pipeline);
1633
+ factx.vtcm_s_tiles = VTCM_LAYOUT_PTR(__fp16, base, L.off_s_tiles);
1634
+ factx.vtcm_p_tiles = VTCM_LAYOUT_PTR(__fp16, base, L.off_p_tiles);
1635
+ factx.vtcm_d_tiles = VTCM_LAYOUT_PTR(__fp16, base, L.off_d_tiles);
1636
+ factx.vtcm_m_vec = VTCM_LAYOUT_PTR(HVX_Vector, base, L.off_m_vec);
1637
+ factx.vtcm_l_vec = VTCM_LAYOUT_PTR(HVX_Vector, base, L.off_l_vec);
1638
+ factx.vtcm_s_rowmax = VTCM_LAYOUT_PTR(HVX_Vector, base, L.off_s_rowmax);
1639
+ factx.vtcm_p_rowsum = VTCM_LAYOUT_PTR(HVX_Vector, base, L.off_p_rowsum);
1640
+ factx.vtcm_row_bufs = VTCM_LAYOUT_PTR(HVX_Vector, base, L.off_row_bufs);
1641
+ factx.row_buf_stride = L.row_buf_stride;
1642
+ factx.vtcm_hmx_scales_id = VTCM_LAYOUT_PTR(uint8_t, base, L.off_hmx_scales_id);
1643
+ factx.vtcm_hmx_scales_qk = VTCM_LAYOUT_PTR(uint8_t, base, L.off_hmx_scales_qk);
1644
+ factx.vtcm_mask_buf = VTCM_LAYOUT_PTR(__fp16, base, L.off_mask_buf);
1645
+ factx.mask_buf_row_stride = L.mask_buf_row_stride;
1646
+ factx.q_tile_bytes = L.q_tile_bytes;
1647
+ factx.o_tile_bytes = L.o_tile_bytes;
1648
+ factx.col_vec_bytes = L.col_vec_bytes;
1649
+ factx.d_tile_bytes = L.d_tile_bytes;
1650
+ factx.vtcm_slopes = VTCM_LAYOUT_PTR(__fp16, base, L.off_slopes);
1651
+
1652
+ const size_t m_line_bytes = L.m_line_bytes; // used by the mask DMAs in the KV loop
1653
+
1654
+ dma_cache_init(&factx.m_cache, (uint8_t *) factx.vtcm_mask_buf, L.m_buf_slot_bytes, HMX_FA_DMA_CACHE_SIZE);
1655
+
1656
+ // ======== Initialize HMX output scales ========
1657
+ hmx_init_column_scales(factx.vtcm_hmx_scales_id, Q6_V_vsplat_R(0x3c00)); // 1.0
1658
+ hmx_init_column_scales(factx.vtcm_hmx_scales_qk, hvx_vec_splat_f16(factx.scale));
1659
+
1660
+ // ======== Skip compute if profiling ========
1661
+ if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) {
1662
+ return HTP_STATUS_OK;
1663
+ }
1664
+
1665
+ // ======== DMA setup ========
1666
+ dma_queue * const dma = ctx->dma[0];
1667
+
1668
+ const size_t n_row_tiles_g_br = g_br / HMX_FP16_TILE_N_ROWS;
1669
+ const size_t n_tiles_per_bc = Bc / HMX_FP16_TILE_N_COLS;
1670
+
1671
+ const size_t qo_element_size = factx.is_q_fp32 ? sizeof(float) : sizeof(__fp16);
1672
+
1673
+ // ======== Reusable job descriptors for pipeline ========
1674
+ hmx_fa_qk_job_t qk_job;
1675
+ hmx_fa_o_update_job_t ou_job;
1676
+ hmx_fa_o_norm_job_t on_job;
1677
+
1678
+ // ======== Main loop ========
1679
+ for (uint32_t ib3 = 0; ib3 < neq3; ++ib3) {
1680
+ const uint32_t im3 = mask ? fastmodulo(ib3, mask->ne[3], &factx.src3_div3) : 0;
1681
+ for (uint32_t q_start = 0; q_start < neq1; q_start += Br) {
1682
+ const uint32_t n_rows_q = hex_smin(Br, neq1 - q_start);
1683
+ const size_t n_rows_g = n_rows_q * G;
1684
+ const size_t g_br_actual = hex_align_up(n_rows_g, HMX_FP16_TILE_N_ROWS);
1685
+ const size_t n_row_tiles = g_br_actual / HMX_FP16_TILE_N_ROWS;
1686
+
1687
+ for (uint32_t kv_head = 0; kv_head < n_kv_heads; ++kv_head) {
1688
+ const uint32_t ik2 = kv_head;
1689
+ const uint32_t ik3 = fastdiv(ib3, &kparams->broadcast_rk3);
1690
+ const uint32_t iv2 = kv_head;
1691
+ const uint32_t iv3 = fastdiv(ib3, &kparams->broadcast_rv3);
1692
+
1693
+ // 1. Push Q DMA (if Q DMA is used)
1694
+ const size_t o_tile_bytes = factx.o_tile_bytes;
1695
+ const bool use_q_dma = (2 * o_tile_bytes >= factx.g_br * factx.DK * (factx.is_q_fp32 ? 4 : 2));
1696
+ if (use_q_dma) {
1697
+ const bool q_transposed = q->nb[1] < q->nb[2];
1698
+ const uint8_t * q_ptr = (const uint8_t *) q->data + q_start * q->nb[1] + (kv_head * factx.G) * q->nb[2] + ib3 * q->nb[3];
1699
+ const size_t el_size = factx.is_q_fp32 ? sizeof(float) : sizeof(__fp16);
1700
+ const size_t q_row_bytes = q_transposed ? n_rows_q * factx.DK * el_size : factx.G * factx.DK * el_size;
1701
+ const size_t src_stride = q_transposed ? q->nb[2] : q->nb[1];
1702
+ const size_t n_rows = q_transposed ? factx.G : n_rows_q;
1703
+ dma_queue_push(dma, dma_make_ptr(factx.vtcm_o_tiles[0], q_ptr), q_row_bytes, hex_smax(src_stride, q_row_bytes), q_row_bytes, n_rows);
1704
+ }
1705
+
1706
+ // 2. Prefetch first KV block
1707
+ if (factx.n_kv_blocks > 0) {
1708
+ const uint32_t kv_rows0 = hex_smin(Bc, nek1);
1709
+
1710
+ const uint8_t * k_src = (const uint8_t *) k->data + ik2 * k->nb[2] + ik3 * k->nb[3];
1711
+ dma_queue_push(dma, dma_make_ptr(factx.vtcm_k_fp16[0], k_src), size_k_row_padded, k->nb[1], size_k_row, kv_rows0);
1712
+
1713
+ const uint8_t * v_src = (const uint8_t *) v->data + iv2 * v->nb[2] + iv3 * v->nb[3];
1714
+ dma_queue_push(dma, dma_make_ptr(factx.vtcm_v_fp16[0], v_src), size_v_row_padded, v->nb[1], size_v_row, kv_rows0);
1715
+ }
1716
+
1717
+ // 3. Pop Q DMA (blocks until Q is loaded)
1718
+ if (use_q_dma) {
1719
+ dma_queue_pop(dma);
1720
+ }
1721
+
1722
+ // ---- Load Q block & Initialize per-block state ----
1723
+ fa_phase_q_load(&factx, q, q_start, kv_head, ib3, n_rows_g);
1724
+
1725
+ __fp16 * o_tile_prev = factx.vtcm_o_tiles[0];
1726
+ __fp16 * o_tile_curr = factx.vtcm_o_tiles[1];
1727
+
1728
+ // ---- KV block loop with DMA double-buffering ----
1729
+ size_t buf_idx = 0;
1730
+
1731
+ htp_trace_event_start(tr_hvx, HTP_TRACE_EVT_HVX_A_PREP, (uint16_t) q_start);
1732
+ fa_compute_slopes(&factx, kv_head, n_rows_g);
1733
+ htp_trace_event_stop(tr_hvx, HTP_TRACE_EVT_HVX_A_PREP, (uint16_t) q_start);
1734
+
1735
+ const size_t k_src_stride = size_k_row_padded / sizeof(__fp16);
1736
+ const size_t v_src_stride = size_v_row_padded / sizeof(__fp16);
1737
+
1738
+ struct hmx_queue * hmx_q = ctx->hmx_queue;
1739
+
1740
+ if (factx.pipeline) {
1741
+ // Pipeline path
1742
+ for (uint32_t kv_blk = 0; kv_blk < factx.n_kv_blocks; ++kv_blk) {
1743
+ const uint32_t kv_start = kv_blk * Bc;
1744
+ const uint32_t kv_rows = hex_smin(Bc, nek1 - kv_start);
1745
+ const size_t n_col_tiles = hmx_ceil_div(kv_rows, HMX_FP16_TILE_N_COLS);
1746
+
1747
+ // Push mask DMA
1748
+ if (mask) {
1749
+ if (__builtin_expect(factx.mask_broadcast, true)) {
1750
+ const uint8_t * ms_src = (const uint8_t *) mask->data + q_start * mask->nb[1] + im3 * mask->nb[3] + kv_start * sizeof(__fp16);
1751
+ dma_cache_push(dma, &factx.m_cache, ms_src, m_line_bytes, mask->nb[1], kv_rows * sizeof(__fp16), n_rows_q);
1752
+ } else {
1753
+ fa_push_mask_dma_gqa(dma, mask, q_start, im3, kv_start, kv_head, G, m_line_bytes, kv_rows, n_rows_q, &factx);
1754
+ }
1755
+ }
1756
+
1757
+ // Prefetch next KV block early
1758
+ if (kv_blk + 1 < factx.n_kv_blocks) {
1759
+ const uint32_t prefetch_start = (kv_blk + 1) * Bc;
1760
+ const uint32_t prefetch_rows = hex_smin(Bc, nek1 - prefetch_start);
1761
+ const size_t prefetch_buf = 1 - buf_idx;
1762
+ const uint8_t * k_prefetch_src = (const uint8_t *) k->data + prefetch_start * k->nb[1] + ik2 * k->nb[2] + ik3 * k->nb[3];
1763
+ dma_queue_push(dma, dma_make_ptr(factx.vtcm_k_fp16[prefetch_buf], k_prefetch_src), size_k_row_padded, k->nb[1], size_k_row, prefetch_rows);
1764
+ const uint8_t * v_prefetch_src = (const uint8_t *) v->data + prefetch_start * v->nb[1] + iv2 * v->nb[2] + iv3 * v->nb[3];
1765
+ dma_queue_push(dma, dma_make_ptr(factx.vtcm_v_fp16[prefetch_buf], v_prefetch_src), size_v_row_padded, v->nb[1], size_v_row, prefetch_rows);
1766
+ }
1767
+
1768
+ // ---- Phase 1: K_int ----
1769
+ if (kv_blk > 0) {
1770
+ ou_job.o_curr = o_tile_curr;
1771
+ ou_job.o_prev = o_tile_prev;
1772
+ ou_job.p_tiles = factx.vtcm_p_tiles;
1773
+ ou_job.v_tiles = factx.vtcm_v_tiles[1 - buf_idx];
1774
+ ou_job.d_tiles = factx.vtcm_d_tiles;
1775
+ ou_job.hmx_scales = factx.vtcm_hmx_scales_id;
1776
+ ou_job.n_row_tiles = n_row_tiles;
1777
+ ou_job.n_col_tiles = hmx_ceil_div(hex_smin(Bc, nek1 - (kv_blk - 1) * Bc), HMX_FP16_TILE_N_COLS);
1778
+ ou_job.n_row_tiles_g_br = n_row_tiles_g_br;
1779
+ ou_job.n_tiles_per_bc = n_tiles_per_bc;
1780
+ ou_job.DV = DV;
1781
+ hmx_queue_push(hmx_q, hmx_queue_make_desc(hmx_fa_o_update_worker, &ou_job));
1782
+ }
1783
+
1784
+ // Wait for current K DMA and interleave
1785
+ void * curr_k = dma_queue_pop(dma).dst;
1786
+ fa_phase_k_interleave(&factx, kv_rows, k_src_stride, curr_k, kv_start);
1787
+
1788
+ // ---- Phase 2: qk_dot ----
1789
+ qk_job.q_tiles = factx.vtcm_q_tiles;
1790
+ qk_job.k_tiles = factx.vtcm_k_tiles;
1791
+ qk_job.s_tiles = factx.vtcm_s_tiles;
1792
+ qk_job.n_row_tiles = n_row_tiles;
1793
+ qk_job.n_col_tiles = n_col_tiles;
1794
+ qk_job.n_dot_tiles = DK / 32;
1795
+ qk_job.n_tiles_per_bc = n_tiles_per_bc;
1796
+ qk_job.hmx_scales = factx.vtcm_hmx_scales_qk;
1797
+ hmx_queue_push(hmx_q, hmx_queue_make_desc(hmx_fa_qk_dot_worker, &qk_job));
1798
+
1799
+ // Wait for current V DMA and interleave
1800
+ void * curr_v = dma_queue_pop(dma).dst;
1801
+ fa_phase_v_interleave(&factx, kv_rows, v_src_stride, curr_v, factx.vtcm_v_tiles[buf_idx], n_tiles_per_bc, kv_start);
1802
+
1803
+ if (kv_blk > 0) {
1804
+ hmx_queue_pop(hmx_q);
1805
+ hex_swap_ptr((void **) &o_tile_curr, (void **) &o_tile_prev);
1806
+ }
1807
+
1808
+ hmx_queue_pop(hmx_q);
1809
+
1810
+ // ---- Phase 3: softmax + build_D ----
1811
+ __fp16 * current_mask_vtcm = NULL;
1812
+ if (mask) {
1813
+ if (__builtin_expect(factx.mask_broadcast, true)) {
1814
+ current_mask_vtcm = (__fp16 *) dma_queue_pop(dma).dst;
1815
+ } else {
1816
+ fa_pop_mask_dma_gqa(dma, G);
1817
+ current_mask_vtcm = factx.vtcm_mask_buf;
1818
+ }
1819
+ }
1820
+
1821
+ fa_softmax_args_t sargs;
1822
+ memset(&sargs, 0, sizeof(sargs));
1823
+ sargs.factx = &factx;
1824
+ sargs.kv_rows = kv_rows;
1825
+ sargs.n_rows_g = n_rows_g;
1826
+ sargs.n_col_tiles = n_col_tiles;
1827
+ sargs.n_tiles_per_bc = n_tiles_per_bc;
1828
+ sargs.n_row_tiles = n_row_tiles;
1829
+ sargs.n_row_tiles_g_br = n_row_tiles_g_br;
1830
+ sargs.Bc = Bc;
1831
+ sargs.G = G;
1832
+ sargs.kv_head = kv_head;
1833
+ sargs.kv_start = kv_start;
1834
+ sargs.q_start = q_start;
1835
+ sargs.ib3 = ib3;
1836
+ sargs.has_alibi = (factx.max_bias != 0.0f);
1837
+ sargs.mask = mask;
1838
+ sargs.mask_vtcm = current_mask_vtcm;
1839
+ sargs.mask_vtcm_row_stride = factx.mask_buf_row_stride;
1840
+ sargs.slopes = factx.vtcm_slopes;
1841
+ fa_phase_softmax_and_build_d(&factx, &sargs, n_row_tiles, n_row_tiles_g_br);
1842
+
1843
+ buf_idx = 1 - buf_idx;
1844
+ }
1845
+
1846
+ // Epilogue
1847
+ if (factx.n_kv_blocks > 0) {
1848
+ const uint32_t last_blk = factx.n_kv_blocks - 1;
1849
+ const size_t last_cols = hmx_ceil_div(hex_smin(Bc, nek1 - last_blk * Bc), HMX_FP16_TILE_N_COLS);
1850
+ ou_job.o_curr = o_tile_curr;
1851
+ ou_job.o_prev = o_tile_prev;
1852
+ ou_job.p_tiles = factx.vtcm_p_tiles;
1853
+ ou_job.v_tiles = factx.vtcm_v_tiles[1 - buf_idx];
1854
+ ou_job.d_tiles = factx.vtcm_d_tiles;
1855
+ ou_job.hmx_scales = factx.vtcm_hmx_scales_id;
1856
+ ou_job.n_row_tiles = n_row_tiles;
1857
+ ou_job.n_col_tiles = last_cols;
1858
+ ou_job.n_row_tiles_g_br = n_row_tiles_g_br;
1859
+ ou_job.n_tiles_per_bc = n_tiles_per_bc;
1860
+ ou_job.DV = DV;
1861
+ hmx_queue_push(hmx_q, hmx_queue_make_desc(hmx_fa_o_update_worker, &ou_job));
1862
+ hmx_queue_pop(hmx_q);
1863
+
1864
+ hex_swap_ptr((void **) &o_tile_curr, (void **) &o_tile_prev);
1865
+ }
1866
+
1867
+ } else {
1868
+ // Fallback path
1869
+ for (uint32_t kv_blk = 0; kv_blk < factx.n_kv_blocks; ++kv_blk) {
1870
+ const uint32_t kv_start = kv_blk * Bc;
1871
+ const uint32_t kv_rows = hex_smin(Bc, nek1 - kv_start);
1872
+ const size_t n_col_tiles = hmx_ceil_div(kv_rows, HMX_FP16_TILE_N_COLS);
1873
+
1874
+ if (mask) {
1875
+ if (__builtin_expect(factx.mask_broadcast, true)) {
1876
+ const uint8_t * ms_src = (const uint8_t *) mask->data + q_start * mask->nb[1] + im3 * mask->nb[3] + kv_start * sizeof(__fp16);
1877
+ dma_cache_push(dma, &factx.m_cache, ms_src, m_line_bytes, mask->nb[1], kv_rows * sizeof(__fp16), n_rows_q);
1878
+ } else {
1879
+ fa_push_mask_dma_gqa(dma, mask, q_start, im3, kv_start, kv_head, G, m_line_bytes, kv_rows, n_rows_q, &factx);
1880
+ }
1881
+ }
1882
+
1883
+ if (kv_blk + 1 < factx.n_kv_blocks) {
1884
+ const uint32_t prefetch_start = (kv_blk + 1) * Bc;
1885
+ const uint32_t prefetch_rows = hex_smin(Bc, nek1 - prefetch_start);
1886
+ const size_t prefetch_buf = 1 - buf_idx;
1887
+ const uint8_t * k_prefetch_src = (const uint8_t *) k->data + prefetch_start * k->nb[1] + ik2 * k->nb[2] + ik3 * k->nb[3];
1888
+ dma_queue_push(dma, dma_make_ptr(factx.vtcm_k_fp16[prefetch_buf], k_prefetch_src), size_k_row_padded, k->nb[1], size_k_row, prefetch_rows);
1889
+ const uint8_t * v_prefetch_src = (const uint8_t *) v->data + prefetch_start * v->nb[1] + iv2 * v->nb[2] + iv3 * v->nb[3];
1890
+ dma_queue_push(dma, dma_make_ptr(factx.vtcm_v_fp16[prefetch_buf], v_prefetch_src), size_v_row_padded, v->nb[1], size_v_row, prefetch_rows);
1891
+ }
1892
+
1893
+ // Wait for current K DMA and interleave
1894
+ void * curr_k = dma_queue_pop(dma).dst;
1895
+ fa_phase_k_interleave(&factx, kv_rows, k_src_stride, curr_k, kv_start);
1896
+
1897
+ {
1898
+ qk_job.q_tiles = factx.vtcm_q_tiles;
1899
+ qk_job.k_tiles = factx.vtcm_k_tiles;
1900
+ qk_job.s_tiles = factx.vtcm_s_tiles;
1901
+ qk_job.n_row_tiles = n_row_tiles;
1902
+ qk_job.n_col_tiles = n_col_tiles;
1903
+ qk_job.n_dot_tiles = (size_t) (DK / 32);
1904
+ qk_job.n_tiles_per_bc = n_tiles_per_bc;
1905
+ qk_job.hmx_scales = factx.vtcm_hmx_scales_qk;
1906
+
1907
+ hmx_queue_push(ctx->hmx_queue, hmx_queue_make_desc(hmx_fa_qk_dot_worker, &qk_job));
1908
+ hmx_queue_pop(ctx->hmx_queue);
1909
+ }
1910
+
1911
+ // Wait for current V DMA and interleave
1912
+ void * curr_v = dma_queue_pop(dma).dst;
1913
+ fa_phase_v_interleave(&factx, kv_rows, v_src_stride, curr_v, factx.vtcm_v_tiles[0], n_tiles_per_bc, kv_start);
1914
+
1915
+ // ---- Phase 3: softmax + build_D ----
1916
+ __fp16 * current_mask_vtcm = NULL;
1917
+ if (mask) {
1918
+ if (__builtin_expect(factx.mask_broadcast, true)) {
1919
+ current_mask_vtcm = (__fp16 *) dma_queue_pop(dma).dst;
1920
+ } else {
1921
+ fa_pop_mask_dma_gqa(dma, G);
1922
+ current_mask_vtcm = factx.vtcm_mask_buf;
1923
+ }
1924
+ }
1925
+
1926
+ fa_softmax_args_t sargs;
1927
+ memset(&sargs, 0, sizeof(sargs));
1928
+ sargs.factx = &factx;
1929
+ sargs.kv_rows = kv_rows;
1930
+ sargs.n_rows_g = n_rows_g;
1931
+ sargs.n_col_tiles = n_col_tiles;
1932
+ sargs.n_tiles_per_bc = n_tiles_per_bc;
1933
+ sargs.n_row_tiles = n_row_tiles;
1934
+ sargs.n_row_tiles_g_br = n_row_tiles_g_br;
1935
+ sargs.Bc = Bc;
1936
+ sargs.G = G;
1937
+ sargs.kv_head = kv_head;
1938
+ sargs.kv_start = kv_start;
1939
+ sargs.q_start = q_start;
1940
+ sargs.ib3 = ib3;
1941
+ sargs.has_alibi = (factx.max_bias != 0.0f);
1942
+ sargs.mask = mask;
1943
+ sargs.mask_vtcm = current_mask_vtcm;
1944
+ sargs.mask_vtcm_row_stride = factx.mask_buf_row_stride;
1945
+ sargs.slopes = factx.vtcm_slopes;
1946
+ fa_phase_softmax_and_build_d(&factx, &sargs, n_row_tiles, n_row_tiles_g_br);
1947
+
1948
+ {
1949
+ ou_job.o_curr = o_tile_curr;
1950
+ ou_job.o_prev = o_tile_prev;
1951
+ ou_job.p_tiles = factx.vtcm_p_tiles;
1952
+ ou_job.v_tiles = factx.vtcm_v_tiles[0];
1953
+ ou_job.d_tiles = factx.vtcm_d_tiles;
1954
+ ou_job.hmx_scales = factx.vtcm_hmx_scales_id;
1955
+ ou_job.n_row_tiles = n_row_tiles;
1956
+ ou_job.n_col_tiles = n_col_tiles;
1957
+ ou_job.n_row_tiles_g_br = n_row_tiles_g_br;
1958
+ ou_job.n_tiles_per_bc = n_tiles_per_bc;
1959
+ ou_job.DV = DV;
1960
+
1961
+ hmx_queue_push(ctx->hmx_queue, hmx_queue_make_desc(hmx_fa_o_update_worker, &ou_job));
1962
+ hmx_queue_pop(ctx->hmx_queue);
1963
+
1964
+ hex_swap_ptr((void **) &o_tile_curr, (void **) &o_tile_prev);
1965
+ }
1966
+
1967
+ buf_idx = 1 - buf_idx;
1968
+ }
1969
+ }
1970
+
1971
+ // ---- Final normalization ----
1972
+ {
1973
+ htp_trace_event_start(tr_hvx, HTP_TRACE_EVT_HVX_O_PROC, (uint16_t) q_start);
1974
+ fa_build_d_diag_inv_l(&factx, n_row_tiles, n_row_tiles_g_br);
1975
+ htp_trace_event_stop(tr_hvx, HTP_TRACE_EVT_HVX_O_PROC, (uint16_t) q_start);
1976
+
1977
+ on_job.o_curr = o_tile_curr;
1978
+ on_job.o_prev = o_tile_prev;
1979
+ on_job.d_tiles = factx.vtcm_d_tiles;
1980
+ on_job.hmx_scales = factx.vtcm_hmx_scales_id;
1981
+ on_job.n_row_tiles = n_row_tiles;
1982
+ on_job.n_row_tiles_g_br = n_row_tiles_g_br;
1983
+ on_job.DV = DV;
1984
+ hmx_queue_push(ctx->hmx_queue, hmx_queue_make_desc(hmx_fa_o_norm_worker, &on_job));
1985
+ hmx_queue_pop(ctx->hmx_queue);
1986
+ }
1987
+
1988
+ // ---- Store O block ----
1989
+ fa_phase_o_store(&factx, dst, o_tile_curr, q_start, kv_head, ib3, n_rows_g);
1990
+ }
1991
+ }
1992
+ }
1993
+
1994
+ return HTP_STATUS_OK;
1995
+ }
1996
+
620
1997
  int op_flash_attn_ext(struct htp_ops_context * octx) {
621
1998
  const struct htp_tensor * q = octx->src[0];
622
1999
  const struct htp_tensor * k = octx->src[1];
@@ -629,110 +2006,83 @@ int op_flash_attn_ext(struct htp_ops_context * octx) {
629
2006
  return HTP_STATUS_NO_SUPPORT;
630
2007
  }
631
2008
 
632
- #ifdef HTP_HAS_HMX
633
- // HMX path: head_dim multiple of 64, F16 KV, and no sinks
634
- if (k->type == HTP_TYPE_F16 && v->type == HTP_TYPE_F16 && k->ne[0] % 64 == 0 && v->ne[0] % 64 == 0 && octx->src[4] == NULL) {
635
- int ret = hmx_flash_attn_ext(octx);
636
- if (ret == HTP_STATUS_OK) {
637
- return ret;
638
- }
639
- // VTCM too small or other failure -> fall through to HVX path
2009
+ const struct htp_fa_kernel_params * kparams = (const struct htp_fa_kernel_params *) octx->kernel_params;
2010
+
2011
+ if (kparams->kernel_type == HTP_FA_KERNEL_UNSUPPORTED) {
2012
+ return HTP_STATUS_NO_SUPPORT;
2013
+ }
2014
+
2015
+ if (kparams->kernel_type == HTP_FA_KERNEL_HMX) {
2016
+ return hmx_flash_attn_ext(octx);
640
2017
  }
641
- #endif
642
2018
 
643
2019
  struct htp_fa_context factx;
644
2020
  factx.octx = octx;
645
2021
 
646
2022
  factx.t_start = HAP_perf_get_qtimer_count();
647
2023
 
648
- factx.src0_div21 = init_fastdiv_values(q->ne[2] * q->ne[1]);
649
- factx.src0_div1 = init_fastdiv_values(q->ne[1]);
2024
+ factx.src0_div21 = kparams->u.hvx.src0_div21;
2025
+ factx.src0_div1 = kparams->u.hvx.src0_div1;
650
2026
 
651
- factx.broadcast_rk2 = init_fastdiv_values(q->ne[2]/k->ne[2]);
652
- factx.broadcast_rk3 = init_fastdiv_values(q->ne[3]/k->ne[3]);
653
- factx.broadcast_rv2 = init_fastdiv_values(q->ne[2]/v->ne[2]);
654
- factx.broadcast_rv3 = init_fastdiv_values(q->ne[3]/v->ne[3]);
2027
+ factx.broadcast_rk2 = kparams->broadcast_rk2;
2028
+ factx.broadcast_rk3 = kparams->broadcast_rk3;
2029
+ factx.broadcast_rv2 = kparams->broadcast_rv2;
2030
+ factx.broadcast_rv3 = kparams->broadcast_rv3;
655
2031
 
656
2032
  if (mask) {
657
- factx.src3_div2 = init_fastdiv_values(mask->ne[2]);
658
- factx.src3_div3 = init_fastdiv_values(mask->ne[3]);
2033
+ factx.src3_div2 = kparams->src3_div2;
2034
+ factx.src3_div3 = kparams->src3_div3;
659
2035
  }
660
2036
 
661
- factx.is_q_fp32 = (q->type == HTP_TYPE_F32);
662
- factx.size_q_row_padded = hex_round_up(q->ne[0] * (factx.is_q_fp32 ? 4 : 2), 128);
663
- factx.size_k_row_padded = hex_round_up(k->ne[0] * sizeof(__fp16), 128);
664
- factx.size_v_row_padded = hex_round_up(v->ne[0] * sizeof(__fp16), 128);
2037
+ factx.is_q_fp32 = (kparams->is_q_fp32 != 0);
2038
+ factx.size_q_row_padded = kparams->u.hvx.size_q_row_padded;
2039
+ factx.size_k_row_padded = kparams->u.hvx.size_k_row_padded;
2040
+ factx.size_v_row_padded = kparams->u.hvx.size_v_row_padded;
665
2041
 
666
2042
  size_t size_q_block = factx.size_q_row_padded * 1; // single row for now
667
2043
  factx.size_k_block = factx.size_k_row_padded * FLASH_ATTN_BLOCK_SIZE;
668
2044
  factx.size_v_block = factx.size_v_row_padded * FLASH_ATTN_BLOCK_SIZE;
669
2045
  factx.size_m_block = hex_round_up(FLASH_ATTN_BLOCK_SIZE * sizeof(__fp16), 128);
670
2046
 
671
- factx.n_blocks = (k->ne[1] + FLASH_ATTN_BLOCK_SIZE - 1) / FLASH_ATTN_BLOCK_SIZE;
672
-
673
- float scale = 1.0f;
674
- float max_bias = 0.0f;
675
- float logit_softcap = 0.0f;
676
-
677
- memcpy(&scale, (float *) octx->op_params + 0, sizeof(float));
678
- memcpy(&max_bias, (float *) octx->op_params + 1, sizeof(float));
679
- memcpy(&logit_softcap, (float *) octx->op_params + 2, sizeof(float));
2047
+ factx.n_blocks = kparams->n_kv_blocks;
680
2048
 
681
- if (logit_softcap != 0.0f) {
682
- scale /= logit_softcap;
683
- }
684
-
685
- factx.scale = scale;
686
- factx.max_bias = max_bias;
687
- factx.logit_softcap = logit_softcap;
2049
+ factx.scale = kparams->scale;
2050
+ factx.max_bias = kparams->max_bias;
2051
+ factx.logit_softcap = (__fp16) kparams->logit_softcap;
688
2052
 
689
- uint32_t n_head = q->ne[2];
690
- factx.n_head_log2 = 1u << (uint32_t) floor(log2(n_head));
691
- factx.m0 = powf(2.0f, -(max_bias ) / factx.n_head_log2);
692
- factx.m1 = powf(2.0f, -(max_bias / 2.0f) / factx.n_head_log2);
2053
+ factx.n_head_log2 = kparams->n_head_log2;
2054
+ factx.m0 = kparams->m0;
2055
+ factx.m1 = kparams->m1;
693
2056
 
2057
+ const uint32_t n_head = q->ne[2];
694
2058
  if (n_head > 512) {
695
2059
  return HTP_STATUS_NO_SUPPORT;
696
2060
  }
697
2061
  for (uint32_t h = 0; h < n_head; ++h) {
698
- factx.slopes[h] = (max_bias > 0.0f) ? alibi_slope(h, factx.n_head_log2, factx.m0, factx.m1) : 1.0f;
2062
+ factx.slopes[h] = (__fp16) ((kparams->max_bias > 0.0f) ? alibi_slope(h, factx.n_head_log2, factx.m0, factx.m1) : 1.0f);
699
2063
  }
700
2064
 
701
2065
  // total rows in q
702
- const uint32_t neq0 = q->ne[0];
703
- const uint32_t neq1 = q->ne[1];
704
- const uint32_t neq2 = q->ne[2];
705
- const uint32_t neq3 = q->ne[3];
706
-
707
- factx.qrows = neq1*neq2*neq3;
708
- factx.qrows_per_thread = (factx.qrows + octx->n_threads - 1) / octx->n_threads;
2066
+ factx.qrows = kparams->qrows;
2067
+ factx.qrows_per_thread = kparams->qrows_per_thread;
709
2068
 
710
2069
  size_t size_vkq_acc = hex_round_up(v->ne[0] * sizeof(float), 128); // VKQ32
711
2070
 
712
- octx->src0_spad.size_per_thread = size_q_block * 1;
713
- octx->src1_spad.size_per_thread = factx.size_k_block * 2;
714
- octx->src2_spad.size_per_thread = factx.size_v_block * 2;
715
- octx->src3_spad.size_per_thread = mask ? factx.size_m_block * DMA_CACHE_MAX_SIZE : 0;
716
- octx->dst_spad.size_per_thread = size_vkq_acc;
2071
+ factx.size_q_block = size_q_block;
2072
+ factx.size_vkq_acc = size_vkq_acc;
717
2073
 
718
- octx->src0_spad.size = octx->src0_spad.size_per_thread * octx->n_threads;
719
- octx->src1_spad.size = octx->src1_spad.size_per_thread * octx->n_threads;
720
- octx->src2_spad.size = octx->src2_spad.size_per_thread * octx->n_threads;
721
- octx->src3_spad.size = octx->src3_spad.size_per_thread * octx->n_threads;
722
- octx->dst_spad.size = octx->dst_spad.size_per_thread * octx->n_threads;
2074
+ uint8_t * vtcm_cur = octx->ctx->vtcm_base;
723
2075
 
724
- size_t total_spad = octx->src0_spad.size + octx->src1_spad.size + octx->src2_spad.size + octx->src3_spad.size + octx->dst_spad.size;
2076
+ factx.spad_q = vtcm_seq_alloc(&vtcm_cur, size_q_block * octx->n_threads);
2077
+ factx.spad_k = vtcm_seq_alloc(&vtcm_cur, factx.size_k_block * 2 * octx->n_threads);
2078
+ factx.spad_v = vtcm_seq_alloc(&vtcm_cur, factx.size_v_block * 2 * octx->n_threads);
2079
+ factx.spad_m = vtcm_seq_alloc(&vtcm_cur, (mask ? factx.size_m_block * HVX_FA_DMA_CACHE_SIZE : 0) * octx->n_threads);
2080
+ factx.spad_a = vtcm_seq_alloc(&vtcm_cur, size_vkq_acc * octx->n_threads);
725
2081
 
726
- if (octx->ctx->vtcm_size < total_spad) {
2082
+ if ((size_t) (vtcm_cur - octx->ctx->vtcm_base) > octx->ctx->vtcm_size) {
727
2083
  return HTP_STATUS_VTCM_TOO_SMALL;
728
2084
  }
729
2085
 
730
- octx->src0_spad.data = octx->ctx->vtcm_base; octx->src0_spad.src = NULL;
731
- octx->src1_spad.data = octx->src0_spad.data + octx->src0_spad.size; octx->src1_spad.src = NULL;
732
- octx->src2_spad.data = octx->src1_spad.data + octx->src1_spad.size; octx->src2_spad.src = NULL;
733
- octx->src3_spad.data = octx->src2_spad.data + octx->src2_spad.size; octx->src3_spad.src = NULL;
734
- octx->dst_spad.data = octx->src3_spad.data + octx->src3_spad.size; octx->dst_spad.src = NULL;
735
-
736
2086
  if (!(octx->flags & HTP_OPFLAGS_SKIP_COMPUTE)) {
737
2087
  worker_pool_run_func(octx->ctx->worker_pool, flash_attn_ext_f16_thread, &factx, octx->n_threads);
738
2088
  }