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
@@ -104,8 +104,8 @@ static __global__ void dequantize_block_q4_0(const void * __restrict__ vx, dst_t
104
104
  const uint8_t * q = x->qs + 4*il;
105
105
 
106
106
  for (int l = 0; l < 4; ++l) {
107
- y[l+ 0] = d * (q[l] & 0xF) + dm;
108
- y[l+16] = d * (q[l] >> 4) + dm;
107
+ y[l+ 0] = ggml_cuda_cast<dst_t>(d * (q[l] & 0xF) + dm);
108
+ y[l+16] = ggml_cuda_cast<dst_t>(d * (q[l] >> 4) + dm);
109
109
  }
110
110
  }
111
111
 
@@ -131,8 +131,8 @@ static __global__ void dequantize_block_q4_1(const void * __restrict__ vx, dst_t
131
131
  const uint8_t * q = x->qs + 4*il;
132
132
 
133
133
  for (int l = 0; l < 4; ++l) {
134
- y[l+ 0] = d.x * (q[l] & 0xF) + d.y;
135
- y[l+16] = d.x * (q[l] >> 4) + d.y;
134
+ y[l+ 0] = ggml_cuda_cast<dst_t>(d.x * (q[l] & 0xF) + d.y);
135
+ y[l+16] = ggml_cuda_cast<dst_t>(d.x * (q[l] >> 4) + d.y);
136
136
  }
137
137
  }
138
138
 
@@ -154,10 +154,10 @@ static __global__ void dequantize_block_q2_K(const void * __restrict__ vx, dst_t
154
154
 
155
155
  float dall = __low2half(x[i].dm);
156
156
  float dmin = __high2half(x[i].dm);
157
- y[l+ 0] = dall * (x[i].scales[is+0] & 0xF) * ((q >> 0) & 3) - dmin * (x[i].scales[is+0] >> 4);
158
- y[l+32] = dall * (x[i].scales[is+2] & 0xF) * ((q >> 2) & 3) - dmin * (x[i].scales[is+2] >> 4);
159
- y[l+64] = dall * (x[i].scales[is+4] & 0xF) * ((q >> 4) & 3) - dmin * (x[i].scales[is+4] >> 4);
160
- y[l+96] = dall * (x[i].scales[is+6] & 0xF) * ((q >> 6) & 3) - dmin * (x[i].scales[is+6] >> 4);
157
+ y[l+ 0] = ggml_cuda_cast<dst_t>(dall * (x[i].scales[is+0] & 0xF) * ((q >> 0) & 3) - dmin * (x[i].scales[is+0] >> 4));
158
+ y[l+32] = ggml_cuda_cast<dst_t>(dall * (x[i].scales[is+2] & 0xF) * ((q >> 2) & 3) - dmin * (x[i].scales[is+2] >> 4));
159
+ y[l+64] = ggml_cuda_cast<dst_t>(dall * (x[i].scales[is+4] & 0xF) * ((q >> 4) & 3) - dmin * (x[i].scales[is+4] >> 4));
160
+ y[l+96] = ggml_cuda_cast<dst_t>(dall * (x[i].scales[is+6] & 0xF) * ((q >> 6) & 3) - dmin * (x[i].scales[is+6] >> 4));
161
161
  }
162
162
 
163
163
  template<typename dst_t>
@@ -188,7 +188,9 @@ static __global__ void dequantize_block_q3_K(const void * __restrict__ vx, dst_t
188
188
  const uint8_t * q = x[i].qs + 32*n;
189
189
  const uint8_t * hm = x[i].hmask;
190
190
 
191
- for (int l = l0; l < l0+4; ++l) y[l] = dl * ((int8_t)((q[l] >> shift) & 3) - ((hm[l] & m) ? 0 : 4));
191
+ for (int l = l0; l < l0+4; ++l) {
192
+ y[l] = ggml_cuda_cast<dst_t>(dl * ((int8_t)((q[l] >> shift) & 3) - ((hm[l] & m) ? 0 : 4)));
193
+ }
192
194
  }
193
195
 
