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
|
@@ -5,35 +5,31 @@ void main() {
|
|
|
5
5
|
return;
|
|
6
6
|
}
|
|
7
7
|
|
|
8
|
-
const uint
|
|
9
|
-
const uint
|
|
8
|
+
const uint i23 = fastdiv(i, p.ne2_012mp, p.ne2_012L);
|
|
9
|
+
const uint i23_offset = i23 * p.ne22*p.ne21*p.ne20;
|
|
10
|
+
const uint i22 = fastdiv(i - i23_offset, p.ne2_01mp, p.ne2_01L);
|
|
11
|
+
const uint i22_offset = i22*p.ne21*p.ne20;
|
|
12
|
+
const uint i21 = fastdiv(i - i23_offset - i22_offset, p.ne2_0mp, p.ne2_0L);
|
|
13
|
+
const uint i20 = i - i23_offset - i22_offset - i21*p.ne20;
|
|
10
14
|
|
|
11
|
-
const uint
|
|
12
|
-
const uint
|
|
13
|
-
const uint
|
|
14
|
-
const uint src_idx = i3 * p.nb03 + i2 * p.nb02 + i1 * p.nb01 + col;
|
|
15
|
-
|
|
16
|
-
const uint dst_i3 = row / (p.ne11 * p.ne12);
|
|
17
|
-
const uint dst_i2 = (row % (p.ne11 * p.ne12)) / p.ne11;
|
|
18
|
-
const uint dst_i1 = row % p.ne11;
|
|
19
|
-
const uint dst_idx = dst_i3 * p.nb13 + dst_i2 * p.nb12 + dst_i1 * p.nb11 + col;
|
|
15
|
+
const uint src_idx_a = get_aoffset() + i23 * p.nb03 + i22 * p.nb02 + i21 * p.nb01 + i20 * p.nb00;
|
|
16
|
+
const uint src_idx_b = get_boffset() + i23 * p.nb13 + i22 * p.nb12 + i21 * p.nb11 + i20 * p.nb10;
|
|
17
|
+
const uint dst_idx = get_doffset() + i23 * p.nb23 + i22 * p.nb22 + i21 * p.nb21 + i20 * p.nb20;
|
|
20
18
|
|
|
21
19
|
if (p.mode == 0) {
|
|
22
20
|
// Default
|
|
23
|
-
const uint offset = p.ne00 / 2;
|
|
24
|
-
const uint idx =
|
|
21
|
+
const uint offset = (p.ne00 / 2) * p.nb00;
|
|
22
|
+
const uint idx = src_idx_a;
|
|
25
23
|
|
|
26
24
|
data_d[dst_idx] = D_TYPE(op(float(data_a[idx]), float(data_a[idx + offset])));
|
|
27
25
|
} else if (p.mode == 1) {
|
|
28
26
|
// Swapped
|
|
29
|
-
const uint offset = p.ne00 / 2;
|
|
30
|
-
const uint idx =
|
|
27
|
+
const uint offset = (p.ne00 / 2) * p.nb00;
|
|
28
|
+
const uint idx = src_idx_a;
|
|
31
29
|
|
|
32
30
|
data_d[dst_idx] = D_TYPE(op(float(data_a[idx + offset]), float(data_a[idx])));
|
|
33
31
|
} else {
|
|
34
32
|
// Split
|
|
35
|
-
|
|
36
|
-
|
|
37
|
-
data_d[dst_idx] = D_TYPE(op(float(data_a[idx]), float(data_b[idx])));
|
|
33
|
+
data_d[dst_idx] = D_TYPE(op(float(data_a[src_idx_a]), float(data_b[src_idx_b])));
|
|
38
34
|
}
|
|
39
35
|
}
|
|
@@ -14,16 +14,13 @@ void main() {
|
|
|
14
14
|
const uint row = gl_WorkGroupID.z * 262144 + gl_WorkGroupID.y * 512 + gl_WorkGroupID.x;
|
|
15
15
|
const uint tid = gl_LocalInvocationID.x;
|
|
16
16
|
|
|
17
|
-
const uint
|
|
18
|
-
const uint
|
|
19
|
-
const uint i2 = (row - i3_offset) / p.ne11;
|
|
20
|
-
const uint i2_offset = i2 * p.ne11;
|
|
21
|
-
const uint i1 = row - i3_offset - i2_offset;
|
|
17
|
+
const uint a_base = get_aoffset() + src0_idx(row * p.ne00);
|
|
18
|
+
const uint d_base = get_doffset() + dst_idx(row * p.ne10);
|
|
22
19
|
|
|
23
20
|
sum[tid] = FLOAT_TYPE(0.0f); // partial sum for thread in warp
|
|
24
21
|
|
|
25
22
|
[[unroll]] for (uint i0 = tid; i0 < p.ne00; i0 += BLOCK_SIZE) {
|
|
26
|
-
const FLOAT_TYPE xi = FLOAT_TYPE(data_a[
|
|
23
|
+
const FLOAT_TYPE xi = FLOAT_TYPE(data_a[a_base + i0*p.nb00]);
|
|
27
24
|
sum[tid] += xi * xi;
|
|
28
25
|
}
|
|
29
26
|
|
|
@@ -39,6 +36,6 @@ void main() {
|
|
|
39
36
|
const FLOAT_TYPE scale = 1.0f / max(sqrt(sum[0]), FLOAT_TYPE(p.param1));
|
|
40
37
|
|
|
41
38
|
[[unroll]] for (uint i0 = tid; i0 < p.ne00; i0 += BLOCK_SIZE) {
|
|
42
|
-
data_d[
|
|
39
|
+
data_d[d_base + i0*p.nb10] = D_TYPE(scale * FLOAT_TYPE(data_a[a_base + i0*p.nb00]));
|
|
43
40
|
}
|
|
44
41
|
}
|
|
@@ -28,13 +28,10 @@ vec2 cache_b_ds;
|
|
|
28
28
|
|
|
29
29
|
#include "mul_mat_vecq_funcs.glsl"
|
|
30
30
|
|
|
31
|
-
void iter(inout FLOAT_TYPE temp[NUM_COLS][NUM_ROWS], const uint first_row, const uint num_rows, const uint
|
|
31
|
+
void iter(inout FLOAT_TYPE temp[NUM_COLS][NUM_ROWS], const uint first_row, const uint num_rows, const uint col, const uint b_qs_idx) {
|
|
32
32
|
[[unroll]] for (uint j = 0; j < NUM_COLS; ++j) {
|
|
33
|
-
const uint col = i*BLOCK_SIZE + tid*K_PER_ITER;
|
|
34
|
-
|
|
35
33
|
// Preload data_b block
|
|
36
34
|
const uint b_block_idx = (j*p.batch_stride_b + col) / QUANT_K_Q8_1 + b_offset;
|
|
37
|
-
const uint b_qs_idx = tid % (32 / K_PER_ITER);
|
|
38
35
|
const uint b_block_idx_outer = b_block_idx / 4;
|
|
39
36
|
const uint b_block_idx_inner = b_block_idx % 4;
|
|
40
37
|
cache_b_ds = vec2(data_b[b_block_idx_outer].ds[b_block_idx_inner]);
|
|
@@ -91,35 +88,35 @@ void compute_outputs(const uint32_t first_row, const uint32_t num_rows) {
|
|
|
91
88
|
}
|
|
92
89
|
}
|
|
93
90
|
|
|
94
|
-
uint
|
|
95
|
-
|
|
91
|
+
const uint col_stride = K_PER_ITER * BLOCK_SIZE;
|
|
92
|
+
uint num_iters = p.ncols / col_stride;
|
|
93
|
+
if (num_iters * col_stride + K_PER_ITER * tid < p.ncols) {
|
|
96
94
|
num_iters++;
|
|
97
95
|
}
|
|
98
|
-
int unroll_count = 4;
|
|
99
|
-
uint unrolled_iters = num_iters & ~(unroll_count - 1);
|
|
100
96
|
|
|
101
|
-
uint
|
|
102
|
-
|
|
97
|
+
const uint b_qs_idx = tid % (32 / K_PER_ITER);
|
|
98
|
+
uint col = tid * K_PER_ITER;
|
|
99
|
+
while (num_iters >= 4) {
|
|
103
100
|
// Manually partially unroll the loop
|
|
104
|
-
[[unroll]] for (uint k = 0; k <
|
|
105
|
-
iter(temp, first_row, num_rows,
|
|
106
|
-
|
|
101
|
+
[[unroll]] for (uint k = 0; k < 4; ++k) {
|
|
102
|
+
iter(temp, first_row, num_rows, col, b_qs_idx);
|
|
103
|
+
col += col_stride;
|
|
107
104
|
}
|
|
108
|
-
}
|
|
109
105
|
|
|
110
|
-
|
|
111
|
-
|
|
106
|
+
num_iters -= 4;
|
|
107
|
+
}
|
|
112
108
|
|
|
113
|
-
|
|
109
|
+
if (num_iters >= 2) {
|
|
114
110
|
// Manually partially unroll the loop
|
|
115
|
-
|
|
116
|
-
|
|
117
|
-
|
|
118
|
-
|
|
111
|
+
iter(temp, first_row, num_rows, col, b_qs_idx);
|
|
112
|
+
col += col_stride;
|
|
113
|
+
iter(temp, first_row, num_rows, col, b_qs_idx);
|
|
114
|
+
col += col_stride;
|
|
115
|
+
num_iters -= 2;
|
|
119
116
|
}
|
|
120
|
-
|
|
121
|
-
|
|
122
|
-
|
|
117
|
+
|
|
118
|
+
if (num_iters > 0) {
|
|
119
|
+
iter(temp, first_row, num_rows, col, b_qs_idx);
|
|
123
120
|
}
|
|
124
121
|
|
|
125
122
|
reduce_result(temp, d_offset, first_row, num_rows, tid);
|
|
@@ -38,17 +38,7 @@
|
|
|
38
38
|
#define LOAD_VEC_B 1
|
|
39
39
|
#endif
|
|
40
40
|
|
|
41
|
-
|
|
42
|
-
#if (defined(DATA_A_F32) || defined(DATA_A_F16) || defined(DATA_A_BF16)) && !defined(ALIGNED)
|
|
43
|
-
#define LOAD_VEC_BATCH_A 2
|
|
44
|
-
#else
|
|
45
|
-
#define LOAD_VEC_BATCH_A 1
|
|
46
|
-
#endif
|
|
47
|
-
#if !defined(ALIGNED)
|
|
48
|
-
#define LOAD_VEC_BATCH_B 2
|
|
49
|
-
#else
|
|
50
|
-
#define LOAD_VEC_BATCH_B 1
|
|
51
|
-
#endif
|
|
41
|
+
layout (constant_id = 11) const uint ALIGNED = 0;
|
|
52
42
|
|
|
53
43
|
#if !defined(TO_FLOAT_TYPE)
|
|
54
44
|
#define TO_FLOAT_TYPE FLOAT_TYPE
|
|
@@ -57,6 +47,13 @@
|
|
|
57
47
|
layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in;
|
|
58
48
|
|
|
59
49
|
layout (binding = 0) readonly buffer A {A_TYPE data_a[];};
|
|
50
|
+
#if defined(DATA_A_F32)
|
|
51
|
+
layout (binding = 0) readonly buffer A_SCALAR {float data_a_scalar[];};
|
|
52
|
+
#elif defined(DATA_A_F16)
|
|
53
|
+
layout (binding = 0) readonly buffer A_SCALAR {float16_t data_a_scalar[];};
|
|
54
|
+
#elif defined(DATA_A_BF16)
|
|
55
|
+
layout (binding = 0) readonly buffer A_SCALAR {uint16_t data_a_scalar[];};
|
|
56
|
+
#endif
|
|
60
57
|
#if defined(A_TYPE_PACKED16)
|
|
61
58
|
layout (binding = 0) readonly buffer A_PACKED16 {A_TYPE_PACKED16 data_a_packed16[];};
|
|
62
59
|
#endif
|
|
@@ -65,6 +62,7 @@ layout (binding = 0) readonly buffer A_PACKED32 {A_TYPE_PACKED32 data_a_packed32
|
|
|
65
62
|
#endif
|
|
66
63
|
|
|
67
64
|
layout (binding = 1) readonly buffer B {B_TYPE data_b[];};
|
|
65
|
+
layout (binding = 1) readonly buffer B_SCALAR {B_TYPE_SCALAR data_b_scalar[];};
|
|
68
66
|
layout (binding = 2) writeonly buffer D {D_TYPE data_d[];};
|
|
69
67
|
|
|
70
68
|
#ifdef MUL_MAT_ID
|
|
@@ -194,13 +192,23 @@ void main() {
|
|
|
194
192
|
const uint warp_r = warp_i % (BM / WM);
|
|
195
193
|
const uint warp_c = warp_i / (BM / WM);
|
|
196
194
|
|
|
197
|
-
|
|
198
|
-
const uint
|
|
199
|
-
const uint
|
|
200
|
-
|
|
195
|
+
#if defined(DATA_A_F32) || defined(DATA_A_F16) || defined(DATA_A_BF16)
|
|
196
|
+
const uint LOAD_VEC_A_EFF = (ALIGNED != 0) ? LOAD_VEC_A : 1;
|
|
197
|
+
const uint LOAD_VEC_BATCH_A = (ALIGNED != 0) ? 1 : 2;
|
|
198
|
+
#else
|
|
199
|
+
const uint LOAD_VEC_A_EFF = LOAD_VEC_A;
|
|
200
|
+
const uint LOAD_VEC_BATCH_A = 1;
|
|
201
|
+
#endif
|
|
202
|
+
const uint LOAD_VEC_B_EFF = (ALIGNED != 0) ? LOAD_VEC_B : 1;
|
|
203
|
+
const uint LOAD_VEC_BATCH_B = (ALIGNED != 0) ? 1 : 2;
|
|
204
|
+
|
|
205
|
+
const uint loadr_a = gl_LocalInvocationID.x % (BK / LOAD_VEC_A_EFF / LOAD_VEC_BATCH_A);
|
|
206
|
+
const uint loadc_a = gl_LocalInvocationID.x / (BK / LOAD_VEC_A_EFF / LOAD_VEC_BATCH_A);
|
|
207
|
+
const uint loadr_b = gl_LocalInvocationID.x % (BK / LOAD_VEC_B_EFF / LOAD_VEC_BATCH_B);
|
|
208
|
+
const uint loadc_b = gl_LocalInvocationID.x / (BK / LOAD_VEC_B_EFF / LOAD_VEC_BATCH_B);
|
|
201
209
|
|
|
202
|
-
const uint loadstride_a = gl_WorkGroupSize.x *
|
|
203
|
-
const uint loadstride_b = gl_WorkGroupSize.x *
|
|
210
|
+
const uint loadstride_a = gl_WorkGroupSize.x * LOAD_VEC_A_EFF * LOAD_VEC_BATCH_A / BK;
|
|
211
|
+
const uint loadstride_b = gl_WorkGroupSize.x * LOAD_VEC_B_EFF * LOAD_VEC_BATCH_B / BK;
|
|
204
212
|
|
|
205
213
|
#ifdef MUL_MAT_ID
|
|
206
214
|
#ifdef MUL_MAT_ID_USE_SUBGROUPS
|
|
@@ -239,15 +247,15 @@ void main() {
|
|
|
239
247
|
|
|
240
248
|
uint pos_a =
|
|
241
249
|
#ifdef MUL_MAT_ID
|
|
242
|
-
expert_idx * (p.batch_stride_a /
|
|
250
|
+
expert_idx * (p.batch_stride_a / LOAD_VEC_A_EFF) +
|
|
243
251
|
#else
|
|
244
|
-
batch_idx_a * (p.batch_stride_a /
|
|
252
|
+
batch_idx_a * (p.batch_stride_a / LOAD_VEC_A_EFF) +
|
|
245
253
|
#endif
|
|
246
|
-
(ir * BM * p.stride_a + start_k) /
|
|
254
|
+
(ir * BM * p.stride_a + start_k) / LOAD_VEC_A_EFF;
|
|
247
255
|
#ifdef MUL_MAT_ID
|
|
248
256
|
uint pos_b = 0;
|
|
249
257
|
#else
|
|
250
|
-
uint pos_b = (batch_idx * p.batch_stride_b + ic * BN * p.stride_b + start_k) /
|
|
258
|
+
uint pos_b = (batch_idx * p.batch_stride_b + ic * BN * p.stride_b + start_k) / LOAD_VEC_B_EFF;
|
|
251
259
|
#endif
|
|
252
260
|
|
|
253
261
|
#ifdef COOPMAT
|
|
@@ -287,8 +295,8 @@ void main() {
|
|
|
287
295
|
|
|
288
296
|
barrier();
|
|
289
297
|
|
|
290
|
-
pos_a += BK /
|
|
291
|
-
pos_b += BK /
|
|
298
|
+
pos_a += BK / LOAD_VEC_A_EFF;
|
|
299
|
+
pos_b += BK / LOAD_VEC_B_EFF;
|
|
292
300
|
|
|
293
301
|
#ifdef COOPMAT
|
|
294
302
|
[[unroll]] for (uint i = 0; i < BK; i += TK) {
|
|
@@ -36,6 +36,7 @@ layout (constant_id = 3) const uint BK = 16; // Assumed to be 32 if working wit
|
|
|
36
36
|
layout (constant_id = 4) const bool enable_smaller_matrices = false;
|
|
37
37
|
const uint BNover2 = enable_smaller_matrices ? (BN / 2) : BN;
|
|
38
38
|
const uint BNover4 = enable_smaller_matrices ? (BN / 4) : BN;
|
|
39
|
+
layout (constant_id = 5) const uint ALIGNED = 0;
|
|
39
40
|
|
|
40
41
|
layout (push_constant) uniform parameter
|
|
41
42
|
{
|
|
@@ -111,7 +112,7 @@ layout(buffer_reference, std430, buffer_reference_align = 2) buffer decodeBufB {
|
|
|
111
112
|
};
|
|
112
113
|
|
|
113
114
|
uint _ne1;
|
|
114
|
-
layout (constant_id =
|
|
115
|
+
layout (constant_id = 6) const uint subgroup_size = 32;
|
|
115
116
|
shared uvec4 ballots_sh[BLOCK_SIZE / subgroup_size];
|
|
116
117
|
|
|
117
118
|
B_TYPE decodeFuncB(const in decodeBufB bl, const in uint blockCoords[2], const in uint coordInBlock[2])
|
|
@@ -297,12 +298,12 @@ void main() {
|
|
|
297
298
|
|
|
298
299
|
// Hint to the compiler that values are aligned (want 16B alignment).
|
|
299
300
|
// Quants are always block-aligned, no alignment needed.
|
|
300
|
-
|
|
301
|
+
if (ALIGNED != 0) {
|
|
301
302
|
#if QUANT_K == 1
|
|
302
|
-
|
|
303
|
-
#endif
|
|
304
|
-
stride_b &= ~7;
|
|
303
|
+
stride_a &= ~7;
|
|
305
304
|
#endif
|
|
305
|
+
stride_b &= ~7;
|
|
306
|
+
}
|
|
306
307
|
|
|
307
308
|
// Create layouts for both clamped and unclamped accesses
|
|
308
309
|
tensorLayoutNV<2> tensorLayoutA = createTensorLayoutNV(2);
|
|
@@ -1,50 +1,57 @@
|
|
|
1
1
|
void load_a_to_shmem(const uint pos_a, const uint row, const uint col, const uint idx_m, const uint block, const uint end_k) {
|
|
2
2
|
#if defined(DATA_A_F32) || defined(DATA_A_F16)
|
|
3
3
|
#if LOAD_VEC_A == 8
|
|
4
|
-
|
|
5
|
-
|
|
6
|
-
|
|
7
|
-
|
|
8
|
-
|
|
9
|
-
|
|
10
|
-
|
|
4
|
+
if (ALIGNED != 0) {
|
|
5
|
+
const uint idx = pos_a + col * p.stride_a / LOAD_VEC_A + row;
|
|
6
|
+
const uint buf_idx = col * SHMEM_STRIDE + row * LOAD_VEC_A / 2;
|
|
7
|
+
FLOAT_TYPEV8 aa = FLOAT_TYPEV8(data_a[idx]);
|
|
8
|
+
buf_a[buf_idx ] = aa[0].xy;
|
|
9
|
+
buf_a[buf_idx + 1] = aa[0].zw;
|
|
10
|
+
buf_a[buf_idx + 2] = aa[1].xy;
|
|
11
|
+
buf_a[buf_idx + 3] = aa[1].zw;
|
|
12
|
+
return;
|
|
13
|
+
}
|
|
11
14
|
#elif LOAD_VEC_A == 4
|
|
12
|
-
|
|
13
|
-
|
|
14
|
-
|
|
15
|
-
|
|
16
|
-
|
|
17
|
-
|
|
15
|
+
if (ALIGNED != 0) {
|
|
16
|
+
const uint idx = pos_a + col * p.stride_a / LOAD_VEC_A + row;
|
|
17
|
+
const uint buf_idx = col * SHMEM_STRIDE + row * LOAD_VEC_A / 2;
|
|
18
|
+
FLOAT_TYPEV4 aa = FLOAT_TYPEV4(data_a[idx]);
|
|
19
|
+
buf_a[buf_idx ] = aa.xy;
|
|
20
|
+
buf_a[buf_idx + 1] = aa.zw;
|
|
21
|
+
return;
|
|
22
|
+
}
|
|
23
|
+
#endif
|
|
18
24
|
const uint idx = pos_a + col * p.stride_a + row * 2;
|
|
19
25
|
const uint buf_idx = col * SHMEM_STRIDE + row;
|
|
20
26
|
if (idx_m < p.M && block + row * 2 + 1 < end_k) {
|
|
21
|
-
buf_a[buf_idx] = FLOAT_TYPEV2(
|
|
22
|
-
|
|
27
|
+
buf_a[buf_idx] = FLOAT_TYPEV2(data_a_scalar[idx],
|
|
28
|
+
data_a_scalar[idx + 1]);
|
|
23
29
|
} else if (idx_m < p.M && block + row * 2 < end_k) {
|
|
24
|
-
buf_a[buf_idx] = FLOAT_TYPEV2(
|
|
30
|
+
buf_a[buf_idx] = FLOAT_TYPEV2(data_a_scalar[idx], 0.0f);
|
|
25
31
|
} else {
|
|
26
32
|
buf_a[buf_idx] = FLOAT_TYPEV2(0.0f);
|
|
27
33
|
}
|
|
28
|
-
#endif
|
|
29
34
|
#elif defined(DATA_A_BF16)
|
|
30
35
|
#if LOAD_VEC_A == 4
|
|
31
|
-
|
|
32
|
-
|
|
33
|
-
|
|
34
|
-
|
|
35
|
-
|
|
36
|
-
|
|
36
|
+
if (ALIGNED != 0) {
|
|
37
|
+
const uint idx = pos_a + col * p.stride_a / LOAD_VEC_A + row;
|
|
38
|
+
const uint buf_idx = col * SHMEM_STRIDE + row * LOAD_VEC_A / 2;
|
|
39
|
+
FLOAT_TYPEV4 aa = FLOAT_TYPEV4(TO_FLOAT_TYPE(data_a[idx]));
|
|
40
|
+
buf_a[buf_idx ] = aa.xy;
|
|
41
|
+
buf_a[buf_idx + 1] = aa.zw;
|
|
42
|
+
return;
|
|
43
|
+
}
|
|
44
|
+
#endif
|
|
37
45
|
const uint idx = pos_a + col * p.stride_a + row * 2;
|
|
38
46
|
const uint buf_idx = col * SHMEM_STRIDE + row;
|
|
39
47
|
if (idx_m < p.M && block + row * 2 + 1 < end_k) {
|
|
40
|
-
buf_a[buf_idx] = FLOAT_TYPEV2(TO_FLOAT_TYPE(
|
|
41
|
-
TO_FLOAT_TYPE(
|
|
48
|
+
buf_a[buf_idx] = FLOAT_TYPEV2(TO_FLOAT_TYPE(data_a_scalar[idx]),
|
|
49
|
+
TO_FLOAT_TYPE(data_a_scalar[idx + 1]));
|
|
42
50
|
} else if (idx_m < p.M && block + row * 2 < end_k) {
|
|
43
|
-
buf_a[buf_idx] = FLOAT_TYPEV2(TO_FLOAT_TYPE(
|
|
51
|
+
buf_a[buf_idx] = FLOAT_TYPEV2(TO_FLOAT_TYPE(data_a_scalar[idx]), 0.0f);
|
|
44
52
|
} else {
|
|
45
53
|
buf_a[buf_idx] = FLOAT_TYPEV2(0.0f);
|
|
46
54
|
}
|
|
47
|
-
#endif
|
|
48
55
|
#elif defined(DATA_A_Q4_0)
|
|
49
56
|
const uint idx = pos_a + col * p.stride_a / LOAD_VEC_A + row;
|
|
50
57
|
const uint buf_idx = col * SHMEM_STRIDE + row * LOAD_VEC_A / 4;
|
|
@@ -526,75 +533,85 @@ void load_a_to_shmem(const uint pos_a, const uint row, const uint col, const uin
|
|
|
526
533
|
#if !defined(MUL_MAT_ID)
|
|
527
534
|
void load_b_to_shmem(const uint pos_b, const uint row, const uint col, const uint idx_n, const uint block, const uint end_k) {
|
|
528
535
|
#if LOAD_VEC_B == 8
|
|
529
|
-
|
|
530
|
-
|
|
531
|
-
|
|
532
|
-
|
|
533
|
-
|
|
534
|
-
|
|
535
|
-
|
|
536
|
-
|
|
536
|
+
if (ALIGNED != 0) {
|
|
537
|
+
// Not supported for b_type bf16 because bf16mat2x4 does not exist
|
|
538
|
+
const uint idx = pos_b + col * p.stride_b / LOAD_VEC_B + row;
|
|
539
|
+
const uint buf_idx = col * SHMEM_STRIDE + row * LOAD_VEC_B / 2;
|
|
540
|
+
FLOAT_TYPEV8 bb = FLOAT_TYPEV8(data_b[idx]);
|
|
541
|
+
buf_b[buf_idx + 0] = bb[0].xy;
|
|
542
|
+
buf_b[buf_idx + 1] = bb[0].zw;
|
|
543
|
+
buf_b[buf_idx + 2] = bb[1].xy;
|
|
544
|
+
buf_b[buf_idx + 3] = bb[1].zw;
|
|
545
|
+
return;
|
|
546
|
+
}
|
|
537
547
|
#elif LOAD_VEC_B == 4
|
|
538
|
-
|
|
539
|
-
|
|
548
|
+
if (ALIGNED != 0) {
|
|
549
|
+
const uint idx = pos_b + col * p.stride_b / LOAD_VEC_B + row;
|
|
550
|
+
const uint buf_idx = col * SHMEM_STRIDE + row * LOAD_VEC_B / 2;
|
|
540
551
|
#if defined(DATA_B_BF16)
|
|
541
|
-
|
|
552
|
+
FLOAT_TYPEV4 bb = FLOAT_TYPEV4(TO_FLOAT_TYPE(data_b[idx]));
|
|
542
553
|
#else
|
|
543
|
-
|
|
554
|
+
FLOAT_TYPEV4 bb = FLOAT_TYPEV4(data_b[idx]);
|
|
555
|
+
#endif
|
|
556
|
+
buf_b[buf_idx + 0] = bb.xy;
|
|
557
|
+
buf_b[buf_idx + 1] = bb.zw;
|
|
558
|
+
return;
|
|
559
|
+
}
|
|
544
560
|
#endif
|
|
545
|
-
buf_b[buf_idx + 0] = bb.xy;
|
|
546
|
-
buf_b[buf_idx + 1] = bb.zw;
|
|
547
|
-
#else // LOAD_VEC_BATCH_B == 2
|
|
548
561
|
const uint idx = pos_b + col * p.stride_b + row * 2;
|
|
549
562
|
const uint buf_idx = col * SHMEM_STRIDE + row;
|
|
550
563
|
if (idx_n < p.N && block + row * 2 + 1 < end_k) {
|
|
551
|
-
buf_b[buf_idx] = FLOAT_TYPEV2(TO_FLOAT_TYPE(
|
|
552
|
-
TO_FLOAT_TYPE(
|
|
564
|
+
buf_b[buf_idx] = FLOAT_TYPEV2(TO_FLOAT_TYPE(data_b_scalar[idx]),
|
|
565
|
+
TO_FLOAT_TYPE(data_b_scalar[idx + 1]));
|
|
553
566
|
} else if (idx_n < p.N && block + row * 2 < end_k) {
|
|
554
|
-
buf_b[buf_idx] = FLOAT_TYPEV2(TO_FLOAT_TYPE(
|
|
567
|
+
buf_b[buf_idx] = FLOAT_TYPEV2(TO_FLOAT_TYPE(data_b_scalar[idx]), 0.0f);
|
|
555
568
|
} else {
|
|
556
569
|
buf_b[buf_idx] = FLOAT_TYPEV2(0.0f);
|
|
557
570
|
}
|
|
558
|
-
#endif
|
|
559
571
|
}
|
|
560
572
|
#else
|
|
561
573
|
void load_b_to_shmem(const uint pos_b, const uint row, const uint col, const uint ic, const uint _ne1, const uint block, const uint end_k) {
|
|
562
574
|
#if LOAD_VEC_B == 8
|
|
563
|
-
|
|
564
|
-
|
|
565
|
-
|
|
566
|
-
|
|
567
|
-
|
|
568
|
-
|
|
569
|
-
|
|
570
|
-
|
|
571
|
-
|
|
575
|
+
if (ALIGNED != 0) {
|
|
576
|
+
// Not supported for b_type bf16 because bf16mat2x4 does not exist
|
|
577
|
+
const u16vec2 row_idx = row_ids[col];
|
|
578
|
+
const uint idx = pos_b + row_idx.y * p.batch_stride_b / LOAD_VEC_B + (row_idx.x % p.ne11) * p.stride_b / LOAD_VEC_B + row;
|
|
579
|
+
const uint buf_idx = col * SHMEM_STRIDE + row * LOAD_VEC_B / 2;
|
|
580
|
+
FLOAT_TYPEV8 bb = FLOAT_TYPEV8(data_b[idx]);
|
|
581
|
+
buf_b[buf_idx + 0] = bb[0].xy;
|
|
582
|
+
buf_b[buf_idx + 1] = bb[0].zw;
|
|
583
|
+
buf_b[buf_idx + 2] = bb[1].xy;
|
|
584
|
+
buf_b[buf_idx + 3] = bb[1].zw;
|
|
585
|
+
return;
|
|
586
|
+
}
|
|
572
587
|
#elif LOAD_VEC_B == 4
|
|
573
|
-
|
|
574
|
-
|
|
575
|
-
|
|
588
|
+
if (ALIGNED != 0) {
|
|
589
|
+
const u16vec2 row_idx = row_ids[col];
|
|
590
|
+
const uint idx = pos_b + row_idx.y * p.batch_stride_b / LOAD_VEC_B + (row_idx.x % p.ne11) * p.stride_b / LOAD_VEC_B + row;
|
|
591
|
+
const uint buf_idx = col * SHMEM_STRIDE + row * LOAD_VEC_B / 2;
|
|
576
592
|
#if defined(DATA_B_BF16)
|
|
577
|
-
|
|
593
|
+
FLOAT_TYPEV4 bb = FLOAT_TYPEV4(TO_FLOAT_TYPE(data_b[idx]));
|
|
578
594
|
#else
|
|
579
|
-
|
|
595
|
+
FLOAT_TYPEV4 bb = FLOAT_TYPEV4(data_b[idx]);
|
|
596
|
+
#endif
|
|
597
|
+
buf_b[buf_idx + 0] = bb.xy;
|
|
598
|
+
buf_b[buf_idx + 1] = bb.zw;
|
|
599
|
+
return;
|
|
600
|
+
}
|
|
580
601
|
#endif
|
|
581
|
-
buf_b[buf_idx + 0] = bb.xy;
|
|
582
|
-
buf_b[buf_idx + 1] = bb.zw;
|
|
583
|
-
#else // LOAD_VEC_BATCH_B == 2
|
|
584
602
|
const uint row_i = ic * BN + col;
|
|
585
603
|
const uint buf_idx = col * SHMEM_STRIDE + row;
|
|
586
604
|
if (row_i < _ne1 && block + row * 2 + 1 < end_k) {
|
|
587
605
|
const u16vec2 row_idx = row_ids[col];
|
|
588
606
|
const uint idx = pos_b + row_idx.y * p.batch_stride_b + (row_idx.x % p.ne11) * p.stride_b + row * 2;
|
|
589
|
-
buf_b[buf_idx] = FLOAT_TYPEV2(TO_FLOAT_TYPE(
|
|
590
|
-
TO_FLOAT_TYPE(
|
|
607
|
+
buf_b[buf_idx] = FLOAT_TYPEV2(TO_FLOAT_TYPE(data_b_scalar[idx]),
|
|
608
|
+
TO_FLOAT_TYPE(data_b_scalar[idx + 1]));
|
|
591
609
|
} else if (row_i < _ne1 && block + row * 2 < end_k) {
|
|
592
610
|
const u16vec2 row_idx = row_ids[col];
|
|
593
611
|
const uint idx = pos_b + row_idx.y * p.batch_stride_b + (row_idx.x % p.ne11) * p.stride_b + row * 2;
|
|
594
|
-
buf_b[buf_idx] = FLOAT_TYPEV2(TO_FLOAT_TYPE(
|
|
612
|
+
buf_b[buf_idx] = FLOAT_TYPEV2(TO_FLOAT_TYPE(data_b_scalar[idx]), 0.0f);
|
|
595
613
|
} else {
|
|
596
614
|
buf_b[buf_idx] = FLOAT_TYPEV2(0.0f);
|
|
597
615
|
}
|
|
598
|
-
#endif
|
|
599
616
|
}
|
|
600
617
|
#endif
|
|
@@ -1,26 +1,26 @@
|
|
|
1
1
|
#version 450
|
|
2
2
|
|
|
3
|
-
#include "generic_head.glsl"
|
|
4
3
|
#include "types.glsl"
|
|
4
|
+
#include "generic_unary_head.glsl"
|
|
5
5
|
|
|
6
6
|
#extension GL_EXT_control_flow_attributes : enable
|
|
7
7
|
#define BLOCK_SIZE 512
|
|
8
8
|
|
|
9
9
|
layout(local_size_x = BLOCK_SIZE, local_size_y = 1, local_size_z = 1) in;
|
|
10
10
|
|
|
11
|
-
layout (binding = 0) readonly buffer X {A_TYPE data_a[];};
|
|
12
|
-
layout (binding = 1) writeonly buffer D {D_TYPE data_d[];};
|
|
13
|
-
|
|
14
11
|
shared vec2 sum[BLOCK_SIZE];
|
|
15
12
|
|
|
16
13
|
void main() {
|
|
17
14
|
const uint row = gl_WorkGroupID.z * 262144 + gl_WorkGroupID.y * 512 + gl_WorkGroupID.x;
|
|
18
15
|
const uint tid = gl_LocalInvocationID.x;
|
|
19
16
|
|
|
17
|
+
const uint a_base = get_aoffset() + src0_idx(row * p.ne00);
|
|
18
|
+
const uint d_base = get_doffset() + dst_idx(row * p.ne10);
|
|
19
|
+
|
|
20
20
|
sum[tid] = vec2(0.0f, 0.0f);
|
|
21
21
|
|
|
22
|
-
[[unroll]] for (uint
|
|
23
|
-
const float xi = float(data_a[
|
|
22
|
+
[[unroll]] for (uint i0 = tid; i0 < p.ne00; i0 += BLOCK_SIZE) {
|
|
23
|
+
const float xi = float(data_a[a_base + i0*p.nb00]);
|
|
24
24
|
sum[tid].x += xi;
|
|
25
25
|
sum[tid].y += xi * xi;
|
|
26
26
|
}
|
|
@@ -34,11 +34,11 @@ void main() {
|
|
|
34
34
|
barrier();
|
|
35
35
|
}
|
|
36
36
|
|
|
37
|
-
const float mean = sum[0].x / p.
|
|
38
|
-
const float var = sum[0].y / p.
|
|
37
|
+
const float mean = sum[0].x / p.ne00;
|
|
38
|
+
const float var = sum[0].y / p.ne00 - mean * mean;
|
|
39
39
|
const float inv_std = inversesqrt(var + p.param1);
|
|
40
40
|
|
|
41
|
-
[[unroll]] for (uint
|
|
42
|
-
data_d[
|
|
41
|
+
[[unroll]] for (uint i0 = tid; i0 < p.ne00; i0 += BLOCK_SIZE) {
|
|
42
|
+
data_d[d_base + i0*p.nb10] = D_TYPE((float(data_a[a_base + i0*p.nb00]) - mean) * inv_std);
|
|
43
43
|
}
|
|
44
44
|
}
|
|
@@ -13,11 +13,11 @@ void main() {
|
|
|
13
13
|
}
|
|
14
14
|
|
|
15
15
|
// Destination multi-index (inlined dst_idx)
|
|
16
|
-
const uint i13 = fastdiv(idx, p.ne1_012mp, p.
|
|
16
|
+
const uint i13 = fastdiv(idx, p.ne1_012mp, fastdiv_L(p.ne1_Ls, 0));
|
|
17
17
|
const uint i13_offset = i13 * p.ne12*p.ne11*p.ne10;
|
|
18
|
-
const uint i12 = fastdiv(idx - i13_offset, p.ne1_01mp, p.
|
|
18
|
+
const uint i12 = fastdiv(idx - i13_offset, p.ne1_01mp, fastdiv_L(p.ne1_Ls, 1));
|
|
19
19
|
const uint i12_offset = i12*p.ne11*p.ne10;
|
|
20
|
-
const uint i11 = fastdiv(idx - i13_offset - i12_offset, p.ne1_0mp, p.
|
|
20
|
+
const uint i11 = fastdiv(idx - i13_offset - i12_offset, p.ne1_0mp, fastdiv_L(p.ne1_Ls, 2));
|
|
21
21
|
const uint i10 = idx - i13_offset - i12_offset - i11*p.ne10;
|
|
22
22
|
const uint d_idx = i13*p.nb13 + i12*p.nb12 + i11*p.nb11 + i10*p.nb10;
|
|
23
23
|
|
|
@@ -20,11 +20,11 @@ void main() {
|
|
|
20
20
|
return;
|
|
21
21
|
}
|
|
22
22
|
|
|
23
|
-
const uint i3 = fastdiv(idx, p.ne1_012mp, p.
|
|
23
|
+
const uint i3 = fastdiv(idx, p.ne1_012mp, fastdiv_L(p.ne1_Ls, 0));
|
|
24
24
|
const uint i3_offset = i3 * p.ne12*p.ne11*p.ne10;
|
|
25
|
-
const uint i2 = fastdiv(idx - i3_offset, p.ne1_01mp, p.
|
|
25
|
+
const uint i2 = fastdiv(idx - i3_offset, p.ne1_01mp, fastdiv_L(p.ne1_Ls, 1));
|
|
26
26
|
const uint i2_offset = i2*p.ne11*p.ne10;
|
|
27
|
-
const uint i1 = fastdiv(idx - i3_offset - i2_offset, p.ne1_0mp, p.
|
|
27
|
+
const uint i1 = fastdiv(idx - i3_offset - i2_offset, p.ne1_0mp, fastdiv_L(p.ne1_Ls, 2));
|
|
28
28
|
const uint i0 = idx - i3_offset - i2_offset - i1*p.ne10;
|
|
29
29
|
|
|
30
30
|
const uint p1 = floatBitsToUint(p.param1);
|
|
@@ -17,11 +17,11 @@ void main() {
|
|
|
17
17
|
return;
|
|
18
18
|
}
|
|
19
19
|
|
|
20
|
-
const uint i03 = fastdiv(idx, p.ne0_012mp, p.
|
|
20
|
+
const uint i03 = fastdiv(idx, p.ne0_012mp, fastdiv_L(p.ne0_Ls, 0));
|
|
21
21
|
const uint i03_offset = i03 * p.ne02*p.ne01*p.ne00;
|
|
22
|
-
const uint i02 = fastdiv(idx - i03_offset, p.ne0_01mp, p.
|
|
22
|
+
const uint i02 = fastdiv(idx - i03_offset, p.ne0_01mp, fastdiv_L(p.ne0_Ls, 1));
|
|
23
23
|
const uint i02_offset = i02*p.ne01*p.ne00;
|
|
24
|
-
const uint i01 = fastdiv(idx - i03_offset - i02_offset, p.ne0_0mp, p.
|
|
24
|
+
const uint i01 = fastdiv(idx - i03_offset - i02_offset, p.ne0_0mp, fastdiv_L(p.ne0_Ls, 2));
|
|
25
25
|
const uint i00 = idx - i03_offset - i02_offset - i01*p.ne00;
|
|
26
26
|
|
|
27
27
|
int param = floatBitsToInt(p.param1);
|