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
|
@@ -27,6 +27,8 @@
|
|
|
27
27
|
#define QR5_1 2
|
|
28
28
|
#define QK8_0 32
|
|
29
29
|
#define QR8_0 1
|
|
30
|
+
#define QK1_0 128
|
|
31
|
+
#define QR1_0 1
|
|
30
32
|
#define QK_K 256
|
|
31
33
|
#define K_SCALE_SIZE (3 * QK_K / 64)
|
|
32
34
|
#define K_QUANTS_PER_ITERATION 2
|
|
@@ -38,6 +40,14 @@ typedef ushort uint16_t;
|
|
|
38
40
|
typedef int int32_t;
|
|
39
41
|
typedef uint uint32_t;
|
|
40
42
|
|
|
43
|
+
//------------------------------------------------------------------------------
|
|
44
|
+
// block_q1_0
|
|
45
|
+
//------------------------------------------------------------------------------
|
|
46
|
+
typedef struct {
|
|
47
|
+
half d; // delta
|
|
48
|
+
uchar qs[QK1_0/8]; // 1-bit signs (16 bytes)
|
|
49
|
+
} block_q1_0;
|
|
50
|
+
|
|
41
51
|
//------------------------------------------------------------------------------
|
|
42
52
|
// block_q4_0
|
|
43
53
|
//------------------------------------------------------------------------------
|
|
@@ -159,6 +169,42 @@ kernel void kernel_convert_f16_to_bf16(
|
|
|
159
169
|
}
|
|
160
170
|
}
|
|
161
171
|
|
|
172
|
+
//------------------------------------------------------------------------------
|
|
173
|
+
// kernel_convert_block_q1_0
|
|
174
|
+
// Convert block_q1_0 (AOS) to 2 separate arrays (SOA): quant bytes + scales.
|
|
175
|
+
// q1_0 bits are stored in natural order (bit j of byte i -> weight 8*i + j)
|
|
176
|
+
//------------------------------------------------------------------------------
|
|
177
|
+
kernel void kernel_convert_block_q1_0(
|
|
178
|
+
global block_q1_0 * src0,
|
|
179
|
+
global uchar * dst_q,
|
|
180
|
+
global half * dst_d
|
|
181
|
+
) {
|
|
182
|
+
global block_q1_0 * b = (global block_q1_0 *) src0 + get_global_id(0);
|
|
183
|
+
global uchar * q = (global uchar *) dst_q + (QK1_0/8)*get_global_id(0);
|
|
184
|
+
global half * d = (global half *) dst_d + get_global_id(0);
|
|
185
|
+
|
|
186
|
+
*d = b->d;
|
|
187
|
+
|
|
188
|
+
for (int i = 0; i < QK1_0/8; ++i) {
|
|
189
|
+
q[i] = b->qs[i];
|
|
190
|
+
}
|
|
191
|
+
}
|
|
192
|
+
|
|
193
|
+
kernel void kernel_restore_block_q1_0(
|
|
194
|
+
global uchar * src_q,
|
|
195
|
+
global half * src_d,
|
|
196
|
+
global block_q1_0 * dst
|
|
197
|
+
) {
|
|
198
|
+
global block_q1_0 * b = (global block_q1_0 *) dst + get_global_id(0);
|
|
199
|
+
global uchar * q = (global uchar *) src_q + (QK1_0/8)*get_global_id(0);
|
|
200
|
+
global half * d = (global half *) src_d + get_global_id(0);
|
|
201
|
+
|
|
202
|
+
b->d = *d;
|
|
203
|
+
for (int i = 0; i < QK1_0/8; ++i) {
|
|
204
|
+
b->qs[i] = q[i];
|
|
205
|
+
}
|
|
206
|
+
}
|
|
207
|
+
|
|
162
208
|
//------------------------------------------------------------------------------
|
|
163
209
|
// kernel_convert_block_q4_0
|
|
164
210
|
// Convert the block_q4_0 format to 2 separate arrays (AOS -> SOA).
|
|
@@ -1582,6 +1628,158 @@ kernel void kernel_restore_block_q8_0(
|
|
|
1582
1628
|
}
|
|
1583
1629
|
}
|
|
1584
1630
|
|
|
1631
|
+
// View-aware AoS q8_0 -> f32 dequant (f32/f32 FA path).
|
|
1632
|
+
kernel void kernel_dequant_q8_0_f32_view_aos(
|
|
1633
|
+
global char * src,
|
|
1634
|
+
ulong src_offset,
|
|
1635
|
+
ulong src_nb1,
|
|
1636
|
+
ulong src_nb2,
|
|
1637
|
+
ulong src_nb3,
|
|
1638
|
+
int nblk0,
|
|
1639
|
+
int ne1,
|
|
1640
|
+
int ne2,
|
|
1641
|
+
int ne3,
|
|
1642
|
+
global float * dst
|
|
1643
|
+
) {
|
|
1644
|
+
int blk_i0 = get_global_id(0);
|
|
1645
|
+
int i1 = get_global_id(1);
|
|
1646
|
+
int batch = get_global_id(2);
|
|
1647
|
+
|
|
1648
|
+
if (blk_i0 >= nblk0) return;
|
|
1649
|
+
if (i1 >= ne1) return;
|
|
1650
|
+
|
|
1651
|
+
int i2 = batch % ne2;
|
|
1652
|
+
int i3 = batch / ne2;
|
|
1653
|
+
if (i3 >= ne3) return;
|
|
1654
|
+
|
|
1655
|
+
global char * block = src + src_offset + (ulong)i3*src_nb3 + (ulong)i2*src_nb2 + (ulong)i1*src_nb1 + (ulong)blk_i0 * (2 + QK8_0);
|
|
1656
|
+
float d = vload_half(0, (global half *)block);
|
|
1657
|
+
global char * qs = block + 2;
|
|
1658
|
+
|
|
1659
|
+
ulong dst_row_base = ((ulong)i3 * ne2 * ne1 + (ulong)i2 * ne1 + (ulong)i1) * nblk0;
|
|
1660
|
+
global float * out = dst + (dst_row_base + blk_i0) * QK8_0;
|
|
1661
|
+
|
|
1662
|
+
for (int i = 0; i < QK8_0; ++i) {
|
|
1663
|
+
out[i] = d * (float)qs[i];
|
|
1664
|
+
}
|
|
1665
|
+
}
|
|
1666
|
+
|
|
1667
|
+
// View-aware AoS q8_0 -> f16 dequant. Rows tight, batch strides may be gapped.
|
|
1668
|
+
kernel void kernel_dequant_q8_0_f16_view_aos(
|
|
1669
|
+
global char * src,
|
|
1670
|
+
ulong src_offset,
|
|
1671
|
+
ulong src_nb1,
|
|
1672
|
+
ulong src_nb2,
|
|
1673
|
+
ulong src_nb3,
|
|
1674
|
+
int nblk0,
|
|
1675
|
+
int ne1,
|
|
1676
|
+
int ne2,
|
|
1677
|
+
int ne3,
|
|
1678
|
+
global half * dst
|
|
1679
|
+
) {
|
|
1680
|
+
int blk_i0 = get_global_id(0);
|
|
1681
|
+
int i1 = get_global_id(1);
|
|
1682
|
+
int batch = get_global_id(2);
|
|
1683
|
+
|
|
1684
|
+
if (blk_i0 >= nblk0) return;
|
|
1685
|
+
if (i1 >= ne1) return;
|
|
1686
|
+
|
|
1687
|
+
int i2 = batch % ne2;
|
|
1688
|
+
int i3 = batch / ne2;
|
|
1689
|
+
if (i3 >= ne3) return;
|
|
1690
|
+
|
|
1691
|
+
global char * block = src + src_offset + (ulong)i3*src_nb3 + (ulong)i2*src_nb2 + (ulong)i1*src_nb1 + (ulong)blk_i0 * (2 + QK8_0);
|
|
1692
|
+
float d = vload_half(0, (global half *)block);
|
|
1693
|
+
global char * qs = block + 2;
|
|
1694
|
+
|
|
1695
|
+
ulong dst_row_base = ((ulong)i3 * ne2 * ne1 + (ulong)i2 * ne1 + (ulong)i1) * nblk0;
|
|
1696
|
+
global half * out = dst + (dst_row_base + blk_i0) * QK8_0;
|
|
1697
|
+
|
|
1698
|
+
for (int i = 0; i < QK8_0; ++i) {
|
|
1699
|
+
out[i] = (half)(d * (float)qs[i]);
|
|
1700
|
+
}
|
|
1701
|
+
}
|
|
1702
|
+
|
|
1703
|
+
// View-aware AoS q4_0 -> f32 dequant (mirrors the q8_0 view variant).
|
|
1704
|
+
kernel void kernel_dequant_q4_0_f32_view_aos(
|
|
1705
|
+
global char * src,
|
|
1706
|
+
ulong src_offset,
|
|
1707
|
+
ulong src_nb1,
|
|
1708
|
+
ulong src_nb2,
|
|
1709
|
+
ulong src_nb3,
|
|
1710
|
+
int nblk0,
|
|
1711
|
+
int ne1,
|
|
1712
|
+
int ne2,
|
|
1713
|
+
int ne3,
|
|
1714
|
+
global float * dst
|
|
1715
|
+
) {
|
|
1716
|
+
int blk_i0 = get_global_id(0);
|
|
1717
|
+
int i1 = get_global_id(1);
|
|
1718
|
+
int batch = get_global_id(2);
|
|
1719
|
+
|
|
1720
|
+
if (blk_i0 >= nblk0) return;
|
|
1721
|
+
if (i1 >= ne1) return;
|
|
1722
|
+
|
|
1723
|
+
int i2 = batch % ne2;
|
|
1724
|
+
int i3 = batch / ne2;
|
|
1725
|
+
if (i3 >= ne3) return;
|
|
1726
|
+
|
|
1727
|
+
global char * block = src + src_offset + (ulong)i3*src_nb3 + (ulong)i2*src_nb2 + (ulong)i1*src_nb1 + (ulong)blk_i0 * (2 + QK4_0/2);
|
|
1728
|
+
float d = vload_half(0, (global half *)block);
|
|
1729
|
+
global uchar * qs = (global uchar *)(block + 2);
|
|
1730
|
+
|
|
1731
|
+
ulong dst_row_base = ((ulong)i3 * ne2 * ne1 + (ulong)i2 * ne1 + (ulong)i1) * nblk0;
|
|
1732
|
+
global float * out = dst + (dst_row_base + blk_i0) * QK4_0;
|
|
1733
|
+
|
|
1734
|
+
for (int i = 0; i < QK4_0/2; ++i) {
|
|
1735
|
+
uchar byte = qs[i];
|
|
1736
|
+
int q0 = (int)(byte & 0x0F) - 8;
|
|
1737
|
+
int q1 = (int)(byte >> 4) - 8;
|
|
1738
|
+
out[i] = d * (float)q0;
|
|
1739
|
+
out[i + QK4_0/2] = d * (float)q1;
|
|
1740
|
+
}
|
|
1741
|
+
}
|
|
1742
|
+
|
|
1743
|
+
// View-aware AoS q4_0 -> f16 dequant (mirrors the q8_0 view variant).
|
|
1744
|
+
kernel void kernel_dequant_q4_0_f16_view_aos(
|
|
1745
|
+
global char * src,
|
|
1746
|
+
ulong src_offset,
|
|
1747
|
+
ulong src_nb1,
|
|
1748
|
+
ulong src_nb2,
|
|
1749
|
+
ulong src_nb3,
|
|
1750
|
+
int nblk0,
|
|
1751
|
+
int ne1,
|
|
1752
|
+
int ne2,
|
|
1753
|
+
int ne3,
|
|
1754
|
+
global half * dst
|
|
1755
|
+
) {
|
|
1756
|
+
int blk_i0 = get_global_id(0);
|
|
1757
|
+
int i1 = get_global_id(1);
|
|
1758
|
+
int batch = get_global_id(2);
|
|
1759
|
+
|
|
1760
|
+
if (blk_i0 >= nblk0) return;
|
|
1761
|
+
if (i1 >= ne1) return;
|
|
1762
|
+
|
|
1763
|
+
int i2 = batch % ne2;
|
|
1764
|
+
int i3 = batch / ne2;
|
|
1765
|
+
if (i3 >= ne3) return;
|
|
1766
|
+
|
|
1767
|
+
global char * block = src + src_offset + (ulong)i3*src_nb3 + (ulong)i2*src_nb2 + (ulong)i1*src_nb1 + (ulong)blk_i0 * (2 + QK4_0/2);
|
|
1768
|
+
float d = vload_half(0, (global half *)block);
|
|
1769
|
+
global uchar * qs = (global uchar *)(block + 2);
|
|
1770
|
+
|
|
1771
|
+
ulong dst_row_base = ((ulong)i3 * ne2 * ne1 + (ulong)i2 * ne1 + (ulong)i1) * nblk0;
|
|
1772
|
+
global half * out = dst + (dst_row_base + blk_i0) * QK4_0;
|
|
1773
|
+
|
|
1774
|
+
for (int i = 0; i < QK4_0/2; ++i) {
|
|
1775
|
+
uchar byte = qs[i];
|
|
1776
|
+
int q0 = (int)(byte & 0x0F) - 8;
|
|
1777
|
+
int q1 = (int)(byte >> 4) - 8;
|
|
1778
|
+
out[i] = (half)(d * (float)q0);
|
|
1779
|
+
out[i + QK4_0/2] = (half)(d * (float)q1);
|
|
1780
|
+
}
|
|
1781
|
+
}
|
|
1782
|
+
|
|
1585
1783
|
kernel void kernel_restore_block_q8_0_trans(
|
|
1586
1784
|
global uchar * src_q,
|
|
1587
1785
|
global half * src_d,
|
|
@@ -4,13 +4,30 @@
|
|
|
4
4
|
#define ACC_TYPE4 float4
|
|
5
5
|
#define DATA_TYPE half
|
|
6
6
|
#define DATA_TYPE4 half4
|
|
7
|
-
#define CONVERT_ACC4(x)
|
|
8
|
-
#define CONVERT_DATA4(x)
|
|
7
|
+
#define CONVERT_ACC4(x) ((float4)((float)(x).s0, (float)(x).s1, (float)(x).s2, (float)(x).s3))
|
|
8
|
+
#define CONVERT_DATA4(x) ((half4)((half)(x).s0, (half)(x).s1, (half)(x).s2, (half)(x).s3))
|
|
9
9
|
|
|
10
10
|
#define DK_VEC (DK/4)
|
|
11
11
|
#define DV_VEC (DV/4)
|
|
12
12
|
#define WG_SIZE (BLOCK_M)
|
|
13
|
-
|
|
13
|
+
// q1 reduces over a Q1_WG_SIZE-wide WG via work-group barriers; the launch WG
|
|
14
|
+
// must match. Defaults to the Adreno sg (64); host passes -D FA_SG=32 on Intel.
|
|
15
|
+
#ifndef FA_SG
|
|
16
|
+
#define FA_SG 64
|
|
17
|
+
#endif
|
|
18
|
+
#define Q1_WG_SIZE FA_SG
|
|
19
|
+
|
|
20
|
+
// The kernels are built with -cl-finite-math-only. On some older Adreno GPUs,
|
|
21
|
+
// infinite operand can cause undefined behavior and miscompilation for exp.
|
|
22
|
+
// Therefore, a large negative value is used instead.
|
|
23
|
+
#define FA_M_INIT (-3.0e38f)
|
|
24
|
+
|
|
25
|
+
// Drop full unroll at DK>=192 — Adreno compiler host-memory budget.
|
|
26
|
+
#if DK >= 192
|
|
27
|
+
#define FA_UNROLL
|
|
28
|
+
#else
|
|
29
|
+
#define FA_UNROLL _Pragma("unroll")
|
|
30
|
+
#endif
|
|
14
31
|
|
|
15
32
|
inline float get_alibi_slope(
|
|
16
33
|
const float max_bias, const uint h, const uint n_head_log2, const float m0, const float m1
|
|
@@ -81,18 +98,18 @@ __kernel void flash_attn_f16(
|
|
|
81
98
|
if (my_query_row < n_q) {
|
|
82
99
|
const ulong q_row_offset = batch_idx * q_nb3 + head_idx * q_nb2 + my_query_row * q_nb1;
|
|
83
100
|
const global DATA_TYPE4* q_ptr = (const global DATA_TYPE4*)(q_base + q_row_offset);
|
|
84
|
-
|
|
101
|
+
FA_UNROLL
|
|
85
102
|
for (int i = 0; i < DK_VEC; ++i) {
|
|
86
103
|
q_priv[i] = CONVERT_ACC4(q_ptr[i]);
|
|
87
104
|
}
|
|
88
105
|
}
|
|
89
106
|
|
|
90
107
|
ACC_TYPE4 o_acc[DV_VEC];
|
|
91
|
-
|
|
108
|
+
FA_UNROLL
|
|
92
109
|
for (int i = 0; i < DV_VEC; ++i) {
|
|
93
110
|
o_acc[i] = (ACC_TYPE4)(0.0f);
|
|
94
111
|
}
|
|
95
|
-
ACC_TYPE m_i =
|
|
112
|
+
ACC_TYPE m_i = FA_M_INIT;
|
|
96
113
|
ACC_TYPE l_i = 0.0f;
|
|
97
114
|
|
|
98
115
|
float slope = get_alibi_slope(max_bias, head_idx, n_head_log2, m0, m1);
|
|
@@ -125,49 +142,72 @@ __kernel void flash_attn_f16(
|
|
|
125
142
|
continue;
|
|
126
143
|
}
|
|
127
144
|
|
|
128
|
-
for (int j = 0; j < BLOCK_N; j +=
|
|
145
|
+
for (int j = 0; j < BLOCK_N; j += 4) {
|
|
129
146
|
const int k_row0 = k_start + j;
|
|
130
147
|
const int k_row1 = k_start + j + 1;
|
|
148
|
+
const int k_row2 = k_start + j + 2;
|
|
149
|
+
const int k_row3 = k_start + j + 3;
|
|
131
150
|
|
|
132
151
|
ACC_TYPE4 dot_acc0 = (ACC_TYPE4)(0.0f);
|
|
133
152
|
ACC_TYPE4 dot_acc1 = (ACC_TYPE4)(0.0f);
|
|
134
|
-
|
|
153
|
+
ACC_TYPE4 dot_acc2 = (ACC_TYPE4)(0.0f);
|
|
154
|
+
ACC_TYPE4 dot_acc3 = (ACC_TYPE4)(0.0f);
|
|
155
|
+
FA_UNROLL
|
|
135
156
|
for (int k = 0; k < DK_VEC; k++) {
|
|
136
|
-
|
|
137
|
-
|
|
157
|
+
const ACC_TYPE4 qk = q_priv[k];
|
|
158
|
+
dot_acc0 = mad(qk, CONVERT_ACC4(l_k[j][k]), dot_acc0);
|
|
159
|
+
dot_acc1 = mad(qk, CONVERT_ACC4(l_k[j+1][k]), dot_acc1);
|
|
160
|
+
dot_acc2 = mad(qk, CONVERT_ACC4(l_k[j+2][k]), dot_acc2);
|
|
161
|
+
dot_acc3 = mad(qk, CONVERT_ACC4(l_k[j+3][k]), dot_acc3);
|
|
138
162
|
}
|
|
139
|
-
ACC_TYPE
|
|
140
|
-
ACC_TYPE
|
|
163
|
+
ACC_TYPE s0 = (dot_acc0.s0 + dot_acc0.s1 + dot_acc0.s2 + dot_acc0.s3) * scale;
|
|
164
|
+
ACC_TYPE s1 = (dot_acc1.s0 + dot_acc1.s1 + dot_acc1.s2 + dot_acc1.s3) * scale;
|
|
165
|
+
ACC_TYPE s2 = (dot_acc2.s0 + dot_acc2.s1 + dot_acc2.s2 + dot_acc2.s3) * scale;
|
|
166
|
+
ACC_TYPE s3 = (dot_acc3.s0 + dot_acc3.s1 + dot_acc3.s2 + dot_acc3.s3) * scale;
|
|
141
167
|
|
|
142
168
|
if (is_causal) {
|
|
143
|
-
|
|
144
|
-
if (
|
|
169
|
+
const int causal_limit = n_kv - n_q + my_query_row;
|
|
170
|
+
if (k_row0 > causal_limit) s0 = FA_M_INIT;
|
|
171
|
+
if (k_row1 > causal_limit) s1 = FA_M_INIT;
|
|
172
|
+
if (k_row2 > causal_limit) s2 = FA_M_INIT;
|
|
173
|
+
if (k_row3 > causal_limit) s3 = FA_M_INIT;
|
|
145
174
|
}
|
|
146
|
-
|
|
147
|
-
if (
|
|
148
|
-
if (
|
|
175
|
+
if (k_row0 >= n_kv) s0 = FA_M_INIT;
|
|
176
|
+
if (k_row1 >= n_kv) s1 = FA_M_INIT;
|
|
177
|
+
if (k_row2 >= n_kv) s2 = FA_M_INIT;
|
|
178
|
+
if (k_row3 >= n_kv) s3 = FA_M_INIT;
|
|
149
179
|
|
|
150
180
|
if (mask_base != NULL) {
|
|
151
181
|
const global DATA_TYPE* mask_ptr = (const global DATA_TYPE*)(mask_base + my_query_row * mask_nb1);
|
|
152
|
-
if (k_row0 < n_kv)
|
|
153
|
-
if (k_row1 < n_kv)
|
|
182
|
+
if (k_row0 < n_kv) s0 += slope * (ACC_TYPE)mask_ptr[k_row0];
|
|
183
|
+
if (k_row1 < n_kv) s1 += slope * (ACC_TYPE)mask_ptr[k_row1];
|
|
184
|
+
if (k_row2 < n_kv) s2 += slope * (ACC_TYPE)mask_ptr[k_row2];
|
|
185
|
+
if (k_row3 < n_kv) s3 += slope * (ACC_TYPE)mask_ptr[k_row3];
|
|
154
186
|
}
|
|
155
187
|
|
|
156
188
|
if (logit_softcap > 0.0f) {
|
|
157
|
-
|
|
158
|
-
|
|
189
|
+
s0 = logit_softcap * tanh(s0 / logit_softcap);
|
|
190
|
+
s1 = logit_softcap * tanh(s1 / logit_softcap);
|
|
191
|
+
s2 = logit_softcap * tanh(s2 / logit_softcap);
|
|
192
|
+
s3 = logit_softcap * tanh(s3 / logit_softcap);
|
|
159
193
|
}
|
|
160
194
|
|
|
161
|
-
const ACC_TYPE m_new
|
|
162
|
-
const ACC_TYPE
|
|
163
|
-
const ACC_TYPE
|
|
164
|
-
const ACC_TYPE
|
|
195
|
+
const ACC_TYPE m_new = max(m_i, max(max(s0, s1), max(s2, s3)));
|
|
196
|
+
const ACC_TYPE scale_prev = native_exp(m_i - m_new);
|
|
197
|
+
const ACC_TYPE p0 = native_exp(s0 - m_new);
|
|
198
|
+
const ACC_TYPE p1 = native_exp(s1 - m_new);
|
|
199
|
+
const ACC_TYPE p2 = native_exp(s2 - m_new);
|
|
200
|
+
const ACC_TYPE p3 = native_exp(s3 - m_new);
|
|
165
201
|
|
|
166
|
-
|
|
202
|
+
FA_UNROLL
|
|
167
203
|
for (int i = 0; i < DV_VEC; ++i) {
|
|
168
|
-
o_acc[i] =
|
|
204
|
+
o_acc[i] = mad(p3, CONVERT_ACC4(l_v[j+3][i]),
|
|
205
|
+
mad(p2, CONVERT_ACC4(l_v[j+2][i]),
|
|
206
|
+
mad(p1, CONVERT_ACC4(l_v[j+1][i]),
|
|
207
|
+
mad(p0, CONVERT_ACC4(l_v[j][i]),
|
|
208
|
+
o_acc[i] * scale_prev))));
|
|
169
209
|
}
|
|
170
|
-
l_i = l_i * scale_prev + p0 + p1;
|
|
210
|
+
l_i = l_i * scale_prev + p0 + p1 + p2 + p3;
|
|
171
211
|
m_i = m_new;
|
|
172
212
|
}
|
|
173
213
|
}
|
|
@@ -179,7 +219,7 @@ __kernel void flash_attn_f16(
|
|
|
179
219
|
const ACC_TYPE m_final = max(m_i, m_sink);
|
|
180
220
|
|
|
181
221
|
const ACC_TYPE scale_o = exp(m_i - m_final);
|
|
182
|
-
|
|
222
|
+
FA_UNROLL
|
|
183
223
|
for (int i = 0; i < DV_VEC; ++i) {
|
|
184
224
|
o_acc[i] *= scale_o;
|
|
185
225
|
}
|
|
@@ -191,12 +231,12 @@ __kernel void flash_attn_f16(
|
|
|
191
231
|
global DATA_TYPE4 *o_row = (global DATA_TYPE4 *)(o_base + o_row_offset);
|
|
192
232
|
if (l_i > 0.0f) {
|
|
193
233
|
const ACC_TYPE l_inv = 1.0f / l_i;
|
|
194
|
-
|
|
234
|
+
FA_UNROLL
|
|
195
235
|
for (int i = 0; i < DV_VEC; ++i) {
|
|
196
236
|
o_row[i] = CONVERT_DATA4(o_acc[i] * l_inv);
|
|
197
237
|
}
|
|
198
238
|
} else {
|
|
199
|
-
|
|
239
|
+
FA_UNROLL
|
|
200
240
|
for (int i = 0; i < DV_VEC; ++i) {
|
|
201
241
|
o_row[i] = (DATA_TYPE4)(0.0f);
|
|
202
242
|
}
|
|
@@ -258,7 +298,7 @@ __kernel void flash_attn_f16_q1(
|
|
|
258
298
|
ACC_TYPE4 q_priv[DK_VEC];
|
|
259
299
|
const ulong q_row_offset = batch_idx * q_nb3 + head_idx * q_nb2;
|
|
260
300
|
const global DATA_TYPE4* q_ptr = (const global DATA_TYPE4*)(q_base + q_row_offset);
|
|
261
|
-
|
|
301
|
+
FA_UNROLL
|
|
262
302
|
for (int i = 0; i < DK_VEC; ++i) {
|
|
263
303
|
q_priv[i] = CONVERT_ACC4(q_ptr[i]);
|
|
264
304
|
}
|
|
@@ -270,12 +310,12 @@ __kernel void flash_attn_f16_q1(
|
|
|
270
310
|
sinks_ptr = (const global ACC_TYPE*)((const global char*)sinks_void + sinks_offset);
|
|
271
311
|
}
|
|
272
312
|
|
|
273
|
-
ACC_TYPE m_i = (sinks_ptr != NULL) ? sinks_ptr[head_idx] :
|
|
313
|
+
ACC_TYPE m_i = (sinks_ptr != NULL) ? sinks_ptr[head_idx] : FA_M_INIT;
|
|
274
314
|
for (int k_idx = tid; k_idx < n_kv; k_idx += Q1_WG_SIZE) {
|
|
275
315
|
const ulong k_row_offset = batch_idx * k_nb3 + head_kv_idx * k_nb2 + k_idx * k_nb1;
|
|
276
316
|
const global DATA_TYPE4* k_ptr = (const global DATA_TYPE4*)(k_base + k_row_offset);
|
|
277
317
|
ACC_TYPE4 dot_acc = (ACC_TYPE4)(0.0f);
|
|
278
|
-
|
|
318
|
+
FA_UNROLL
|
|
279
319
|
for (int k = 0; k < DK_VEC; k++) {
|
|
280
320
|
dot_acc = mad(q_priv[k], CONVERT_ACC4(k_ptr[k]), dot_acc);
|
|
281
321
|
}
|
|
@@ -293,7 +333,7 @@ __kernel void flash_attn_f16_q1(
|
|
|
293
333
|
__local ACC_TYPE local_m[Q1_WG_SIZE];
|
|
294
334
|
local_m[tid] = m_i;
|
|
295
335
|
barrier(CLK_LOCAL_MEM_FENCE);
|
|
296
|
-
|
|
336
|
+
FA_UNROLL
|
|
297
337
|
for (int s = Q1_WG_SIZE / 2; s > 0; s >>= 1) {
|
|
298
338
|
if (tid < s) local_m[tid] = max(local_m[tid], local_m[tid + s]);
|
|
299
339
|
barrier(CLK_LOCAL_MEM_FENCE);
|
|
@@ -301,7 +341,7 @@ __kernel void flash_attn_f16_q1(
|
|
|
301
341
|
const ACC_TYPE m_final = local_m[0];
|
|
302
342
|
|
|
303
343
|
ACC_TYPE4 o_acc[DV_VEC];
|
|
304
|
-
|
|
344
|
+
FA_UNROLL
|
|
305
345
|
for (int i = 0; i < DV_VEC; ++i) o_acc[i] = (ACC_TYPE4)(0.0f);
|
|
306
346
|
ACC_TYPE l_i = 0.0f;
|
|
307
347
|
|
|
@@ -311,7 +351,7 @@ __kernel void flash_attn_f16_q1(
|
|
|
311
351
|
const global DATA_TYPE4* k_ptr = (const global DATA_TYPE4*)(k_base + k_row_offset);
|
|
312
352
|
const global DATA_TYPE4* v_ptr = (const global DATA_TYPE4*)(v_base + v_row_offset);
|
|
313
353
|
ACC_TYPE4 dot_acc = (ACC_TYPE4)(0.0f);
|
|
314
|
-
|
|
354
|
+
FA_UNROLL
|
|
315
355
|
for (int k = 0; k < DK_VEC; k++) {
|
|
316
356
|
dot_acc = mad(q_priv[k], CONVERT_ACC4(k_ptr[k]), dot_acc);
|
|
317
357
|
}
|
|
@@ -325,7 +365,7 @@ __kernel void flash_attn_f16_q1(
|
|
|
325
365
|
}
|
|
326
366
|
const ACC_TYPE p = exp(score - m_final);
|
|
327
367
|
l_i += p;
|
|
328
|
-
|
|
368
|
+
FA_UNROLL
|
|
329
369
|
for (int i = 0; i < DV_VEC; i++) {
|
|
330
370
|
o_acc[i] = mad(p, CONVERT_ACC4(v_ptr[i]), o_acc[i]);
|
|
331
371
|
}
|
|
@@ -335,7 +375,7 @@ __kernel void flash_attn_f16_q1(
|
|
|
335
375
|
__local ACC_TYPE4 local_o_comp[Q1_WG_SIZE];
|
|
336
376
|
local_l[tid] = l_i;
|
|
337
377
|
barrier(CLK_LOCAL_MEM_FENCE);
|
|
338
|
-
|
|
378
|
+
FA_UNROLL
|
|
339
379
|
for (int s = Q1_WG_SIZE / 2; s > 0; s >>= 1) {
|
|
340
380
|
if (tid < s) local_l[tid] += local_l[tid + s];
|
|
341
381
|
barrier(CLK_LOCAL_MEM_FENCE);
|
|
@@ -354,7 +394,7 @@ __kernel void flash_attn_f16_q1(
|
|
|
354
394
|
for (int i = 0; i < DV_VEC; i++) {
|
|
355
395
|
local_o_comp[tid] = o_acc[i];
|
|
356
396
|
barrier(CLK_LOCAL_MEM_FENCE);
|
|
357
|
-
|
|
397
|
+
FA_UNROLL
|
|
358
398
|
for (int s = Q1_WG_SIZE / 2; s > 0; s >>= 1) {
|
|
359
399
|
if (tid < s) local_o_comp[tid] += local_o_comp[tid + s];
|
|
360
400
|
barrier(CLK_LOCAL_MEM_FENCE);
|
|
@@ -364,7 +404,7 @@ __kernel void flash_attn_f16_q1(
|
|
|
364
404
|
}
|
|
365
405
|
}
|
|
366
406
|
} else if (tid == 0) {
|
|
367
|
-
|
|
407
|
+
FA_UNROLL
|
|
368
408
|
for (int i = 0; i < DV_VEC; ++i) o_row[i] = (DATA_TYPE4)(0.0f);
|
|
369
409
|
}
|
|
370
410
|
}
|