194
196
  static inline __device__ void get_scale_min_k4(int j, const uint8_t * q, uint8_t & d, uint8_t & m) {
@@ -226,8 +228,8 @@ static __global__ void dequantize_block_q4_K(const void * __restrict__ vx, dst_t
226
228
  get_scale_min_k4(is + 1, x[i].scales, sc, m);
227
229
  const float d2 = dall * sc; const float m2 = dmin * m;
228
230
  for (int l = 0; l < n; ++l) {
229
- y[l + 0] = d1 * (q[l] & 0xF) - m1;
230
- y[l +32] = d2 * (q[l] >> 4) - m2;
231
+ y[l + 0] = ggml_cuda_cast<dst_t>(d1 * (q[l] & 0xF) - m1);
232
+ y[l +32] = ggml_cuda_cast<dst_t>(d2 * (q[l] >> 4) - m2);
231
233
  }
232
234
  }
233
235
 
@@ -258,11 +260,11 @@ static __global__ void dequantize_block_q5_K(const void * __restrict__ vx, dst_t
258
260
  const float d2 = dall * sc; const float m2 = dmin * m;
259
261
 
260
262
  uint8_t hm = 1 << (2*il);
261
- y[ 0] = d1 * ((ql[ 0] & 0xF) + (qh[ 0] & hm ? 16 : 0)) - m1;
262
- y[ 1] = d1 * ((ql[ 1] & 0xF) + (qh[ 1] & hm ? 16 : 0)) - m1;
263
+ y[ 0] = ggml_cuda_cast<dst_t>(d1 * ((ql[ 0] & 0xF) + (qh[ 0] & hm ? 16 : 0)) - m1);
264
+ y[ 1] = ggml_cuda_cast<dst_t>(d1 * ((ql[ 1] & 0xF) + (qh[ 1] & hm ? 16 : 0)) - m1);
263
265
  hm <<= 1;
264
- y[32] = d2 * ((ql[ 0] >> 4) + (qh[ 0] & hm ? 16 : 0)) - m2;
265
- y[33] = d2 * ((ql[ 1] >> 4) + (qh[ 1] & hm ? 16 : 0)) - m2;
266
+ y[32] = ggml_cuda_cast<dst_t>(d2 * ((ql[ 0] >> 4) + (qh[ 0] & hm ? 16 : 0)) - m2);
267
+ y[33] = ggml_cuda_cast<dst_t>(d2 * ((ql[ 1] >> 4) + (qh[ 1] & hm ? 16 : 0)) - m2);
266
268
  }
267
269
 
268
270
  template<typename dst_t>
@@ -285,10 +287,10 @@ static __global__ void dequantize_block_q6_K(const void * __restrict__ vx, dst_t
285
287
  const uint8_t qh = x[i].qh[32*ip + il];
286
288
  const int8_t * sc = x[i].scales + is;
287
289
 
288
- y[ 0] = d * sc[0] * ((int8_t)((ql[ 0] & 0xF) | (((qh >> 0) & 3) << 4)) - 32);
289
- y[32] = d * sc[2] * ((int8_t)((ql[32] & 0xF) | (((qh >> 2) & 3) << 4)) - 32);
290
- y[64] = d * sc[4] * ((int8_t)((ql[ 0] >> 4) | (((qh >> 4) & 3) << 4)) - 32);
291
- y[96] = d * sc[6] * ((int8_t)((ql[32] >> 4) | (((qh >> 6) & 3) << 4)) - 32);
290
+ y[ 0] = ggml_cuda_cast<dst_t>(d * sc[0] * ((int8_t)((ql[ 0] & 0xF) | (((qh >> 0) & 3) << 4)) - 32));
291
+ y[32] = ggml_cuda_cast<dst_t>(d * sc[2] * ((int8_t)((ql[32] & 0xF) | (((qh >> 2) & 3) << 4)) - 32));
292
+ y[64] = ggml_cuda_cast<dst_t>(d * sc[4] * ((int8_t)((ql[ 0] >> 4) | (((qh >> 4) & 3) << 4)) - 32));
293
+ y[96] = ggml_cuda_cast<dst_t>(d * sc[6] * ((int8_t)((ql[32] >> 4) | (((qh >> 6) & 3) << 4)) - 32));
292
294
  }
293
295
 
294
296
  template<typename dst_t>
@@ -307,7 +309,9 @@ static __global__ void dequantize_block_iq2_xxs(const void * __restrict__ vx, ds
307
309
  const uint32_t aux32 = q2[2] | (q2[3] << 16);
308
310
  const float d = (float)x[i].d * (0.5f + (aux32 >> 28)) * 0.25f;
309
311
  const uint8_t signs = ksigns_iq2xs[(aux32 >> 7*il) & 127];
310
- for (int j = 0; j < 8; ++j) y[j] = d * grid[j] * (signs & kmask_iq2xs[j] ? -1.f : 1.f);
312
+ for (int j = 0; j < 8; ++j) {
313
+ y[j] = ggml_cuda_cast<dst_t>(d * grid[j] * (signs & kmask_iq2xs[j] ? -1.f : 1.f));
314
+ }
311
315
  }
312
316
 
313
317
  template<typename dst_t>
@@ -324,7 +328,9 @@ static __global__ void dequantize_block_iq2_xs(const void * __restrict__ vx, dst
324
328
  const uint8_t * grid = (const uint8_t *)(iq2xs_grid + (q2[il] & 511));
325
329
  const float d = (float)x[i].d * (0.5f + ((x[i].scales[ib] >> 4*(il/2)) & 0xf)) * 0.25f;
326
330
  const uint8_t signs = ksigns_iq2xs[q2[il] >> 9];
327
- for (int j = 0; j < 8; ++j) y[j] = d * grid[j] * (signs & kmask_iq2xs[j] ? -1.f : 1.f);
331
+ for (int j = 0; j < 8; ++j) {
332
+ y[j] = ggml_cuda_cast<dst_t>(d * grid[j] * (signs & kmask_iq2xs[j] ? -1.f : 1.f));
333
+ }
328
334
  }
329
335
 
330
336
  template<typename dst_t>
@@ -340,7 +346,9 @@ static __global__ void dequantize_block_iq2_s(const void * __restrict__ vx, dst_
340
346
  const uint8_t * grid = (const uint8_t *)(iq2s_grid + (x[i].qs[4*ib+il] | ((x[i].qh[ib] << (8-2*il)) & 0x300)));
341
347
  const float d = (float)x[i].d * (0.5f + ((x[i].scales[ib] >> 4*(il/2)) & 0xf)) * 0.25f;
342
348
  const uint8_t signs = x[i].qs[QK_K/8+4*ib+il];
343
- for (int j = 0; j < 8; ++j) y[j] = d * grid[j] * (signs & kmask_iq2xs[j] ? -1.f : 1.f);
349
+ for (int j = 0; j < 8; ++j) {
350
+ y[j] = ggml_cuda_cast<dst_t>(d * grid[j] * (signs & kmask_iq2xs[j] ? -1.f : 1.f));
351
+ }
344
352
  }
345
353
 
346
354
  template<typename dst_t>
@@ -361,8 +369,8 @@ static __global__ void dequantize_block_iq3_xxs(const void * __restrict__ vx, ds
361
369
  const float d = (float)x[i].d * (0.5f + (aux32 >> 28)) * 0.5f;
362
370
  const uint8_t signs = ksigns_iq2xs[(aux32 >> 7*il) & 127];
363
371
  for (int j = 0; j < 4; ++j) {
364
- y[j+0] = d * grid1[j] * (signs & kmask_iq2xs[j+0] ? -1.f : 1.f);
365
- y[j+4] = d * grid2[j] * (signs & kmask_iq2xs[j+4] ? -1.f : 1.f);
372
+ y[j+0] = ggml_cuda_cast<dst_t>(d * grid1[j] * (signs & kmask_iq2xs[j+0] ? -1.f : 1.f));
373
+ y[j+4] = ggml_cuda_cast<dst_t>(d * grid2[j] * (signs & kmask_iq2xs[j+4] ? -1.f : 1.f));
366
374
  }
367
375
  }
368
376
 
@@ -382,8 +390,8 @@ static __global__ void dequantize_block_iq3_s(const void * __restrict__ vx, dst_
382
390
  const float d = (float)x[i].d * (1 + 2*((x[i].scales[ib/2] >> 4*(ib%2)) & 0xf));
383
391
  const uint8_t signs = x[i].signs[4*ib + il];
384
392
  for (int j = 0; j < 4; ++j) {
385
- y[j+0] = d * grid1[j] * (signs & kmask_iq2xs[j+0] ? -1.f : 1.f);
386
- y[j+4] = d * grid2[j] * (signs & kmask_iq2xs[j+4] ? -1.f : 1.f);
393
+ y[j+0] = ggml_cuda_cast<dst_t>(d * grid1[j] * (signs & kmask_iq2xs[j+0] ? -1.f : 1.f));
394
+ y[j+4] = ggml_cuda_cast<dst_t>(d * grid2[j] * (signs & kmask_iq2xs[j+4] ? -1.f : 1.f));
387
395
  }
388
396
  }
389
397
 
@@ -404,7 +412,7 @@ static __global__ void dequantize_block_iq1_s(const void * __restrict__ vx, dst_
404
412
  grid32[1] = (grid32[0] >> 4) & 0x0f0f0f0f;
405
413
  grid32[0] &= 0x0f0f0f0f;
406
414
  for (int j = 0; j < 8; ++j) {
407
- y[j] = d * (q[j] + delta);
415
+ y[j] = ggml_cuda_cast<dst_t>(d * (q[j] + delta));
408
416
  }
409
417
  }
410
418
 
@@ -429,7 +437,7 @@ static __global__ void dequantize_block_iq1_m(const void * __restrict__ vx, dst_
429
437
  grid32[1] = (grid32[0] >> 4) & 0x0f0f0f0f;
430
438
  grid32[0] &= 0x0f0f0f0f;
431
439
  for (int j = 0; j < 8; ++j) {
432
- y[j] = d * (q[j] + delta);
440
+ y[j] = ggml_cuda_cast<dst_t>(d * (q[j] + delta));
433
441
  }
434
442
  }
435
443
 
@@ -446,8 +454,8 @@ static __global__ void dequantize_block_iq4_nl(const void * __restrict__ vx, dst
446
454
  const uint8_t * q4 = x[ib].qs + 4*il;
447
455
  const float d = (float)x[ib].d;
448
456
  for (int j = 0; j < 4; ++j) {
449
- y[j+ 0] = d * kvalues_iq4nl[q4[j] & 0xf];
450
- y[j+16] = d * kvalues_iq4nl[q4[j] >> 4];
457
+ y[j+ 0] = ggml_cuda_cast<dst_t>(d * kvalues_iq4nl[q4[j] & 0xf]);
458
+ y[j+16] = ggml_cuda_cast<dst_t>(d * kvalues_iq4nl[q4[j] >> 4]);
451
459
  }
452
460
  }
453
461
 
@@ -463,8 +471,8 @@ static __global__ void dequantize_block_iq4_xs(const void * __restrict__ vx, dst
463
471
  const uint8_t * q4 = x[i].qs + 16*ib + 4*il;
464
472
  const float d = (float)x[i].d * ((((x[i].scales_l[ib/2] >> 4*(ib%2)) & 0xf) | (((x[i].scales_h >> 2*ib) & 3) << 4)) - 32);
465
473
  for (int j = 0; j < 4; ++j) {
466
- y[j+ 0] = d * kvalues_iq4nl[q4[j] & 0xf];
467
- y[j+16] = d * kvalues_iq4nl[q4[j] >> 4];
474
+ y[j+ 0] = ggml_cuda_cast<dst_t>(d * kvalues_iq4nl[q4[j] & 0xf]);
475
+ y[j+16] = ggml_cuda_cast<dst_t>(d * kvalues_iq4nl[q4[j] >> 4]);
468
476
  }
469
477
  }
470
478
 
@@ -481,8 +489,8 @@ static __global__ void dequantize_block_mxfp4(const void * __restrict__ vx, dst_
481
489
  const uint8_t * q4 = x[ib].qs + 4*il;
482
490
  const float d = ggml_cuda_e8m0_to_fp32(x[ib].e);
483
491
  for (int j = 0; j < 4; ++j) {
484
- y[j+ 0] = d * kvalues_mxfp4[q4[j] & 0xf]*0.5f;
485
- y[j+16] = d * kvalues_mxfp4[q4[j] >> 4]*0.5f;
492
+ y[j+ 0] = ggml_cuda_cast<dst_t>(d * kvalues_mxfp4[q4[j] & 0xf]*0.5f);
493
+ y[j+16] = ggml_cuda_cast<dst_t>(d * kvalues_mxfp4[q4[j] >> 4]*0.5f);
486
494
  }
487
495
  }
488
496
 
@@ -700,6 +708,50 @@ static void convert_unary_cont_cuda(const void * vx, dst_t * y, const int64_t k,
700
708
 
701
709
  to_bf16_cuda_t ggml_get_to_bf16_cuda(ggml_type type) {
702
710
  switch (type) {
711
+ case GGML_TYPE_Q1_0:
712
+ return dequantize_block_cont_cuda<QK1_0, QR1_0, dequantize_q1_0>;
713
+ case GGML_TYPE_Q4_0:
714
+ return dequantize_row_q4_0_cuda;
715
+ case GGML_TYPE_Q4_1:
716
+ return dequantize_row_q4_1_cuda;
717
+ case GGML_TYPE_Q5_0:
718
+ return dequantize_block_cont_cuda<QK5_0, QR5_0, dequantize_q5_0>;
719
+ case GGML_TYPE_Q5_1:
720
+ return dequantize_block_cont_cuda<QK5_1, QR5_1, dequantize_q5_1>;
721
+ case GGML_TYPE_Q8_0:
722
+ return dequantize_block_cont_cuda<QK8_0, QR8_0, dequantize_q8_0>;
723
+ case GGML_TYPE_Q2_K:
724
+ return dequantize_row_q2_K_cuda;
725
+ case GGML_TYPE_Q3_K:
726
+ return dequantize_row_q3_K_cuda;
727
+ case GGML_TYPE_Q4_K:
728
+ return dequantize_row_q4_K_cuda;
729
+ case GGML_TYPE_Q5_K:
730
+ return dequantize_row_q5_K_cuda;
731
+ case GGML_TYPE_Q6_K:
732
+ return dequantize_row_q6_K_cuda;
733
+ case GGML_TYPE_IQ2_XXS:
734
+ return dequantize_row_iq2_xxs_cuda;
735
+ case GGML_TYPE_IQ2_XS:
736
+ return dequantize_row_iq2_xs_cuda;
737
+ case GGML_TYPE_IQ2_S:
738
+ return dequantize_row_iq2_s_cuda;
739
+ case GGML_TYPE_IQ3_XXS:
740
+ return dequantize_row_iq3_xxs_cuda;
741
+ case GGML_TYPE_IQ1_S:
742
+ return dequantize_row_iq1_s_cuda;
743
+ case GGML_TYPE_IQ1_M:
744
+ return dequantize_row_iq1_m_cuda;
745
+ case GGML_TYPE_IQ4_NL:
746
+ return dequantize_row_iq4_nl_cuda;
747
+ case GGML_TYPE_IQ4_XS:
748
+ return dequantize_row_iq4_xs_cuda;
749
+ case GGML_TYPE_IQ3_S:
750
+ return dequantize_row_iq3_s_cuda;
751
+ case GGML_TYPE_MXFP4:
752
+ return dequantize_row_mxfp4_cuda;
753
+ case GGML_TYPE_NVFP4:
754
+ return dequantize_row_nvfp4_cuda;
703
755
  case GGML_TYPE_F32:
704
756
  return convert_unary_cont_cuda<float>;
705
757
  case GGML_TYPE_F16:
@@ -53,10 +53,10 @@ static __global__ void cpy_scalar_transpose(const char * cx, char * cdst, const
53
53
  const int64_t nmat = ne / (ne00 * ne01);
54
54
  const int64_t n = ne00 * ne01;
55
55
 
56
- const int x = blockIdx.x * CUDA_CPY_TILE_DIM_2D + threadIdx.x;
57
- const int y = blockIdx.y * CUDA_CPY_TILE_DIM_2D + threadIdx.y;
58
- const int tx = blockIdx.y * CUDA_CPY_TILE_DIM_2D + threadIdx.x; // transpose block offset
59
- const int ty = blockIdx.x * CUDA_CPY_TILE_DIM_2D + threadIdx.y;
56
+ const int64_t x = (int64_t) blockIdx.x * CUDA_CPY_TILE_DIM_2D + threadIdx.x;
57
+ const int64_t y = (int64_t) blockIdx.y * CUDA_CPY_TILE_DIM_2D + threadIdx.y;
58
+ const int64_t tx = (int64_t) blockIdx.y * CUDA_CPY_TILE_DIM_2D + threadIdx.x; // transpose block offset
59
+ const int64_t ty = (int64_t) blockIdx.x * CUDA_CPY_TILE_DIM_2D + threadIdx.y;
60
60
 
61
61
  __shared__ float tile[2][CUDA_CPY_TILE_DIM_2D][CUDA_CPY_TILE_DIM_2D+1];
62
62
  int cur_tile_buf = 0;
@@ -197,7 +197,7 @@ static void ggml_cpy_scalar_contiguous_cuda(
197
197
  cudaStream_t stream) {
198
198
 
199
199
  const int64_t num_blocks = (ne + CUDA_CPY_BLOCK_SIZE - 1) / CUDA_CPY_BLOCK_SIZE;
200
- GGML_ASSERT(num_blocks < UINT_MAX);
200
+ GGML_ASSERT(num_blocks <= INT_MAX);
201
201
  const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params((dim3)num_blocks, CUDA_CPY_BLOCK_SIZE, 0, stream);
202
202
  ggml_cuda_kernel_launch(cpy_scalar_contiguous<src_t, dst_t>, launch_params, cx, cdst, ne);
203
203
  }
@@ -208,6 +208,14 @@ static void ggml_cpy_scalar_cuda(
208
208
  const int64_t ne00, const int64_t ne01, const int64_t ne02, const int64_t nb00, const int64_t nb01, const int64_t nb02,
209
209
  const int64_t nb03, const int64_t ne10, const int64_t ne11, const int64_t ne12, const int64_t nb10, const int64_t nb11, const int64_t nb12, const int64_t nb13, cudaStream_t stream) {
210
210
 
211
+ const auto launch_scalar_generic = [&]() {
212
+ const int64_t num_blocks = (ne + CUDA_CPY_BLOCK_SIZE - 1) / CUDA_CPY_BLOCK_SIZE;
213
+ GGML_ASSERT(num_blocks <= INT_MAX);
214
+ const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params((dim3)num_blocks, CUDA_CPY_BLOCK_SIZE, 0, stream);
215
+ ggml_cuda_kernel_launch(cpy_scalar<cpy_1_scalar<src_t, dst_t>>, launch_params,
216
+ cx, cdst, ne, ne00, ne01, ne02, nb00, nb01, nb02, nb03, ne10, ne11, ne12, nb10, nb11, nb12, nb13);
217
+ };
218
+
211
219
  if (transposed) {
212
220
  GGML_ASSERT(ne == ne00*ne01*ne02); // ne[3] is 1 assumed
213
221
  int64_t ne00n, ne01n, ne02n;
@@ -224,20 +232,18 @@ static void ggml_cpy_scalar_cuda(
224
232
  int64_t grid_x = (ne01n + CUDA_CPY_TILE_DIM_2D - 1) / CUDA_CPY_TILE_DIM_2D;
225
233
  int64_t grid_y = (ne00n + CUDA_CPY_TILE_DIM_2D - 1) / CUDA_CPY_TILE_DIM_2D;
226
234
  int64_t grid_z = (ne/(ne01n*ne00n) + CUDA_CPY_BLOCK_NM - 1) / CUDA_CPY_BLOCK_NM;
227
- GGML_ASSERT(grid_x < UINT_MAX);
228
- GGML_ASSERT(grid_y < USHRT_MAX);
229
- GGML_ASSERT(grid_z < USHRT_MAX);
230
- dim3 dimGrid(grid_x, grid_y, grid_z);
231
- dim3 dimBlock(CUDA_CPY_TILE_DIM_2D, CUDA_CPY_BLOCK_ROWS, 1);
232
- const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(dimGrid, dimBlock, 0, stream);
233
- ggml_cuda_kernel_launch(cpy_scalar_transpose<dst_t>, launch_params,
234
- cx, cdst, ne, ne00n, ne01n, ne02n, nb00, nb01, nb02, nb03, ne10, ne11, ne12, nb10, nb11, nb12, nb13);
235
+ GGML_ASSERT(grid_x <= INT_MAX);
236
+ if (grid_y > USHRT_MAX || grid_z > USHRT_MAX) {
237
+ launch_scalar_generic();
238
+ } else {
239
+ dim3 dimGrid(grid_x, grid_y, grid_z);
240
+ dim3 dimBlock(CUDA_CPY_TILE_DIM_2D, CUDA_CPY_BLOCK_ROWS, 1);
241
+ const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(dimGrid, dimBlock, 0, stream);
242
+ ggml_cuda_kernel_launch(cpy_scalar_transpose<dst_t>, launch_params,
243
+ cx, cdst, ne, ne00n, ne01n, ne02n, nb00, nb01, nb02, nb03, ne10, ne11, ne12, nb10, nb11, nb12, nb13);
244
+ }
235
245
  } else {
236
- const int64_t num_blocks = (ne + CUDA_CPY_BLOCK_SIZE - 1) / CUDA_CPY_BLOCK_SIZE;
237
- GGML_ASSERT(num_blocks < UINT_MAX);
238
- const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params((dim3)num_blocks, CUDA_CPY_BLOCK_SIZE, 0, stream);
239
- ggml_cuda_kernel_launch(cpy_scalar<cpy_1_scalar<src_t, dst_t>>, launch_params,
240
- cx, cdst, ne, ne00, ne01, ne02, nb00, nb01, nb02, nb03, ne10, ne11, ne12, nb10, nb11, nb12, nb13);
246
+ launch_scalar_generic();
241
247
  }
242
248
  }
243
249
 
@@ -248,7 +254,7 @@ static void ggml_cpy_f32_q8_0_cuda(
248
254
 
249
255
  GGML_ASSERT(ne % QK8_0 == 0);
250
256
  const int64_t num_blocks = ne / QK8_0;
251
- GGML_ASSERT(num_blocks < UINT_MAX);
257
+ GGML_ASSERT(num_blocks <= INT_MAX);
252
258
  cpy_f32_q<cpy_blck_f32_q8_0, QK8_0><<<num_blocks, 1, 0, stream>>>
253
259
  (cx, cdst, ne, ne00, ne01, ne02, nb00, nb01, nb02, nb03, ne10, ne11, ne12, nb10, nb11, nb12, nb13);
254
260
  }
@@ -259,7 +265,7 @@ static void ggml_cpy_q8_0_f32_cuda(
259
265
  const int64_t nb03, const int64_t ne10, const int64_t ne11, const int64_t ne12, const int64_t nb10, const int64_t nb11, const int64_t nb12, const int64_t nb13, cudaStream_t stream) {
260
266
 
261
267
  const int64_t num_blocks = ne;
262
- GGML_ASSERT(num_blocks < UINT_MAX);
268
+ GGML_ASSERT(num_blocks <= INT_MAX);
263
269
  cpy_q_f32<cpy_blck_q8_0_f32, QK8_0><<<num_blocks, 1, 0, stream>>>
264
270
  (cx, cdst, ne, ne00, ne01, ne02, nb00, nb01, nb02, nb03, ne10, ne11, ne12, nb10, nb11, nb12, nb13);
265
271
  }
@@ -271,7 +277,7 @@ static void ggml_cpy_f32_q4_0_cuda(
271
277
 
272
278
  GGML_ASSERT(ne % QK4_0 == 0);
273
279
  const int64_t num_blocks = ne / QK4_0;
274
- GGML_ASSERT(num_blocks < UINT_MAX);
280
+ GGML_ASSERT(num_blocks <= INT_MAX);
275
281
  cpy_f32_q<cpy_blck_f32_q4_0, QK4_0><<<num_blocks, 1, 0, stream>>>
276
282
  (cx, cdst, ne, ne00, ne01, ne02, nb00, nb01, nb02, nb03, ne10, ne11, ne12, nb10, nb11, nb12, nb13);
277
283
  }
@@ -284,7 +290,7 @@ static void ggml_cpy_q4_0_f32_cuda(
284
290
  const int64_t nb10, const int64_t nb11, const int64_t nb12, const int64_t nb13,
285
291
  cudaStream_t stream) {
286
292
  const int64_t num_blocks = ne;
287
- GGML_ASSERT(num_blocks < UINT_MAX);
293
+ GGML_ASSERT(num_blocks <= INT_MAX);
288
294
  cpy_q_f32<cpy_blck_q_f32<dequantize_q4_0, QK4_0>, QK4_0><<<num_blocks, 1, 0, stream>>>(
289
295
  cx, cdst, ne, ne00, ne01, ne02, nb00, nb01, nb02, nb03,
290
296
  ne10, ne11, ne12, nb10, nb11, nb12, nb13);
@@ -297,7 +303,7 @@ static void ggml_cpy_f32_q4_1_cuda(
297
303
 
298
304
  GGML_ASSERT(ne % QK4_1 == 0);
299
305
  const int64_t num_blocks = ne / QK4_1;
300
- GGML_ASSERT(num_blocks < UINT_MAX);
306
+ GGML_ASSERT(num_blocks <= INT_MAX);
301
307
  cpy_f32_q<cpy_blck_f32_q4_1, QK4_1><<<num_blocks, 1, 0, stream>>>
302
308
  (cx, cdst, ne, ne00, ne01, ne02, nb00, nb01, nb02, nb03, ne10, ne11, ne12, nb10, nb11, nb12, nb13);
303
309
  }
@@ -310,7 +316,7 @@ static void ggml_cpy_q4_1_f32_cuda(
310
316
  const int64_t nb10, const int64_t nb11, const int64_t nb12, const int64_t nb13,
311
317
  cudaStream_t stream) {
312
318
  const int64_t num_blocks = ne;
313
- GGML_ASSERT(num_blocks < UINT_MAX);
319
+ GGML_ASSERT(num_blocks <= INT_MAX);
314
320
  cpy_q_f32<cpy_blck_q_f32<dequantize_q4_1, QK4_1>, QK4_1><<<num_blocks, 1, 0, stream>>>(
315
321
  cx, cdst, ne, ne00, ne01, ne02, nb00, nb01, nb02, nb03,
316
322
  ne10, ne11, ne12, nb10, nb11, nb12, nb13);
@@ -323,7 +329,7 @@ static void ggml_cpy_f32_q5_0_cuda(
323
329
 
324
330
  GGML_ASSERT(ne % QK5_0 == 0);
325
331
  const int64_t num_blocks = ne / QK5_0;
326
- GGML_ASSERT(num_blocks < UINT_MAX);
332
+ GGML_ASSERT(num_blocks <= INT_MAX);
327
333
  cpy_f32_q<cpy_blck_f32_q5_0, QK5_0><<<num_blocks, 1, 0, stream>>>
328
334
  (cx, cdst, ne, ne00, ne01, ne02, nb00, nb01, nb02, nb03, ne10, ne11, ne12, nb10, nb11, nb12, nb13);
329
335
  }
@@ -336,7 +342,7 @@ static void ggml_cpy_q5_0_f32_cuda(
336
342
  const int64_t nb10, const int64_t nb11, const int64_t nb12, const int64_t nb13,
337
343
  cudaStream_t stream) {
338
344
  const int64_t num_blocks = ne;
339
- GGML_ASSERT(num_blocks < UINT_MAX);
345
+ GGML_ASSERT(num_blocks <= INT_MAX);
340
346
  cpy_q_f32<cpy_blck_q_f32<dequantize_q5_0, QK5_0>, QK5_0><<<num_blocks, 1, 0, stream>>>(
341
347
  cx, cdst, ne, ne00, ne01, ne02, nb00, nb01, nb02, nb03,
342
348
  ne10, ne11, ne12, nb10, nb11, nb12, nb13);
@@ -349,7 +355,7 @@ static void ggml_cpy_f32_q5_1_cuda(
349
355
 
350
356
  GGML_ASSERT(ne % QK5_1 == 0);
351
357
  const int64_t num_blocks = ne / QK5_1;
352
- GGML_ASSERT(num_blocks < UINT_MAX);
358
+ GGML_ASSERT(num_blocks <= INT_MAX);
353
359
  cpy_f32_q<cpy_blck_f32_q5_1, QK5_1><<<num_blocks, 1, 0, stream>>>
354
360
  (cx, cdst, ne, ne00, ne01, ne02, nb00, nb01, nb02, nb03, ne10, ne11, ne12, nb10, nb11, nb12, nb13);
355
361
  }
@@ -362,7 +368,7 @@ static void ggml_cpy_q5_1_f32_cuda(
362
368
  const int64_t nb10, const int64_t nb11, const int64_t nb12, const int64_t nb13,
363
369
  cudaStream_t stream) {
364
370
  const int64_t num_blocks = ne;
365
- GGML_ASSERT(num_blocks < UINT_MAX);
371
+ GGML_ASSERT(num_blocks <= INT_MAX);
366
372
  cpy_q_f32<cpy_blck_q_f32<dequantize_q5_1, QK5_1>, QK5_1><<<num_blocks, 1, 0, stream>>>(
367
373
  cx, cdst, ne, ne00, ne01, ne02, nb00, nb01, nb02, nb03,
368
374
  ne10, ne11, ne12, nb10, nb11, nb12, nb13);
@@ -375,11 +381,51 @@ static void ggml_cpy_f32_iq4_nl_cuda(
375
381
 
376
382
  GGML_ASSERT(ne % QK4_NL == 0);
377
383
  const int64_t num_blocks = ne / QK4_NL;
378
- GGML_ASSERT(num_blocks < UINT_MAX);
384
+ GGML_ASSERT(num_blocks <= INT_MAX);
379
385
  cpy_f32_q<cpy_blck_f32_iq4_nl, QK4_NL><<<num_blocks, 1, 0, stream>>>
380
386
  (cx, cdst, ne, ne00, ne01, ne02, nb00, nb01, nb02, nb03, ne10, ne11, ne12, nb10, nb11, nb12, nb13);
381
387
  }
382
388
 
389
+ // check if a same-type copy reduces to a 2D strided copy (height rows of width
390
+ // contiguous bytes), so it can use cudaMemcpy2DAsync instead of the scalar kernel
391
+ static bool ggml_cuda_cpy_as_memcpy_2d(const ggml_tensor * src0, const ggml_tensor * src1,
392
+ size_t & width, size_t & height, size_t & spitch, size_t & dpitch) {
393
+ // require matching shape: a reshaped copy maps elements by flat order, which the
394
+ // prefix walk below does not handle
395
+ if (src0->type != src1->type || !ggml_are_same_shape(src0, src1)) {
396
+ return false;
397
+ }
398
+
399
+ // grow the contiguous prefix block shared by both tensors
400
+ size_t block_nb = ggml_element_size(src0);
401
+ int d = 0;
402
+ for (; d < GGML_MAX_DIMS; ++d) {
403
+ if (src0->nb[d] != block_nb || src1->nb[d] != block_nb) {
404
+ break;
405
+ }
406
+ block_nb *= src0->ne[d];
407
+ }
408
+
409
+ // d == 0: nothing contiguous; d == GGML_MAX_DIMS: fully contiguous (handled by memcpy)
410
+ if (d == 0 || d == GGML_MAX_DIMS) {
411
+ return false;
412
+ }
413
+
414
+ // dim d carries the rows; everything above it must be a single element
415
+ for (int i = d + 1; i < GGML_MAX_DIMS; ++i) {
416
+ if (src0->ne[i] != 1) {
417
+ return false;
418
+ }
419
+ }
420
+
421
+ width = block_nb;
422
+ height = src0->ne[d];
423
+ spitch = src0->nb[d];
424
+ dpitch = src1->nb[d];
425
+
426
+ return spitch >= width && dpitch >= width;
427
+ }
428
+
383
429
  void ggml_cuda_cpy(ggml_backend_cuda_context & ctx, const ggml_tensor * src0, ggml_tensor * src1) {
384
430
  const int64_t ne = ggml_nelements(src0);
385
431
  GGML_ASSERT(ne == ggml_nelements(src1));
@@ -415,6 +461,8 @@ void ggml_cuda_cpy(ggml_backend_cuda_context & ctx, const ggml_tensor * src0, gg
415
461
  const bool can_be_transposed = nb01 == (int64_t)ggml_element_size(src0) &&
416
462
  src0->ne[3] == 1 && nb02 == ne00 * ne01 * (int64_t)ggml_element_size(src0);
417
463
 
464
+ size_t mc_width = 0, mc_height = 0, mc_spitch = 0, mc_dpitch = 0;
465
+
418
466
  if (src0->type == src1->type && contiguous_srcs) {
419
467
  GGML_ASSERT(ggml_nbytes(src0) == ggml_nbytes(src1));
420
468
  #if defined(GGML_USE_MUSA) && defined(GGML_MUSA_MUDNN_COPY)
@@ -425,6 +473,9 @@ void ggml_cuda_cpy(ggml_backend_cuda_context & ctx, const ggml_tensor * src0, gg
425
473
  {
426
474
  CUDA_CHECK(cudaMemcpyAsync(src1_ddc, src0_ddc, ggml_nbytes(src0), cudaMemcpyDeviceToDevice, main_stream));
427
475
  }
476
+ } else if (ggml_cuda_cpy_as_memcpy_2d(src0, src1, mc_width, mc_height, mc_spitch, mc_dpitch)) {
477
+ CUDA_CHECK(cudaMemcpy2DAsync(src1_ddc, mc_dpitch, src0_ddc, mc_spitch,
478
+ mc_width, mc_height, cudaMemcpyDeviceToDevice, main_stream));
428
479
  } else if (src0->type == GGML_TYPE_F32 && src1->type == GGML_TYPE_F32) {
429
480
  if (can_be_transposed) {
430
481
  ggml_cpy_scalar_cuda<float, float, true>
@@ -664,7 +664,10 @@ constexpr __device__ dequantize_V_t get_dequantize_V() {
664
664
  template <int ncols1>
665
665
  __launch_bounds__(FATTN_KQ_STRIDE/2, 1)
666
666
  static __global__ void flash_attn_mask_to_KV_max(
667
- const half2 * __restrict__ mask, int * __restrict__ KV_max, const int ne30, const int s31, const int s33) {
667
+ const half2 * mask_ptr, int * KV_max_ptr, const int ne30, const int64_t s31, const int64_t s33) {
668
+ const half2 * GGML_CUDA_RESTRICT mask = mask_ptr;
669
+ int * GGML_CUDA_RESTRICT KV_max = KV_max_ptr;
670
+
668
671
  const int ne31 = gridDim.x;
669
672
  const int tid = threadIdx.x;
670
673
  const int sequence = blockIdx.y;
@@ -1089,8 +1092,8 @@ void launch_fattn(
1089
1092
  // Only worth the overhead if there is at lease one FATTN_KQ_STRIDE x FATTN_KQ_STRIDE square to be skipped or
1090
1093
  // multiple sequences of possibly different lengths.
1091
1094
  if (mask && K->ne[1] % FATTN_KQ_STRIDE == 0 && (Q->ne[1] >= 1024 || Q->ne[3] > 1)) {
1092
- const int s31 = mask->nb[1] / sizeof(half2);
1093
- const int s33 = mask->nb[3] / sizeof(half2);
1095
+ const int64_t s31 = mask->nb[1] / sizeof(half2);
1096
+ const int64_t s33 = mask->nb[3] / sizeof(half2);
1094
1097
 
1095
1098
  const dim3 blocks_num_KV_max(ntiles_x, Q->ne[3], 1);
1096
1099
  const dim3 block_dim_KV_max(FATTN_KQ_STRIDE/2, 1, 1);
@@ -1099,8 +1102,9 @@ void launch_fattn(
1099
1102
  const int iter_k = K->ne[1] / FATTN_KQ_STRIDE;
1100
1103
 
1101
1104
  KV_max.alloc(ne_KV_max);
1102
- flash_attn_mask_to_KV_max<ncols1><<<blocks_num_KV_max, block_dim_KV_max, 0, main_stream>>>
1103
- ((const half2 *) mask->data, KV_max.ptr, iter_k, s31, s33);
1105
+ ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(blocks_num_KV_max, block_dim_KV_max, 0, main_stream);
1106
+ ggml_cuda_kernel_launch(flash_attn_mask_to_KV_max<ncols1>, launch_params,
1107
+ (const half2 *) mask->data, KV_max.ptr, iter_k, s31, s33);
1104
1108
  CUDA_CHECK(cudaGetLastError());
1105
1109
  }
1106
1110
 
@@ -2003,6 +2003,10 @@ DECL_FATTN_MMA_F16_CASE_ALL_NCOLS2(112, 112, 64)
2003
2003
  DECL_FATTN_MMA_F16_CASE_ALL_NCOLS2(128, 128, 64)
2004
2004
  DECL_FATTN_MMA_F16_CASE_ALL_NCOLS2(256, 256, 64)
2005
2005
 
2006
+ extern DECL_FATTN_MMA_F16_CASE(512, 512, 4, 2);
2007
+ extern DECL_FATTN_MMA_F16_CASE(512, 512, 8, 2);
2008
+ extern DECL_FATTN_MMA_F16_CASE(512, 512, 16, 2);
2009
+ extern DECL_FATTN_MMA_F16_CASE(512, 512, 32, 2);
2006
2010
  extern DECL_FATTN_MMA_F16_CASE(512, 512, 2, 4);
2007
2011
  extern DECL_FATTN_MMA_F16_CASE(512, 512, 4, 4);
2008
2012
  extern DECL_FATTN_MMA_F16_CASE(512, 512, 8, 4);
@@ -76,6 +76,7 @@ static constexpr __host__ __device__ uint32_t ggml_cuda_fattn_tile_get_config_nv
76
76
 
77
77
  GGML_CUDA_FATTN_TILE_CONFIG_CASE(320, 256, 16, 256, 2, 64, 64)
78
78
 
79
+ GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 2, 64, 2, 64, 64)
79
80
  GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 4, 128, 2, 64, 64)
80
81
  GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 8, 256, 2, 64, 64)
81
82
  GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 16, 256, 2, 64, 64)
@@ -144,6 +145,7 @@ static constexpr __host__ __device__ uint32_t ggml_cuda_fattn_tile_get_config_nv
144
145
 
145
146
  GGML_CUDA_FATTN_TILE_CONFIG_CASE(320, 256, 16, 256, 2, 32, 64)
146
147
 
148
+ GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 2, 64, 2, 32, 64)
147
149
  GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 4, 128, 2, 32, 64)
148
150
  GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 8, 256, 2, 32, 64)
149
151
  GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 16, 256, 2, 32, 64)
@@ -219,6 +221,7 @@ static constexpr __host__ __device__ uint32_t ggml_cuda_fattn_tile_get_config_am
219
221
 
220
222
  GGML_CUDA_FATTN_TILE_CONFIG_CASE(320, 256, 32, 512, 1, 128, 64)
221
223
 
224
+ GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 2, 64, 2, 64, 64)
222
225
  GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 4, 128, 2, 64, 64)
223
226
  GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 8, 256, 2, 64, 64)
224
227
  GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 16, 256, 2, 64, 64)
@@ -296,6 +299,7 @@ static constexpr __host__ __device__ uint32_t ggml_cuda_fattn_tile_get_config_am
296
299
 
297
300
  GGML_CUDA_FATTN_TILE_CONFIG_CASE(320, 256, 32, 256, 2, 128, 64)
298
301
 
302
+ GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 2, 64, 2, 64, 64)
299
303
  GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 4, 128, 2, 64, 64)
300
304
  GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 8, 256, 2, 64, 64)
301
305
  GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 16, 256, 4, 64, 64)
@@ -1308,12 +1312,12 @@ static void launch_fattn_tile_switch_ncols2(ggml_backend_cuda_context & ctx, ggm
1308
1312
  return;
1309
1313
  }
1310
1314
 
1311
- if constexpr (DV <= 256) {
1312
- if (use_gqa_opt && gqa_ratio % 2 == 0) {
1313
- launch_fattn_tile_switch_ncols1<DKQ, DV, 2, use_logit_softcap>(ctx, dst);
1314
- return;
1315
- }
1315
+ if (use_gqa_opt && gqa_ratio % 2 == 0) {
1316
+ launch_fattn_tile_switch_ncols1<DKQ, DV, 2, use_logit_softcap>(ctx, dst);
1317
+ return;
1318
+ }
1316
1319
 
1320
+ if constexpr (DV <= 256) {
1317
1321
  launch_fattn_tile_switch_ncols1<DKQ, DV, 1, use_logit_softcap>(ctx, dst);
1318
1322
  return;
1319
1323
  }
@@ -99,12 +99,12 @@ static void ggml_cuda_flash_attn_ext_mma_f16_switch_ncols2(ggml_backend_cuda_con
99
99
  return;
100
100
  }
101
101
 
102
- if constexpr (DKQ <= 256) {
103
- if (use_gqa_opt && gqa_ratio > 1) {
104
- ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1<DKQ, DV, 2>(ctx, dst);
105
- return;
106
- }
102
+ if (use_gqa_opt && gqa_ratio > 1) {
103
+ ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1<DKQ, DV, 2>(ctx, dst);
104
+ return;
105
+ }
107
106
 
107
+ if constexpr (DKQ <= 256) {
108
108
  ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1<DKQ, DV, 1>(ctx, dst);
109
109
  } else {
110
110
  GGML_ABORT("fatal error");
@@ -337,6 +337,26 @@ enum best_fattn_kernel {
337
337
  BEST_FATTN_KERNEL_MMA_F16 = 400,
338
338
  };
339
339
 
340
+ static bool ggml_cuda_fattn_kv_type_supported(ggml_type type) {
341
+ switch (type) {
342
+ case GGML_TYPE_F32:
343
+ case GGML_TYPE_F16:
344
+ return true;
345
+ case GGML_TYPE_Q4_1:
346
+ case GGML_TYPE_Q5_0:
347
+ case GGML_TYPE_Q5_1:
348
+ #ifndef GGML_CUDA_FA_ALL_QUANTS
349
+ return false;
350
+ #endif // GGML_CUDA_FA_ALL_QUANTS
351
+ case GGML_TYPE_Q4_0:
352
+ case GGML_TYPE_Q8_0:
353
+ case GGML_TYPE_BF16:
354
+ return true;
355
+ default:
356
+ return false;
357
+ }
358
+ }
359
+
340
360
  static best_fattn_kernel ggml_cuda_get_best_fattn_kernel(const int device, const ggml_tensor * dst) {
341
361
  #ifndef FLASH_ATTN_AVAILABLE
342
362
  GGML_UNUSED(device); GGML_UNUSED(dst);
@@ -427,22 +447,8 @@ static best_fattn_kernel ggml_cuda_get_best_fattn_kernel(const int device, const
427
447
  }
428
448
  #endif // GGML_CUDA_FA_ALL_QUANTS
429
449
 
430
- switch (K->type) {
431
- case GGML_TYPE_F32:
432
- case GGML_TYPE_F16:
433
- break;
434
- case GGML_TYPE_Q4_1:
435
- case GGML_TYPE_Q5_0:
436
- case GGML_TYPE_Q5_1:
437
- #ifndef GGML_CUDA_FA_ALL_QUANTS
438
- return BEST_FATTN_KERNEL_NONE;
439
- #endif // GGML_CUDA_FA_ALL_QUANTS
440
- case GGML_TYPE_Q4_0:
441
- case GGML_TYPE_Q8_0:
442
- case GGML_TYPE_BF16:
443
- break;
444
- default:
445
- return BEST_FATTN_KERNEL_NONE;
450
+ if (!ggml_cuda_fattn_kv_type_supported(K->type) || !ggml_cuda_fattn_kv_type_supported(V->type)) {
451
+ return BEST_FATTN_KERNEL_NONE;
446
452
  }
447
453
 
448
454
  if (mask && mask->ne[2] != 1) {