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.
- checksums.yaml +4 -4
- data/README.md +5 -4
- data/ext/options.rb +1 -1
- data/ext/ruby_whisper.c +0 -1
- data/ext/ruby_whisper.h +7 -1
- data/ext/ruby_whisper_context.c +50 -1
- data/ext/ruby_whisper_log_settable.h +1 -2
- data/ext/ruby_whisper_params.c +9 -8
- data/ext/ruby_whisper_transcribe.cpp +0 -19
- data/ext/ruby_whisper_vad_context.c +30 -10
- data/ext/ruby_whisper_vad_context_detect.cpp +8 -9
- data/ext/ruby_whisper_vad_params.c +4 -4
- data/ext/ruby_whisper_vad_segment.c +2 -2
- data/ext/sources/CMakeLists.txt +2 -1
- data/ext/sources/cmake/parakeet.pc.in +2 -2
- data/ext/sources/cmake/whisper.pc.in +2 -2
- data/ext/sources/examples/cli/cli.cpp +9 -1
- data/ext/sources/examples/common-ggml.cpp +2 -0
- data/ext/sources/examples/vad-speech-segments/speech.cpp +3 -2
- data/ext/sources/ggml/CMakeLists.txt +3 -4
- data/ext/sources/ggml/include/ggml-cuda.h +0 -3
- data/ext/sources/ggml/include/ggml-sycl.h +8 -0
- data/ext/sources/ggml/include/ggml.h +3 -1
- data/ext/sources/ggml/src/CMakeLists.txt +8 -1
- data/ext/sources/ggml/src/ggml-backend-meta.cpp +7 -4
- data/ext/sources/ggml/src/ggml-common.h +13 -2
- data/ext/sources/ggml/src/ggml-cpu/CMakeLists.txt +1 -1
- data/ext/sources/ggml/src/ggml-cpu/amx/mmq.cpp +5 -6
- data/ext/sources/ggml/src/ggml-cpu/arch/arm/quants.c +78 -4
- data/ext/sources/ggml/src/ggml-cpu/arch/x86/quants.c +142 -4
- data/ext/sources/ggml/src/ggml-cpu/arch-fallback.h +7 -2
- data/ext/sources/ggml/src/ggml-cpu/ggml-cpu.c +14 -0
- data/ext/sources/ggml/src/ggml-cpu/llamafile/sgemm.cpp +26 -19
- data/ext/sources/ggml/src/ggml-cpu/ops.cpp +129 -46
- data/ext/sources/ggml/src/ggml-cpu/quants.c +51 -0
- data/ext/sources/ggml/src/ggml-cpu/quants.h +3 -0
- data/ext/sources/ggml/src/ggml-cpu/simd-gemm.h +1 -1
- data/ext/sources/ggml/src/ggml-cpu/simd-mappings.h +11 -0
- data/ext/sources/ggml/src/ggml-cpu/vec.cpp +2 -2
- data/ext/sources/ggml/src/ggml-cuda/binbcast.cu +90 -46
- data/ext/sources/ggml/src/ggml-cuda/col2im-1d.cu +81 -0
- data/ext/sources/ggml/src/ggml-cuda/col2im-1d.cuh +3 -0
- data/ext/sources/ggml/src/ggml-cuda/common.cuh +4 -0
- data/ext/sources/ggml/src/ggml-cuda/concat.cu +33 -21
- data/ext/sources/ggml/src/ggml-cuda/conv-transpose-1d.cu +14 -12
- data/ext/sources/ggml/src/ggml-cuda/convert.cu +86 -34
- data/ext/sources/ggml/src/ggml-cuda/cpy.cu +80 -29
- data/ext/sources/ggml/src/ggml-cuda/fattn-common.cuh +9 -5
- data/ext/sources/ggml/src/ggml-cuda/fattn-mma-f16.cuh +4 -0
- data/ext/sources/ggml/src/ggml-cuda/fattn-tile.cuh +9 -5
- data/ext/sources/ggml/src/ggml-cuda/fattn.cu +27 -21
- data/ext/sources/ggml/src/ggml-cuda/gated_delta_net.cu +40 -25
- data/ext/sources/ggml/src/ggml-cuda/gated_delta_net.cuh +10 -0
- data/ext/sources/ggml/src/ggml-cuda/getrows.cu +15 -12
- data/ext/sources/ggml/src/ggml-cuda/ggml-cuda.cu +718 -1248
- data/ext/sources/ggml/src/ggml-cuda/mmq.cu +7 -0
- data/ext/sources/ggml/src/ggml-cuda/mmvq.cu +77 -40
- data/ext/sources/ggml/src/ggml-cuda/out-prod.cu +55 -12
- data/ext/sources/ggml/src/ggml-cuda/set-rows.cu +64 -4
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_16-ncols2_2.cu +1 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_32-ncols2_2.cu +1 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_4-ncols2_2.cu +1 -0
- data/ext/sources/ggml/src/ggml-cuda/template-instances/fattn-mma-f16-instance-ncols1_8-ncols2_2.cu +1 -0
- data/ext/sources/ggml/src/ggml-cuda/topk-moe.cu +7 -1
- data/ext/sources/ggml/src/ggml-cuda/vendors/hip.h +1 -0
- data/ext/sources/ggml/src/ggml-cuda/vendors/musa.h +1 -0
- data/ext/sources/ggml/src/ggml-hexagon/CMakeLists.txt +0 -5
- data/ext/sources/ggml/src/ggml-hexagon/ggml-hexagon.cpp +1634 -1293
- data/ext/sources/ggml/src/ggml-hexagon/htp/CMakeLists.txt +11 -40
- data/ext/sources/ggml/src/ggml-hexagon/htp/cmake-toolchain.cmake +13 -15
- data/ext/sources/ggml/src/ggml-hexagon/htp/concat-ops.c +1 -1
- data/ext/sources/ggml/src/ggml-hexagon/htp/flash-attn-ops.c +1749 -399
- data/ext/sources/ggml/src/ggml-hexagon/htp/flash-attn-ops.h +303 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hex-common.h +80 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hex-dma.h +26 -23
- data/ext/sources/ggml/src/ggml-hexagon/htp/hex-profile.h +64 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hex-utils.h +1 -83
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-fa-kernels.h +555 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-mm-kernels-tiled.h +1303 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-queue.c +9 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-queue.h +27 -4
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-utils.h +59 -37
- data/ext/sources/ggml/src/ggml-hexagon/htp/htp-ctx.h +11 -3
- data/ext/sources/ggml/src/ggml-hexagon/htp/htp-ops.h +52 -12
- data/ext/sources/ggml/src/ggml-hexagon/htp/htp-vtcm.h +19 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/htp_iface.idl +2 -1
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-base.h +14 -30
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-exp.h +39 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-fa-kernels.h +232 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-flat.h +1511 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-mm-kernels-tiled.h +1200 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/hvx-sigmoid.h +39 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/main.c +127 -32
- data/ext/sources/ggml/src/ggml-hexagon/htp/matmul-ops.c +3023 -4425
- data/ext/sources/ggml/src/ggml-hexagon/htp/matmul-ops.h +650 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp/rope-ops.c +48 -13
- data/ext/sources/ggml/src/ggml-hexagon/htp/ssm-conv.c +10 -9
- data/ext/sources/ggml/src/ggml-hexagon/htp/worker-pool.c +15 -3
- data/ext/sources/ggml/src/ggml-hexagon/htp/worker-pool.h +8 -0
- data/ext/sources/ggml/src/ggml-hexagon/htp-opnode.h +168 -50
- data/ext/sources/ggml/src/ggml-hexagon/libggml-htp.inf +0 -4
- data/ext/sources/ggml/src/ggml-hip/CMakeLists.txt +5 -0
- data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.cpp +69 -5
- data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.h +4 -1
- data/ext/sources/ggml/src/ggml-metal/ggml-metal-device.m +27 -6
- data/ext/sources/ggml/src/ggml-metal/ggml-metal-impl.h +38 -0
- data/ext/sources/ggml/src/ggml-metal/ggml-metal-ops.cpp +132 -2
- data/ext/sources/ggml/src/ggml-metal/ggml-metal-ops.h +2 -0
- data/ext/sources/ggml/src/ggml-metal/ggml-metal.metal +345 -87
- data/ext/sources/ggml/src/ggml-opencl/CMakeLists.txt +13 -0
- data/ext/sources/ggml/src/ggml-opencl/fa_tune.h +92 -0
- data/ext/sources/ggml/src/ggml-opencl/ggml-opencl.cpp +4060 -357
- data/ext/sources/ggml/src/ggml-opencl/kernels/cvt.cl +198 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f16.cl +81 -41
- data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32.cl +88 -39
- data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_f16.cl +1995 -96
- data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_q4_0.cl +1615 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_f32_q8_0.cl +1486 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/flash_attn_pre_f16.cl +156 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_mxfp4_f32_ns.cl +74 -6
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_0_f32_ns.cl +74 -6
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_1_f32_ns.cl +74 -6
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q4_k_f32_ns.cl +71 -6
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_0_f32_ns.cl +74 -6
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_1_f32_ns.cl +74 -6
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q5_k_f32_ns.cl +74 -6
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_moe_q6_k_f32_ns.cl +74 -6
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemm_noshuffle_q1_0_f32.cl +94 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q1_0_f32.cl +121 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q8_0_f32.cl +1 -1
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mm_q1_0_f32_l4_lm.cl +156 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_f16_f32_l4.cl +1149 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q1_0_f32.cl +141 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/mul_mv_q1_0_f32_flat.cl +190 -0
- data/ext/sources/ggml/src/ggml-opencl/kernels/norm.cl +5 -2
- data/ext/sources/ggml/src/ggml-opencl/kernels/set_rows.cl +500 -0
- data/ext/sources/ggml/src/ggml-opencl/libdl.h +79 -0
- data/ext/sources/ggml/src/ggml-openvino/.clang-format +0 -5
- data/ext/sources/ggml/src/ggml-openvino/CMakeLists.txt +2 -4
- data/ext/sources/ggml/src/ggml-openvino/ggml-decoder.cpp +733 -130
- data/ext/sources/ggml/src/ggml-openvino/ggml-decoder.h +76 -23
- data/ext/sources/ggml/src/ggml-openvino/ggml-openvino-extra.cpp +57 -3
- data/ext/sources/ggml/src/ggml-openvino/ggml-openvino-extra.h +29 -8
- data/ext/sources/ggml/src/ggml-openvino/ggml-openvino.cpp +307 -59
- data/ext/sources/ggml/src/ggml-openvino/ggml-quants.cpp +66 -0
- data/ext/sources/ggml/src/ggml-openvino/ggml-quants.h +10 -4
- data/ext/sources/ggml/src/ggml-openvino/openvino/decoder.h +56 -16
- data/ext/sources/ggml/src/ggml-openvino/openvino/frontend.h +1 -1
- data/ext/sources/ggml/src/ggml-openvino/openvino/input_model.h +4 -4
- data/ext/sources/ggml/src/ggml-openvino/openvino/node_context.h +94 -37
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/add_id.cpp +76 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/argsort.cpp +47 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/clamp.cpp +33 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/concat.cpp +48 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/cont.cpp +8 -16
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/cpy.cpp +14 -1
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/div.cpp +146 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/flash_attn_ext.cpp +108 -21
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/gated_delta_net.cpp +282 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/gated_delta_net.hpp +65 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/get_rows.cpp +2 -9
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/glu_geglu.cpp +21 -7
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/glu_swiglu.cpp +41 -8
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/im2col.cpp +120 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/l2_norm.cpp +44 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/mul_mat_id.cpp +226 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/mulmat.cpp +19 -9
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/norm.cpp +58 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/pad.cpp +95 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/permute.cpp +58 -13
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/repeat.cpp +74 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/reshape.cpp +13 -6
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/rms_norm.cpp +1 -1
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/rope.cpp +134 -38
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/set_rows.cpp +3 -3
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/softmax.cpp +126 -49
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/ssm_conv.cpp +59 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/sum_rows.cpp +27 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/transpose.cpp +32 -1
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_silu.cpp +1 -1
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_softplus.cpp +38 -0
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/view.cpp +90 -25
- data/ext/sources/ggml/src/ggml-openvino/openvino/op_table.cpp +41 -23
- data/ext/sources/ggml/src/ggml-openvino/openvino/op_table.h +18 -5
- data/ext/sources/ggml/src/ggml-openvino/openvino/pass/mark_decompression_convert_constant_folding.h +1 -1
- data/ext/sources/ggml/src/ggml-openvino/openvino/translate_session.cpp +43 -40
- data/ext/sources/ggml/src/ggml-openvino/openvino/translate_session.h +5 -4
- data/ext/sources/ggml/src/ggml-openvino/openvino/utils.cpp +548 -3
- data/ext/sources/ggml/src/ggml-openvino/openvino/utils.h +28 -26
- data/ext/sources/ggml/src/ggml-openvino/utils.cpp +383 -94
- data/ext/sources/ggml/src/ggml-openvino/utils.h +11 -8
- data/ext/sources/ggml/src/ggml-quants.c +76 -0
- data/ext/sources/ggml/src/ggml-quants.h +3 -0
- data/ext/sources/ggml/src/ggml-sycl/CMakeLists.txt +5 -5
- data/ext/sources/ggml/src/ggml-sycl/backend.hpp +2 -0
- data/ext/sources/ggml/src/ggml-sycl/binbcast.cpp +12 -0
- data/ext/sources/ggml/src/ggml-sycl/col2im-1d.cpp +102 -0
- data/ext/sources/ggml/src/ggml-sycl/col2im-1d.hpp +8 -0
- data/ext/sources/ggml/src/ggml-sycl/common.cpp +6 -8
- data/ext/sources/ggml/src/ggml-sycl/common.hpp +19 -2
- data/ext/sources/ggml/src/ggml-sycl/concat.cpp +21 -1
- data/ext/sources/ggml/src/ggml-sycl/conv2d-dw.cpp +158 -0
- data/ext/sources/ggml/src/ggml-sycl/conv2d-dw.hpp +10 -0
- data/ext/sources/ggml/src/ggml-sycl/conv2d-transpose.cpp +125 -0
- data/ext/sources/ggml/src/ggml-sycl/conv2d-transpose.hpp +10 -0
- data/ext/sources/ggml/src/ggml-sycl/conv2d.cpp +150 -0
- data/ext/sources/ggml/src/ggml-sycl/conv2d.hpp +10 -0
- data/ext/sources/ggml/src/ggml-sycl/conv3d.cpp +224 -0
- data/ext/sources/ggml/src/ggml-sycl/conv3d.hpp +8 -0
- data/ext/sources/ggml/src/ggml-sycl/convert.cpp +6 -0
- data/ext/sources/ggml/src/ggml-sycl/cpy.cpp +706 -0
- data/ext/sources/ggml/src/ggml-sycl/cpy.hpp +281 -0
- data/ext/sources/ggml/src/ggml-sycl/cross_entropy_loss.cpp +255 -0
- data/ext/sources/ggml/src/ggml-sycl/cross_entropy_loss.hpp +7 -0
- data/ext/sources/ggml/src/ggml-sycl/dequantize.hpp +15 -0
- data/ext/sources/ggml/src/ggml-sycl/dmmv.cpp +492 -319
- data/ext/sources/ggml/src/ggml-sycl/dpct/helper.hpp +15 -7
- data/ext/sources/ggml/src/ggml-sycl/element_wise.cpp +215 -115
- data/ext/sources/ggml/src/ggml-sycl/element_wise.hpp +2 -0
- data/ext/sources/ggml/src/ggml-sycl/ggml-sycl.cpp +1006 -336
- data/ext/sources/ggml/src/ggml-sycl/mmvq.cpp +252 -67
- data/ext/sources/ggml/src/ggml-sycl/mmvq.hpp +17 -0
- data/ext/sources/ggml/src/ggml-sycl/norm.cpp +103 -49
- data/ext/sources/ggml/src/ggml-sycl/outprod.cpp +45 -9
- data/ext/sources/ggml/src/ggml-sycl/pool.cpp +185 -0
- data/ext/sources/ggml/src/ggml-sycl/pool.hpp +22 -0
- data/ext/sources/ggml/src/ggml-sycl/presets.hpp +3 -1
- data/ext/sources/ggml/src/ggml-sycl/set_rows.cpp +10 -2
- data/ext/sources/ggml/src/ggml-sycl/softmax.cpp +9 -10
- data/ext/sources/ggml/src/ggml-sycl/vecdotq.hpp +35 -0
- data/ext/sources/ggml/src/ggml-vulkan/CMakeLists.txt +5 -0
- data/ext/sources/ggml/src/ggml-vulkan/ggml-vulkan.cpp +833 -215
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/col2im_1d.comp +61 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/conv2d_mm.comp +1 -1
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/conv3d_mm.comp +431 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/diag.comp +3 -3
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp +1 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp +1 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/generic_unary_head.glsl +21 -19
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/get_rows_back.comp +25 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/glu_head.glsl +23 -4
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/glu_main.glsl +14 -18
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/l2_norm.comp +4 -7
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq.comp +21 -24
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp +31 -23
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_cm2.comp +6 -5
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl +84 -67
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/norm.comp +10 -10
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/repeat_back.comp +3 -3
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/roll.comp +3 -3
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/tri.comp +3 -3
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/unary.comp +168 -0
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +121 -74
- data/ext/sources/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp +26 -19
- data/ext/sources/ggml/src/ggml-webgpu/ggml-webgpu.cpp +31 -36
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/common_decls.tmpl +16 -2
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_split.wgsl +7 -7
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/get_rows.wgsl +21 -0
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_decls.tmpl +439 -320
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_id_vec.wgsl +2 -2
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec.wgsl +45 -39
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_acc.tmpl +586 -465
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/mul_mat_vec_q_acc.tmpl +63 -69
- data/ext/sources/ggml/src/ggml-webgpu/wgsl-shaders/quantize_q8.wgsl +14 -9
- data/ext/sources/ggml/src/ggml.c +36 -14
- data/ext/sources/include/whisper.h +21 -0
- data/ext/sources/src/whisper.cpp +164 -14
- data/lib/whisper/log_settable.rb +5 -8
- data/lib/whisper/model/uri.rb +0 -7
- data/sig/whisper.rbs +6 -0
- data/test/test_vad.rb +9 -0
- data/test/test_vad_context.rb +2 -2
- data/whispercpp.gemspec +1 -1
- metadata +62 -37
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-flash-attn-ops.c +0 -1878
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-matmul-ops.c +0 -2066
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-ops.c +0 -6
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-ops.h +0 -88
- data/ext/sources/ggml/src/ggml-hexagon/htp/hmx-profile.h +0 -34
- data/ext/sources/ggml/src/ggml-hexagon/htp/vtcm-utils.h +0 -16
- data/ext/sources/ggml/src/ggml-openvino/openvino/op/unary_gelu.cpp +0 -25
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/abs.comp +0 -21
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/ceil.comp +0 -22
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/clamp.comp +0 -17
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/cos.comp +0 -17
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/elu.comp +0 -27
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/exp.comp +0 -20
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/floor.comp +0 -22
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu.comp +0 -25
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu_erf.comp +0 -39
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/gelu_quick.comp +0 -23
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/hardsigmoid.comp +0 -22
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/hardswish.comp +0 -22
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/leaky_relu.comp +0 -22
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/neg.comp +0 -20
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/relu.comp +0 -21
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/round.comp +0 -29
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sgn.comp +0 -21
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sigmoid.comp +0 -20
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/silu.comp +0 -22
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sin.comp +0 -17
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/softplus.comp +0 -23
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/sqrt.comp +0 -17
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/square.comp +0 -17
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/step.comp +0 -22
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/tanh.comp +0 -20
- data/ext/sources/ggml/src/ggml-vulkan/vulkan-shaders/trunc.comp +0 -22
- 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)
|
|
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)
|
|
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)
|
|
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)
|
|
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
|
|
57
|
-
const
|
|
58
|
-
const
|
|
59
|
-
const
|
|
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
|
|
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
|
|
228
|
-
|
|
229
|
-
|
|
230
|
-
|
|
231
|
-
|
|
232
|
-
|
|
233
|
-
|
|
234
|
-
|
|
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
|
-
|
|
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
|
|
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
|
|
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
|
|
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
|
|
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
|
|
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
|
|
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
|
|
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
|
|
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
|
|
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
|
|
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
|
|
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 *
|
|
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
|
|
1093
|
-
const
|
|
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
|
-
|
|
1103
|
-
|
|
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
|
|
1312
|
-
|
|
1313
|
-
|
|
1314
|
-
|
|
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
|
|
103
|
-
|
|
104
|
-
|
|
105
|
-
|
|
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
|
-
|
|
431
|
-
|
|
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) {
|