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
|
@@ -98,6 +98,46 @@
|
|
|
98
98
|
c_reg.lo += convert_float8(acc.lo); \
|
|
99
99
|
c_reg.hi += convert_float8(acc.hi); \
|
|
100
100
|
|
|
101
|
+
// Quarter-tile variant: computes 8 output columns (one skip-group) into a float8
|
|
102
|
+
// accumulator. Same reduction order / flush cadence as dotx16_reduce8, so the
|
|
103
|
+
// non-skipped path is byte-identical; it just lets the caller skip empty
|
|
104
|
+
// 8-column groups at finer granularity. Uses a private half8 `acc8`.
|
|
105
|
+
#define dotx8_reduce4(a_reg, b_lm, c_reg, lm_offset) \
|
|
106
|
+
acc8.s0 = dot(a_reg.s0123, b_lm[lm_offset + 0]); \
|
|
107
|
+
acc8.s1 = dot(a_reg.s0123, b_lm[lm_offset + 1]); \
|
|
108
|
+
acc8.s2 = dot(a_reg.s0123, b_lm[lm_offset + 2]); \
|
|
109
|
+
acc8.s3 = dot(a_reg.s0123, b_lm[lm_offset + 3]); \
|
|
110
|
+
acc8.s4 = dot(a_reg.s0123, b_lm[lm_offset + 4]); \
|
|
111
|
+
acc8.s5 = dot(a_reg.s0123, b_lm[lm_offset + 5]); \
|
|
112
|
+
acc8.s6 = dot(a_reg.s0123, b_lm[lm_offset + 6]); \
|
|
113
|
+
acc8.s7 = dot(a_reg.s0123, b_lm[lm_offset + 7]); \
|
|
114
|
+
acc8.s0 += dot(a_reg.s4567, b_lm[lm_offset + 32]); \
|
|
115
|
+
acc8.s1 += dot(a_reg.s4567, b_lm[lm_offset + 33]); \
|
|
116
|
+
acc8.s2 += dot(a_reg.s4567, b_lm[lm_offset + 34]); \
|
|
117
|
+
acc8.s3 += dot(a_reg.s4567, b_lm[lm_offset + 35]); \
|
|
118
|
+
acc8.s4 += dot(a_reg.s4567, b_lm[lm_offset + 36]); \
|
|
119
|
+
acc8.s5 += dot(a_reg.s4567, b_lm[lm_offset + 37]); \
|
|
120
|
+
acc8.s6 += dot(a_reg.s4567, b_lm[lm_offset + 38]); \
|
|
121
|
+
acc8.s7 += dot(a_reg.s4567, b_lm[lm_offset + 39]); \
|
|
122
|
+
c_reg += convert_float8(acc8); \
|
|
123
|
+
acc8.s0 = dot(a_reg.s89ab, b_lm[lm_offset + 64]); \
|
|
124
|
+
acc8.s1 = dot(a_reg.s89ab, b_lm[lm_offset + 65]); \
|
|
125
|
+
acc8.s2 = dot(a_reg.s89ab, b_lm[lm_offset + 66]); \
|
|
126
|
+
acc8.s3 = dot(a_reg.s89ab, b_lm[lm_offset + 67]); \
|
|
127
|
+
acc8.s4 = dot(a_reg.s89ab, b_lm[lm_offset + 68]); \
|
|
128
|
+
acc8.s5 = dot(a_reg.s89ab, b_lm[lm_offset + 69]); \
|
|
129
|
+
acc8.s6 = dot(a_reg.s89ab, b_lm[lm_offset + 70]); \
|
|
130
|
+
acc8.s7 = dot(a_reg.s89ab, b_lm[lm_offset + 71]); \
|
|
131
|
+
acc8.s0 += dot(a_reg.scdef, b_lm[lm_offset + 96]); \
|
|
132
|
+
acc8.s1 += dot(a_reg.scdef, b_lm[lm_offset + 97]); \
|
|
133
|
+
acc8.s2 += dot(a_reg.scdef, b_lm[lm_offset + 98]); \
|
|
134
|
+
acc8.s3 += dot(a_reg.scdef, b_lm[lm_offset + 99]); \
|
|
135
|
+
acc8.s4 += dot(a_reg.scdef, b_lm[lm_offset + 100]); \
|
|
136
|
+
acc8.s5 += dot(a_reg.scdef, b_lm[lm_offset + 101]); \
|
|
137
|
+
acc8.s6 += dot(a_reg.scdef, b_lm[lm_offset + 102]); \
|
|
138
|
+
acc8.s7 += dot(a_reg.scdef, b_lm[lm_offset + 103]); \
|
|
139
|
+
c_reg += convert_float8(acc8); \
|
|
140
|
+
|
|
101
141
|
|
|
102
142
|
__attribute__((qcom_wave_pair_mode(1))) // 1=force single 2=force pair
|
|
103
143
|
kernel void kernel_gemm_moe_q5_0_f32_ns(
|
|
@@ -110,7 +150,9 @@ kernel void kernel_gemm_moe_q5_0_f32_ns(
|
|
|
110
150
|
__write_only image1d_buffer_t dst,
|
|
111
151
|
__global int * total_tiles,
|
|
112
152
|
uint ne00,
|
|
113
|
-
uint ne01
|
|
153
|
+
uint ne01,
|
|
154
|
+
uint is_ragged,
|
|
155
|
+
uint skip_gran
|
|
114
156
|
) {
|
|
115
157
|
uint block_id_m = get_global_id(1); // m_tile
|
|
116
158
|
uint block_id_n = get_global_id(2); // n_tile
|
|
@@ -120,6 +162,28 @@ kernel void kernel_gemm_moe_q5_0_f32_ns(
|
|
|
120
162
|
return;
|
|
121
163
|
}
|
|
122
164
|
|
|
165
|
+
// Ragged tile-skip: when is_ragged and the upper 16 token-slots of this tile are all
|
|
166
|
+
// padding (router 0xFFFFFFFF), skip the second (reg_c.hi) dotx16_reduce8 half -> ~half
|
|
167
|
+
// the GEMM dot for sparse tiles. Numerically identical (the skipped lanes are padding).
|
|
168
|
+
// Ragged tile-skip: tokens are packed contiguously per expert (moe_scatter fills
|
|
169
|
+
// lanes 0..V-1, moe_fill pre-pads the rest), so router padding (0xFFFFFFFF) is always
|
|
170
|
+
// trailing. Find the valid-token count V and round it UP to the skip granularity
|
|
171
|
+
// skip_gran (columns per skip-group: 8 = quarter, 16 = half/legacy, 32 = disabled).
|
|
172
|
+
// A 8-column group g is all-padding iff its first column (8*g) >= n_active, so its
|
|
173
|
+
// dotx8_reduce4 is skipped. Numerically identical (skipped lanes are padding).
|
|
174
|
+
uint n_active = TILESIZE_N;
|
|
175
|
+
if (is_ragged && skip_gran < TILESIZE_N) {
|
|
176
|
+
uint n_valid = TILESIZE_N;
|
|
177
|
+
for (uint _t = 0; _t < TILESIZE_N; ++_t) {
|
|
178
|
+
if (src2[block_id_n * TILESIZE_N + _t] == 0xFFFFFFFFu) { n_valid = _t; break; }
|
|
179
|
+
}
|
|
180
|
+
n_active = min((uint)TILESIZE_N, ((n_valid + skip_gran - 1) / skip_gran) * skip_gran);
|
|
181
|
+
}
|
|
182
|
+
// Group 0 (cols 0-7) always runs; groups 1-3 skip when fully padding.
|
|
183
|
+
bool skip_g1 = (8u >= n_active);
|
|
184
|
+
bool skip_g2 = (16u >= n_active);
|
|
185
|
+
bool skip_g3 = (24u >= n_active);
|
|
186
|
+
|
|
123
187
|
__private half16 reg_a;
|
|
124
188
|
__private float32 reg_c = (float32)(0);
|
|
125
189
|
__local half4 shared_b[128];
|
|
@@ -171,9 +235,11 @@ kernel void kernel_gemm_moe_q5_0_f32_ns(
|
|
|
171
235
|
sub_group_barrier(CLK_LOCAL_MEM_FENCE);
|
|
172
236
|
|
|
173
237
|
// 32 16x16 fp16 dot product with 8 elements reduction for better precision
|
|
174
|
-
|
|
175
|
-
|
|
176
|
-
|
|
238
|
+
half8 acc8;
|
|
239
|
+
dotx8_reduce4(reg_a, shared_b, reg_c.lo.lo, 0);
|
|
240
|
+
if (!skip_g1) { dotx8_reduce4(reg_a, shared_b, reg_c.lo.hi, 8); }
|
|
241
|
+
if (!skip_g2) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.lo, 16); }
|
|
242
|
+
if (!skip_g3) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.hi, 24); }
|
|
177
243
|
|
|
178
244
|
// Repeat for second sub-block
|
|
179
245
|
uint half_step = step + TILESIZE_K;
|
|
@@ -198,8 +264,10 @@ kernel void kernel_gemm_moe_q5_0_f32_ns(
|
|
|
198
264
|
sub_group_barrier(CLK_LOCAL_MEM_FENCE);
|
|
199
265
|
|
|
200
266
|
// 32 16x16 fp16 dot product with 3-levels reduction for better precision
|
|
201
|
-
|
|
202
|
-
|
|
267
|
+
dotx8_reduce4(reg_a, shared_b, reg_c.lo.lo, 0);
|
|
268
|
+
if (!skip_g1) { dotx8_reduce4(reg_a, shared_b, reg_c.lo.hi, 8); }
|
|
269
|
+
if (!skip_g2) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.lo, 16); }
|
|
270
|
+
if (!skip_g3) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.hi, 24); }
|
|
203
271
|
}
|
|
204
272
|
|
|
205
273
|
if ((get_global_id(0) + block_id_m * TILESIZE_M) >= ne01) {
|
|
@@ -98,6 +98,46 @@
|
|
|
98
98
|
c_reg.lo += convert_float8(acc.lo); \
|
|
99
99
|
c_reg.hi += convert_float8(acc.hi); \
|
|
100
100
|
|
|
101
|
+
// Quarter-tile variant: computes 8 output columns (one skip-group) into a float8
|
|
102
|
+
// accumulator. Same reduction order / flush cadence as dotx16_reduce8, so the
|
|
103
|
+
// non-skipped path is byte-identical; it just lets the caller skip empty
|
|
104
|
+
// 8-column groups at finer granularity. Uses a private half8 `acc8`.
|
|
105
|
+
#define dotx8_reduce4(a_reg, b_lm, c_reg, lm_offset) \
|
|
106
|
+
acc8.s0 = dot(a_reg.s0123, b_lm[lm_offset + 0]); \
|
|
107
|
+
acc8.s1 = dot(a_reg.s0123, b_lm[lm_offset + 1]); \
|
|
108
|
+
acc8.s2 = dot(a_reg.s0123, b_lm[lm_offset + 2]); \
|
|
109
|
+
acc8.s3 = dot(a_reg.s0123, b_lm[lm_offset + 3]); \
|
|
110
|
+
acc8.s4 = dot(a_reg.s0123, b_lm[lm_offset + 4]); \
|
|
111
|
+
acc8.s5 = dot(a_reg.s0123, b_lm[lm_offset + 5]); \
|
|
112
|
+
acc8.s6 = dot(a_reg.s0123, b_lm[lm_offset + 6]); \
|
|
113
|
+
acc8.s7 = dot(a_reg.s0123, b_lm[lm_offset + 7]); \
|
|
114
|
+
acc8.s0 += dot(a_reg.s4567, b_lm[lm_offset + 32]); \
|
|
115
|
+
acc8.s1 += dot(a_reg.s4567, b_lm[lm_offset + 33]); \
|
|
116
|
+
acc8.s2 += dot(a_reg.s4567, b_lm[lm_offset + 34]); \
|
|
117
|
+
acc8.s3 += dot(a_reg.s4567, b_lm[lm_offset + 35]); \
|
|
118
|
+
acc8.s4 += dot(a_reg.s4567, b_lm[lm_offset + 36]); \
|
|
119
|
+
acc8.s5 += dot(a_reg.s4567, b_lm[lm_offset + 37]); \
|
|
120
|
+
acc8.s6 += dot(a_reg.s4567, b_lm[lm_offset + 38]); \
|
|
121
|
+
acc8.s7 += dot(a_reg.s4567, b_lm[lm_offset + 39]); \
|
|
122
|
+
c_reg += convert_float8(acc8); \
|
|
123
|
+
acc8.s0 = dot(a_reg.s89ab, b_lm[lm_offset + 64]); \
|
|
124
|
+
acc8.s1 = dot(a_reg.s89ab, b_lm[lm_offset + 65]); \
|
|
125
|
+
acc8.s2 = dot(a_reg.s89ab, b_lm[lm_offset + 66]); \
|
|
126
|
+
acc8.s3 = dot(a_reg.s89ab, b_lm[lm_offset + 67]); \
|
|
127
|
+
acc8.s4 = dot(a_reg.s89ab, b_lm[lm_offset + 68]); \
|
|
128
|
+
acc8.s5 = dot(a_reg.s89ab, b_lm[lm_offset + 69]); \
|
|
129
|
+
acc8.s6 = dot(a_reg.s89ab, b_lm[lm_offset + 70]); \
|
|
130
|
+
acc8.s7 = dot(a_reg.s89ab, b_lm[lm_offset + 71]); \
|
|
131
|
+
acc8.s0 += dot(a_reg.scdef, b_lm[lm_offset + 96]); \
|
|
132
|
+
acc8.s1 += dot(a_reg.scdef, b_lm[lm_offset + 97]); \
|
|
133
|
+
acc8.s2 += dot(a_reg.scdef, b_lm[lm_offset + 98]); \
|
|
134
|
+
acc8.s3 += dot(a_reg.scdef, b_lm[lm_offset + 99]); \
|
|
135
|
+
acc8.s4 += dot(a_reg.scdef, b_lm[lm_offset + 100]); \
|
|
136
|
+
acc8.s5 += dot(a_reg.scdef, b_lm[lm_offset + 101]); \
|
|
137
|
+
acc8.s6 += dot(a_reg.scdef, b_lm[lm_offset + 102]); \
|
|
138
|
+
acc8.s7 += dot(a_reg.scdef, b_lm[lm_offset + 103]); \
|
|
139
|
+
c_reg += convert_float8(acc8); \
|
|
140
|
+
|
|
101
141
|
|
|
102
142
|
__attribute__((qcom_wave_pair_mode(1))) // 1=force single 2=force pair
|
|
103
143
|
kernel void kernel_gemm_moe_q5_1_f32_ns(
|
|
@@ -111,7 +151,9 @@ kernel void kernel_gemm_moe_q5_1_f32_ns(
|
|
|
111
151
|
__write_only image1d_buffer_t dst,
|
|
112
152
|
__global int * total_tiles,
|
|
113
153
|
uint ne00,
|
|
114
|
-
uint ne01
|
|
154
|
+
uint ne01,
|
|
155
|
+
uint is_ragged,
|
|
156
|
+
uint skip_gran
|
|
115
157
|
) {
|
|
116
158
|
uint block_id_m = get_global_id(1); // m_tile
|
|
117
159
|
uint block_id_n = get_global_id(2); // n_tile
|
|
@@ -121,6 +163,28 @@ kernel void kernel_gemm_moe_q5_1_f32_ns(
|
|
|
121
163
|
return;
|
|
122
164
|
}
|
|
123
165
|
|
|
166
|
+
// Ragged tile-skip: when is_ragged and the upper 16 token-slots of this tile are all
|
|
167
|
+
// padding (router 0xFFFFFFFF), skip the second (reg_c.hi) dotx16_reduce8 half -> ~half
|
|
168
|
+
// the GEMM dot for sparse tiles. Numerically identical (the skipped lanes are padding).
|
|
169
|
+
// Ragged tile-skip: tokens are packed contiguously per expert (moe_scatter fills
|
|
170
|
+
// lanes 0..V-1, moe_fill pre-pads the rest), so router padding (0xFFFFFFFF) is always
|
|
171
|
+
// trailing. Find the valid-token count V and round it UP to the skip granularity
|
|
172
|
+
// skip_gran (columns per skip-group: 8 = quarter, 16 = half/legacy, 32 = disabled).
|
|
173
|
+
// A 8-column group g is all-padding iff its first column (8*g) >= n_active, so its
|
|
174
|
+
// dotx8_reduce4 is skipped. Numerically identical (skipped lanes are padding).
|
|
175
|
+
uint n_active = TILESIZE_N;
|
|
176
|
+
if (is_ragged && skip_gran < TILESIZE_N) {
|
|
177
|
+
uint n_valid = TILESIZE_N;
|
|
178
|
+
for (uint _t = 0; _t < TILESIZE_N; ++_t) {
|
|
179
|
+
if (src2[block_id_n * TILESIZE_N + _t] == 0xFFFFFFFFu) { n_valid = _t; break; }
|
|
180
|
+
}
|
|
181
|
+
n_active = min((uint)TILESIZE_N, ((n_valid + skip_gran - 1) / skip_gran) * skip_gran);
|
|
182
|
+
}
|
|
183
|
+
// Group 0 (cols 0-7) always runs; groups 1-3 skip when fully padding.
|
|
184
|
+
bool skip_g1 = (8u >= n_active);
|
|
185
|
+
bool skip_g2 = (16u >= n_active);
|
|
186
|
+
bool skip_g3 = (24u >= n_active);
|
|
187
|
+
|
|
124
188
|
__private half16 reg_a;
|
|
125
189
|
__private float32 reg_c = (float32)(0);
|
|
126
190
|
__local half4 shared_b[128];
|
|
@@ -173,9 +237,11 @@ kernel void kernel_gemm_moe_q5_1_f32_ns(
|
|
|
173
237
|
sub_group_barrier(CLK_LOCAL_MEM_FENCE);
|
|
174
238
|
|
|
175
239
|
// 32 16x16 fp16 dot product with 8 elements reduction for better precision
|
|
176
|
-
|
|
177
|
-
|
|
178
|
-
|
|
240
|
+
half8 acc8;
|
|
241
|
+
dotx8_reduce4(reg_a, shared_b, reg_c.lo.lo, 0);
|
|
242
|
+
if (!skip_g1) { dotx8_reduce4(reg_a, shared_b, reg_c.lo.hi, 8); }
|
|
243
|
+
if (!skip_g2) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.lo, 16); }
|
|
244
|
+
if (!skip_g3) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.hi, 24); }
|
|
179
245
|
|
|
180
246
|
// Repeat for second sub-block
|
|
181
247
|
uint half_step = step + TILESIZE_K;
|
|
@@ -200,8 +266,10 @@ kernel void kernel_gemm_moe_q5_1_f32_ns(
|
|
|
200
266
|
sub_group_barrier(CLK_LOCAL_MEM_FENCE);
|
|
201
267
|
|
|
202
268
|
// 32 16x16 fp16 dot product with 3-levels reduction for better precision
|
|
203
|
-
|
|
204
|
-
|
|
269
|
+
dotx8_reduce4(reg_a, shared_b, reg_c.lo.lo, 0);
|
|
270
|
+
if (!skip_g1) { dotx8_reduce4(reg_a, shared_b, reg_c.lo.hi, 8); }
|
|
271
|
+
if (!skip_g2) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.lo, 16); }
|
|
272
|
+
if (!skip_g3) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.hi, 24); }
|
|
205
273
|
}
|
|
206
274
|
|
|
207
275
|
if ((get_global_id(0) + block_id_m * TILESIZE_M) >= ne01) {
|
|
@@ -114,6 +114,46 @@ inline void get_scale_min_k4(
|
|
|
114
114
|
c_reg.lo += convert_float8(acc.lo); \
|
|
115
115
|
c_reg.hi += convert_float8(acc.hi); \
|
|
116
116
|
|
|
117
|
+
// Quarter-tile variant: computes 8 output columns (one skip-group) into a float8
|
|
118
|
+
// accumulator. Same reduction order / flush cadence as dotx16_reduce8, so the
|
|
119
|
+
// non-skipped path is byte-identical; it just lets the caller skip empty
|
|
120
|
+
// 8-column groups at finer granularity. Uses a private half8 `acc8`.
|
|
121
|
+
#define dotx8_reduce4(a_reg, b_lm, c_reg, lm_offset) \
|
|
122
|
+
acc8.s0 = dot(a_reg.s0123, b_lm[lm_offset + 0]); \
|
|
123
|
+
acc8.s1 = dot(a_reg.s0123, b_lm[lm_offset + 1]); \
|
|
124
|
+
acc8.s2 = dot(a_reg.s0123, b_lm[lm_offset + 2]); \
|
|
125
|
+
acc8.s3 = dot(a_reg.s0123, b_lm[lm_offset + 3]); \
|
|
126
|
+
acc8.s4 = dot(a_reg.s0123, b_lm[lm_offset + 4]); \
|
|
127
|
+
acc8.s5 = dot(a_reg.s0123, b_lm[lm_offset + 5]); \
|
|
128
|
+
acc8.s6 = dot(a_reg.s0123, b_lm[lm_offset + 6]); \
|
|
129
|
+
acc8.s7 = dot(a_reg.s0123, b_lm[lm_offset + 7]); \
|
|
130
|
+
acc8.s0 += dot(a_reg.s4567, b_lm[lm_offset + 32]); \
|
|
131
|
+
acc8.s1 += dot(a_reg.s4567, b_lm[lm_offset + 33]); \
|
|
132
|
+
acc8.s2 += dot(a_reg.s4567, b_lm[lm_offset + 34]); \
|
|
133
|
+
acc8.s3 += dot(a_reg.s4567, b_lm[lm_offset + 35]); \
|
|
134
|
+
acc8.s4 += dot(a_reg.s4567, b_lm[lm_offset + 36]); \
|
|
135
|
+
acc8.s5 += dot(a_reg.s4567, b_lm[lm_offset + 37]); \
|
|
136
|
+
acc8.s6 += dot(a_reg.s4567, b_lm[lm_offset + 38]); \
|
|
137
|
+
acc8.s7 += dot(a_reg.s4567, b_lm[lm_offset + 39]); \
|
|
138
|
+
c_reg += convert_float8(acc8); \
|
|
139
|
+
acc8.s0 = dot(a_reg.s89ab, b_lm[lm_offset + 64]); \
|
|
140
|
+
acc8.s1 = dot(a_reg.s89ab, b_lm[lm_offset + 65]); \
|
|
141
|
+
acc8.s2 = dot(a_reg.s89ab, b_lm[lm_offset + 66]); \
|
|
142
|
+
acc8.s3 = dot(a_reg.s89ab, b_lm[lm_offset + 67]); \
|
|
143
|
+
acc8.s4 = dot(a_reg.s89ab, b_lm[lm_offset + 68]); \
|
|
144
|
+
acc8.s5 = dot(a_reg.s89ab, b_lm[lm_offset + 69]); \
|
|
145
|
+
acc8.s6 = dot(a_reg.s89ab, b_lm[lm_offset + 70]); \
|
|
146
|
+
acc8.s7 = dot(a_reg.s89ab, b_lm[lm_offset + 71]); \
|
|
147
|
+
acc8.s0 += dot(a_reg.scdef, b_lm[lm_offset + 96]); \
|
|
148
|
+
acc8.s1 += dot(a_reg.scdef, b_lm[lm_offset + 97]); \
|
|
149
|
+
acc8.s2 += dot(a_reg.scdef, b_lm[lm_offset + 98]); \
|
|
150
|
+
acc8.s3 += dot(a_reg.scdef, b_lm[lm_offset + 99]); \
|
|
151
|
+
acc8.s4 += dot(a_reg.scdef, b_lm[lm_offset + 100]); \
|
|
152
|
+
acc8.s5 += dot(a_reg.scdef, b_lm[lm_offset + 101]); \
|
|
153
|
+
acc8.s6 += dot(a_reg.scdef, b_lm[lm_offset + 102]); \
|
|
154
|
+
acc8.s7 += dot(a_reg.scdef, b_lm[lm_offset + 103]); \
|
|
155
|
+
c_reg += convert_float8(acc8); \
|
|
156
|
+
|
|
117
157
|
|
|
118
158
|
__attribute__((qcom_wave_pair_mode(1)))
|
|
119
159
|
kernel void kernel_gemm_moe_q5_k_f32_ns(
|
|
@@ -128,7 +168,9 @@ kernel void kernel_gemm_moe_q5_k_f32_ns(
|
|
|
128
168
|
__write_only image1d_buffer_t dst,
|
|
129
169
|
__global int * total_tiles,
|
|
130
170
|
uint ne00,
|
|
131
|
-
uint ne01
|
|
171
|
+
uint ne01,
|
|
172
|
+
uint is_ragged,
|
|
173
|
+
uint skip_gran
|
|
132
174
|
) {
|
|
133
175
|
uint block_id_m = get_global_id(1); // m_tile
|
|
134
176
|
uint block_id_n = get_global_id(2); // n_tile
|
|
@@ -138,6 +180,28 @@ kernel void kernel_gemm_moe_q5_k_f32_ns(
|
|
|
138
180
|
return;
|
|
139
181
|
}
|
|
140
182
|
|
|
183
|
+
// Ragged tile-skip: when is_ragged and the upper 16 token-slots of this tile are all
|
|
184
|
+
// padding (router 0xFFFFFFFF), skip the second (reg_c.hi) dotx16_reduce8 half -> ~half
|
|
185
|
+
// the GEMM dot for sparse tiles. Numerically identical (the skipped lanes are padding).
|
|
186
|
+
// Ragged tile-skip: tokens are packed contiguously per expert (moe_scatter fills
|
|
187
|
+
// lanes 0..V-1, moe_fill pre-pads the rest), so router padding (0xFFFFFFFF) is always
|
|
188
|
+
// trailing. Find the valid-token count V and round it UP to the skip granularity
|
|
189
|
+
// skip_gran (columns per skip-group: 8 = quarter, 16 = half/legacy, 32 = disabled).
|
|
190
|
+
// A 8-column group g is all-padding iff its first column (8*g) >= n_active, so its
|
|
191
|
+
// dotx8_reduce4 is skipped. Numerically identical (skipped lanes are padding).
|
|
192
|
+
uint n_active = TILESIZE_N;
|
|
193
|
+
if (is_ragged && skip_gran < TILESIZE_N) {
|
|
194
|
+
uint n_valid = TILESIZE_N;
|
|
195
|
+
for (uint _t = 0; _t < TILESIZE_N; ++_t) {
|
|
196
|
+
if (src2[block_id_n * TILESIZE_N + _t] == 0xFFFFFFFFu) { n_valid = _t; break; }
|
|
197
|
+
}
|
|
198
|
+
n_active = min((uint)TILESIZE_N, ((n_valid + skip_gran - 1) / skip_gran) * skip_gran);
|
|
199
|
+
}
|
|
200
|
+
// Group 0 (cols 0-7) always runs; groups 1-3 skip when fully padding.
|
|
201
|
+
bool skip_g1 = (8u >= n_active);
|
|
202
|
+
bool skip_g2 = (16u >= n_active);
|
|
203
|
+
bool skip_g3 = (24u >= n_active);
|
|
204
|
+
|
|
141
205
|
__private half16 reg_a;
|
|
142
206
|
__private float32 reg_c = (float32)(0);
|
|
143
207
|
__local half4 shared_b[128];
|
|
@@ -204,9 +268,11 @@ kernel void kernel_gemm_moe_q5_k_f32_ns(
|
|
|
204
268
|
|
|
205
269
|
sub_group_barrier(CLK_LOCAL_MEM_FENCE);
|
|
206
270
|
|
|
207
|
-
|
|
208
|
-
|
|
209
|
-
|
|
271
|
+
half8 acc8;
|
|
272
|
+
dotx8_reduce4(reg_a, shared_b, reg_c.lo.lo, 0);
|
|
273
|
+
if (!skip_g1) { dotx8_reduce4(reg_a, shared_b, reg_c.lo.hi, 8); }
|
|
274
|
+
if (!skip_g2) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.lo, 16); }
|
|
275
|
+
if (!skip_g3) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.hi, 24); }
|
|
210
276
|
|
|
211
277
|
// Second half
|
|
212
278
|
uint half_step = step + TILESIZE_K;
|
|
@@ -226,8 +292,10 @@ kernel void kernel_gemm_moe_q5_k_f32_ns(
|
|
|
226
292
|
|
|
227
293
|
sub_group_barrier(CLK_LOCAL_MEM_FENCE);
|
|
228
294
|
|
|
229
|
-
|
|
230
|
-
|
|
295
|
+
dotx8_reduce4(reg_a, shared_b, reg_c.lo.lo, 0);
|
|
296
|
+
if (!skip_g1) { dotx8_reduce4(reg_a, shared_b, reg_c.lo.hi, 8); }
|
|
297
|
+
if (!skip_g2) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.lo, 16); }
|
|
298
|
+
if (!skip_g3) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.hi, 24); }
|
|
231
299
|
}
|
|
232
300
|
|
|
233
301
|
if ((get_global_id(0) + block_id_m * TILESIZE_M) >= ne01) {
|
|
@@ -98,6 +98,46 @@
|
|
|
98
98
|
c_reg.lo += convert_float8(acc.lo); \
|
|
99
99
|
c_reg.hi += convert_float8(acc.hi); \
|
|
100
100
|
|
|
101
|
+
// Quarter-tile variant: computes 8 output columns (one skip-group) into a float8
|
|
102
|
+
// accumulator. Same reduction order / flush cadence as dotx16_reduce8, so the
|
|
103
|
+
// non-skipped path is byte-identical; it just lets the caller skip empty
|
|
104
|
+
// 8-column groups at finer granularity. Uses a private half8 `acc8`.
|
|
105
|
+
#define dotx8_reduce4(a_reg, b_lm, c_reg, lm_offset) \
|
|
106
|
+
acc8.s0 = dot(a_reg.s0123, b_lm[lm_offset + 0]); \
|
|
107
|
+
acc8.s1 = dot(a_reg.s0123, b_lm[lm_offset + 1]); \
|
|
108
|
+
acc8.s2 = dot(a_reg.s0123, b_lm[lm_offset + 2]); \
|
|
109
|
+
acc8.s3 = dot(a_reg.s0123, b_lm[lm_offset + 3]); \
|
|
110
|
+
acc8.s4 = dot(a_reg.s0123, b_lm[lm_offset + 4]); \
|
|
111
|
+
acc8.s5 = dot(a_reg.s0123, b_lm[lm_offset + 5]); \
|
|
112
|
+
acc8.s6 = dot(a_reg.s0123, b_lm[lm_offset + 6]); \
|
|
113
|
+
acc8.s7 = dot(a_reg.s0123, b_lm[lm_offset + 7]); \
|
|
114
|
+
acc8.s0 += dot(a_reg.s4567, b_lm[lm_offset + 32]); \
|
|
115
|
+
acc8.s1 += dot(a_reg.s4567, b_lm[lm_offset + 33]); \
|
|
116
|
+
acc8.s2 += dot(a_reg.s4567, b_lm[lm_offset + 34]); \
|
|
117
|
+
acc8.s3 += dot(a_reg.s4567, b_lm[lm_offset + 35]); \
|
|
118
|
+
acc8.s4 += dot(a_reg.s4567, b_lm[lm_offset + 36]); \
|
|
119
|
+
acc8.s5 += dot(a_reg.s4567, b_lm[lm_offset + 37]); \
|
|
120
|
+
acc8.s6 += dot(a_reg.s4567, b_lm[lm_offset + 38]); \
|
|
121
|
+
acc8.s7 += dot(a_reg.s4567, b_lm[lm_offset + 39]); \
|
|
122
|
+
c_reg += convert_float8(acc8); \
|
|
123
|
+
acc8.s0 = dot(a_reg.s89ab, b_lm[lm_offset + 64]); \
|
|
124
|
+
acc8.s1 = dot(a_reg.s89ab, b_lm[lm_offset + 65]); \
|
|
125
|
+
acc8.s2 = dot(a_reg.s89ab, b_lm[lm_offset + 66]); \
|
|
126
|
+
acc8.s3 = dot(a_reg.s89ab, b_lm[lm_offset + 67]); \
|
|
127
|
+
acc8.s4 = dot(a_reg.s89ab, b_lm[lm_offset + 68]); \
|
|
128
|
+
acc8.s5 = dot(a_reg.s89ab, b_lm[lm_offset + 69]); \
|
|
129
|
+
acc8.s6 = dot(a_reg.s89ab, b_lm[lm_offset + 70]); \
|
|
130
|
+
acc8.s7 = dot(a_reg.s89ab, b_lm[lm_offset + 71]); \
|
|
131
|
+
acc8.s0 += dot(a_reg.scdef, b_lm[lm_offset + 96]); \
|
|
132
|
+
acc8.s1 += dot(a_reg.scdef, b_lm[lm_offset + 97]); \
|
|
133
|
+
acc8.s2 += dot(a_reg.scdef, b_lm[lm_offset + 98]); \
|
|
134
|
+
acc8.s3 += dot(a_reg.scdef, b_lm[lm_offset + 99]); \
|
|
135
|
+
acc8.s4 += dot(a_reg.scdef, b_lm[lm_offset + 100]); \
|
|
136
|
+
acc8.s5 += dot(a_reg.scdef, b_lm[lm_offset + 101]); \
|
|
137
|
+
acc8.s6 += dot(a_reg.scdef, b_lm[lm_offset + 102]); \
|
|
138
|
+
acc8.s7 += dot(a_reg.scdef, b_lm[lm_offset + 103]); \
|
|
139
|
+
c_reg += convert_float8(acc8); \
|
|
140
|
+
|
|
101
141
|
|
|
102
142
|
__attribute__((qcom_wave_pair_mode(1)))
|
|
103
143
|
kernel void kernel_gemm_moe_q6_k_f32_ns(
|
|
@@ -111,7 +151,9 @@ kernel void kernel_gemm_moe_q6_k_f32_ns(
|
|
|
111
151
|
__write_only image1d_buffer_t dst,
|
|
112
152
|
__global int * total_tiles,
|
|
113
153
|
uint ne00,
|
|
114
|
-
uint ne01
|
|
154
|
+
uint ne01,
|
|
155
|
+
uint is_ragged,
|
|
156
|
+
uint skip_gran
|
|
115
157
|
) {
|
|
116
158
|
uint block_id_m = get_global_id(1); // m_tile
|
|
117
159
|
uint block_id_n = get_global_id(2); // n_tile
|
|
@@ -121,6 +163,28 @@ kernel void kernel_gemm_moe_q6_k_f32_ns(
|
|
|
121
163
|
return;
|
|
122
164
|
}
|
|
123
165
|
|
|
166
|
+
// Ragged tile-skip: when is_ragged and the upper 16 token-slots of this tile are all
|
|
167
|
+
// padding (router 0xFFFFFFFF), skip the second (reg_c.hi) dotx16_reduce8 half -> ~half
|
|
168
|
+
// the GEMM dot for sparse tiles. Numerically identical (the skipped lanes are padding).
|
|
169
|
+
// Ragged tile-skip: tokens are packed contiguously per expert (moe_scatter fills
|
|
170
|
+
// lanes 0..V-1, moe_fill pre-pads the rest), so router padding (0xFFFFFFFF) is always
|
|
171
|
+
// trailing. Find the valid-token count V and round it UP to the skip granularity
|
|
172
|
+
// skip_gran (columns per skip-group: 8 = quarter, 16 = half/legacy, 32 = disabled).
|
|
173
|
+
// A 8-column group g is all-padding iff its first column (8*g) >= n_active, so its
|
|
174
|
+
// dotx8_reduce4 is skipped. Numerically identical (skipped lanes are padding).
|
|
175
|
+
uint n_active = TILESIZE_N;
|
|
176
|
+
if (is_ragged && skip_gran < TILESIZE_N) {
|
|
177
|
+
uint n_valid = TILESIZE_N;
|
|
178
|
+
for (uint _t = 0; _t < TILESIZE_N; ++_t) {
|
|
179
|
+
if (src2[block_id_n * TILESIZE_N + _t] == 0xFFFFFFFFu) { n_valid = _t; break; }
|
|
180
|
+
}
|
|
181
|
+
n_active = min((uint)TILESIZE_N, ((n_valid + skip_gran - 1) / skip_gran) * skip_gran);
|
|
182
|
+
}
|
|
183
|
+
// Group 0 (cols 0-7) always runs; groups 1-3 skip when fully padding.
|
|
184
|
+
bool skip_g1 = (8u >= n_active);
|
|
185
|
+
bool skip_g2 = (16u >= n_active);
|
|
186
|
+
bool skip_g3 = (24u >= n_active);
|
|
187
|
+
|
|
124
188
|
__private half16 reg_a;
|
|
125
189
|
__private float32 reg_c = (float32)(0);
|
|
126
190
|
__local half4 shared_b[128];
|
|
@@ -183,9 +247,11 @@ kernel void kernel_gemm_moe_q6_k_f32_ns(
|
|
|
183
247
|
|
|
184
248
|
sub_group_barrier(CLK_LOCAL_MEM_FENCE);
|
|
185
249
|
|
|
186
|
-
|
|
187
|
-
|
|
188
|
-
|
|
250
|
+
half8 acc8;
|
|
251
|
+
dotx8_reduce4(reg_a, shared_b, reg_c.lo.lo, 0);
|
|
252
|
+
if (!skip_g1) { dotx8_reduce4(reg_a, shared_b, reg_c.lo.hi, 8); }
|
|
253
|
+
if (!skip_g2) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.lo, 16); }
|
|
254
|
+
if (!skip_g3) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.hi, 24); }
|
|
189
255
|
|
|
190
256
|
// Second half
|
|
191
257
|
uint half_step = step + TILESIZE_K;
|
|
@@ -205,8 +271,10 @@ kernel void kernel_gemm_moe_q6_k_f32_ns(
|
|
|
205
271
|
|
|
206
272
|
sub_group_barrier(CLK_LOCAL_MEM_FENCE);
|
|
207
273
|
|
|
208
|
-
|
|
209
|
-
|
|
274
|
+
dotx8_reduce4(reg_a, shared_b, reg_c.lo.lo, 0);
|
|
275
|
+
if (!skip_g1) { dotx8_reduce4(reg_a, shared_b, reg_c.lo.hi, 8); }
|
|
276
|
+
if (!skip_g2) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.lo, 16); }
|
|
277
|
+
if (!skip_g3) { dotx8_reduce4(reg_a, shared_b, reg_c.hi.hi, 24); }
|
|
210
278
|
}
|
|
211
279
|
|
|
212
280
|
if ((get_global_id(0) + block_id_m * TILESIZE_M) >= ne01) {
|
|
@@ -0,0 +1,94 @@
|
|
|
1
|
+
#pragma OPENCL EXTENSION cl_khr_fp16 : enable
|
|
2
|
+
#pragma OPENCL EXTENSION cl_qcom_reqd_sub_group_size : enable
|
|
3
|
+
|
|
4
|
+
#ifdef cl_qcom_reqd_sub_group_size
|
|
5
|
+
#pragma OPENCL EXTENSION cl_qcom_reqd_sub_group_size : enable
|
|
6
|
+
#define ADRENO_GPU 1
|
|
7
|
+
#define REQD_SUBGROUP_SIZE_128 __attribute__((qcom_reqd_sub_group_size("full")))
|
|
8
|
+
#endif
|
|
9
|
+
|
|
10
|
+
// each work-item computes a 4 (rows of A / m) x 8 (cols of B / n) output tile.
|
|
11
|
+
#ifdef ADRENO_GPU
|
|
12
|
+
REQD_SUBGROUP_SIZE_128
|
|
13
|
+
#endif
|
|
14
|
+
kernel void kernel_gemm_noshuffle_q1_0_f32(
|
|
15
|
+
global const uint * src0_q,
|
|
16
|
+
global const half * src0_d,
|
|
17
|
+
read_only image1d_buffer_t src1,
|
|
18
|
+
global float * dst,
|
|
19
|
+
int k,
|
|
20
|
+
int m,
|
|
21
|
+
int n,
|
|
22
|
+
int n_no_padding,
|
|
23
|
+
ulong offsetd
|
|
24
|
+
) {
|
|
25
|
+
int n_4 = n >> 2;
|
|
26
|
+
|
|
27
|
+
int gy = get_global_id(0);
|
|
28
|
+
int gx = get_global_id(1);
|
|
29
|
+
int gx_2 = gx << 2;
|
|
30
|
+
dst = (global float *)((global char*)dst + offsetd);
|
|
31
|
+
|
|
32
|
+
half8 c0 = 0, c1 = 0, c2 = 0, c3 = 0;
|
|
33
|
+
half8 B;
|
|
34
|
+
|
|
35
|
+
global const uint* wptr = src0_q + gx_2;
|
|
36
|
+
global const half* sptr = src0_d + gx_2;
|
|
37
|
+
|
|
38
|
+
// 32 weights per uint32, 128 weights (one block / one scale) per 4 uint32.
|
|
39
|
+
for (int i = 0; i < k; i += 32) {
|
|
40
|
+
uint4 pack4 = vload4(0, wptr + (i / 32) * m); // 4 rows, 32 K-values each
|
|
41
|
+
half4 scale = vload4(0, sptr + (i / 128) * m); // 4 rows, one scale per 128
|
|
42
|
+
|
|
43
|
+
for (int j = 0; j < 32; ++j) {
|
|
44
|
+
B.s0123 = read_imageh(src1, gy * 2 + (i + j) * n_4);
|
|
45
|
+
B.s4567 = read_imageh(src1, gy * 2 + (i + j) * n_4 + 1);
|
|
46
|
+
|
|
47
|
+
// sign bit -> +-1 (half arithmetic avoids unsigned underflow)
|
|
48
|
+
half4 wj = (half4)(
|
|
49
|
+
2.0h * (half)((pack4.s0 >> j) & 1u) - 1.0h,
|
|
50
|
+
2.0h * (half)((pack4.s1 >> j) & 1u) - 1.0h,
|
|
51
|
+
2.0h * (half)((pack4.s2 >> j) & 1u) - 1.0h,
|
|
52
|
+
2.0h * (half)((pack4.s3 >> j) & 1u) - 1.0h) * scale;
|
|
53
|
+
|
|
54
|
+
c0 += B * wj.s0;
|
|
55
|
+
c1 += B * wj.s1;
|
|
56
|
+
c2 += B * wj.s2;
|
|
57
|
+
c3 += B * wj.s3;
|
|
58
|
+
}
|
|
59
|
+
}
|
|
60
|
+
|
|
61
|
+
int idx = (gy << 3) * m + (gx << 2);
|
|
62
|
+
|
|
63
|
+
if(idx+3 < m*n_no_padding){
|
|
64
|
+
vstore4((float4)(c0.s0, c1.s0, c2.s0, c3.s0), 0, dst + idx);
|
|
65
|
+
idx += m;
|
|
66
|
+
}
|
|
67
|
+
if(idx+3 < m*n_no_padding){
|
|
68
|
+
vstore4((float4)(c0.s1, c1.s1, c2.s1, c3.s1), 0, dst + idx);
|
|
69
|
+
idx += m;
|
|
70
|
+
}
|
|
71
|
+
if(idx+3 < m*n_no_padding){
|
|
72
|
+
vstore4((float4)(c0.s2, c1.s2, c2.s2, c3.s2), 0, dst + idx);
|
|
73
|
+
idx += m;
|
|
74
|
+
}
|
|
75
|
+
if(idx+3 < m*n_no_padding){
|
|
76
|
+
vstore4((float4)(c0.s3, c1.s3, c2.s3, c3.s3), 0, dst + idx);
|
|
77
|
+
idx += m;
|
|
78
|
+
}
|
|
79
|
+
if(idx+3 < m*n_no_padding){
|
|
80
|
+
vstore4((float4)(c0.s4, c1.s4, c2.s4, c3.s4), 0, dst + idx);
|
|
81
|
+
idx += m;
|
|
82
|
+
}
|
|
83
|
+
if(idx+3 < m*n_no_padding){
|
|
84
|
+
vstore4((float4)(c0.s5, c1.s5, c2.s5, c3.s5), 0, dst + idx);
|
|
85
|
+
idx += m;
|
|
86
|
+
}
|
|
87
|
+
if(idx+3 < m*n_no_padding){
|
|
88
|
+
vstore4((float4)(c0.s6, c1.s6, c2.s6, c3.s6), 0, dst + idx);
|
|
89
|
+
idx += m;
|
|
90
|
+
}
|
|
91
|
+
if(idx+3 < m*n_no_padding){
|
|
92
|
+
vstore4((float4)(c0.s7, c1.s7, c2.s7, c3.s7), 0, dst + idx);
|
|
93
|
+
}
|
|
94
|
+
}
|
|
@@ -0,0 +1,121 @@
|
|
|
1
|
+
#pragma OPENCL EXTENSION cl_khr_fp16 : enable
|
|
2
|
+
#pragma OPENCL EXTENSION cl_khr_subgroups : enable
|
|
3
|
+
|
|
4
|
+
#ifdef cl_qcom_reqd_sub_group_size
|
|
5
|
+
#pragma OPENCL EXTENSION cl_qcom_reqd_sub_group_size : enable
|
|
6
|
+
#define ADRENO_GPU 1
|
|
7
|
+
#define REQD_SUBGROUP_SIZE_64 __attribute__((qcom_reqd_sub_group_size("half")))
|
|
8
|
+
#endif
|
|
9
|
+
|
|
10
|
+
#define QK1_0 128
|
|
11
|
+
#define N_SIMDGROUP 4
|
|
12
|
+
|
|
13
|
+
#define dequantizeBlockAccum_q1(total, bits, scale, regB, lb) \
|
|
14
|
+
total += (2.0f*(float)((bits >> 0) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s0, lb+0); \
|
|
15
|
+
total += (2.0f*(float)((bits >> 1) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s1, lb+0); \
|
|
16
|
+
total += (2.0f*(float)((bits >> 2) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s2, lb+0); \
|
|
17
|
+
total += (2.0f*(float)((bits >> 3) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s3, lb+0); \
|
|
18
|
+
total += (2.0f*(float)((bits >> 4) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s4, lb+0); \
|
|
19
|
+
total += (2.0f*(float)((bits >> 5) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s5, lb+0); \
|
|
20
|
+
total += (2.0f*(float)((bits >> 6) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s6, lb+0); \
|
|
21
|
+
total += (2.0f*(float)((bits >> 7) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s7, lb+0); \
|
|
22
|
+
total += (2.0f*(float)((bits >> 8) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s0, lb+1); \
|
|
23
|
+
total += (2.0f*(float)((bits >> 9) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s1, lb+1); \
|
|
24
|
+
total += (2.0f*(float)((bits >> 10) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s2, lb+1); \
|
|
25
|
+
total += (2.0f*(float)((bits >> 11) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s3, lb+1); \
|
|
26
|
+
total += (2.0f*(float)((bits >> 12) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s4, lb+1); \
|
|
27
|
+
total += (2.0f*(float)((bits >> 13) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s5, lb+1); \
|
|
28
|
+
total += (2.0f*(float)((bits >> 14) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s6, lb+1); \
|
|
29
|
+
total += (2.0f*(float)((bits >> 15) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s7, lb+1); \
|
|
30
|
+
total += (2.0f*(float)((bits >> 16) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s0, lb+2); \
|
|
31
|
+
total += (2.0f*(float)((bits >> 17) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s1, lb+2); \
|
|
32
|
+
total += (2.0f*(float)((bits >> 18) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s2, lb+2); \
|
|
33
|
+
total += (2.0f*(float)((bits >> 19) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s3, lb+2); \
|
|
34
|
+
total += (2.0f*(float)((bits >> 20) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s4, lb+2); \
|
|
35
|
+
total += (2.0f*(float)((bits >> 21) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s5, lb+2); \
|
|
36
|
+
total += (2.0f*(float)((bits >> 22) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s6, lb+2); \
|
|
37
|
+
total += (2.0f*(float)((bits >> 23) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s7, lb+2); \
|
|
38
|
+
total += (2.0f*(float)((bits >> 24) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s0, lb+3); \
|
|
39
|
+
total += (2.0f*(float)((bits >> 25) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s1, lb+3); \
|
|
40
|
+
total += (2.0f*(float)((bits >> 26) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s2, lb+3); \
|
|
41
|
+
total += (2.0f*(float)((bits >> 27) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s3, lb+3); \
|
|
42
|
+
total += (2.0f*(float)((bits >> 28) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s4, lb+3); \
|
|
43
|
+
total += (2.0f*(float)((bits >> 29) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s5, lb+3); \
|
|
44
|
+
total += (2.0f*(float)((bits >> 30) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s6, lb+3); \
|
|
45
|
+
total += (2.0f*(float)((bits >> 31) & 1u) - 1.0f) * scale * sub_group_broadcast(regB.s7, lb+3);
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
#ifdef ADRENO_GPU
|
|
49
|
+
REQD_SUBGROUP_SIZE_64
|
|
50
|
+
#endif
|
|
51
|
+
__kernel void kernel_gemv_noshuffle_q1_0_f32(
|
|
52
|
+
read_only image1d_buffer_t src0_q,
|
|
53
|
+
global half * src0_d,
|
|
54
|
+
read_only image1d_buffer_t src1,
|
|
55
|
+
ulong offset1,
|
|
56
|
+
global float * dst,
|
|
57
|
+
ulong offsetd,
|
|
58
|
+
int ne00,
|
|
59
|
+
int ne01,
|
|
60
|
+
int ne02,
|
|
61
|
+
int ne10,
|
|
62
|
+
int ne12,
|
|
63
|
+
int ne0,
|
|
64
|
+
int ne1,
|
|
65
|
+
int r2,
|
|
66
|
+
int r3)
|
|
67
|
+
{
|
|
68
|
+
uint groupId = get_local_id(1);
|
|
69
|
+
uint gid = get_global_id(0);
|
|
70
|
+
ushort slid = get_sub_group_local_id();
|
|
71
|
+
|
|
72
|
+
uint K = ne00;
|
|
73
|
+
uint M = ne01;
|
|
74
|
+
|
|
75
|
+
uint LINE_STRIDE_A = M;
|
|
76
|
+
uint BLOCK_STRIDE_A = 4 * M;
|
|
77
|
+
|
|
78
|
+
uint4 regA;
|
|
79
|
+
half regS;
|
|
80
|
+
float8 regB;
|
|
81
|
+
|
|
82
|
+
float totalSum = 0.0f;
|
|
83
|
+
|
|
84
|
+
#pragma unroll 1
|
|
85
|
+
for (uint kb = groupId; kb < (K / QK1_0); kb += N_SIMDGROUP) {
|
|
86
|
+
regS = src0_d[gid + kb * LINE_STRIDE_A]; // each fiber loads its row's scale
|
|
87
|
+
|
|
88
|
+
// first 16 fibers load 8 B values each -> 128 activations for this block
|
|
89
|
+
if (slid < 16) {
|
|
90
|
+
regB.s0123 = read_imagef(src1, (slid * 2 + kb * 32));
|
|
91
|
+
regB.s4567 = read_imagef(src1, (1 + slid * 2 + kb * 32));
|
|
92
|
+
}
|
|
93
|
+
|
|
94
|
+
// load this row's 4 uint32 (128 sign bits)
|
|
95
|
+
regA.s0 = read_imageui(src0_q, (gid + kb * BLOCK_STRIDE_A + LINE_STRIDE_A * 0)).x;
|
|
96
|
+
regA.s1 = read_imageui(src0_q, (gid + kb * BLOCK_STRIDE_A + LINE_STRIDE_A * 1)).x;
|
|
97
|
+
regA.s2 = read_imageui(src0_q, (gid + kb * BLOCK_STRIDE_A + LINE_STRIDE_A * 2)).x;
|
|
98
|
+
regA.s3 = read_imageui(src0_q, (gid + kb * BLOCK_STRIDE_A + LINE_STRIDE_A * 3)).x;
|
|
99
|
+
|
|
100
|
+
float scale = (float)regS;
|
|
101
|
+
dequantizeBlockAccum_q1(totalSum, regA.s0, scale, regB, 0);
|
|
102
|
+
dequantizeBlockAccum_q1(totalSum, regA.s1, scale, regB, 4);
|
|
103
|
+
dequantizeBlockAccum_q1(totalSum, regA.s2, scale, regB, 8);
|
|
104
|
+
dequantizeBlockAccum_q1(totalSum, regA.s3, scale, regB, 12);
|
|
105
|
+
}
|
|
106
|
+
|
|
107
|
+
// reduction in local memory, assumes #wave = N_SIMDGROUP = 4
|
|
108
|
+
local float reduceLM[SIMDGROUP_WIDTH * 3];
|
|
109
|
+
if (groupId == 1) reduceLM[SIMDGROUP_WIDTH * 0 + slid] = totalSum;
|
|
110
|
+
if (groupId == 2) reduceLM[SIMDGROUP_WIDTH * 1 + slid] = totalSum;
|
|
111
|
+
if (groupId == 3) reduceLM[SIMDGROUP_WIDTH * 2 + slid] = totalSum;
|
|
112
|
+
barrier(CLK_LOCAL_MEM_FENCE);
|
|
113
|
+
if (groupId == 0) totalSum += reduceLM[SIMDGROUP_WIDTH * 0 + slid];
|
|
114
|
+
if (groupId == 0) totalSum += reduceLM[SIMDGROUP_WIDTH * 1 + slid];
|
|
115
|
+
if (groupId == 0) totalSum += reduceLM[SIMDGROUP_WIDTH * 2 + slid];
|
|
116
|
+
|
|
117
|
+
if (groupId == 0) {
|
|
118
|
+
dst = (global float*)((global char*)dst + offsetd);
|
|
119
|
+
dst[gid] = totalSum;
|
|
120
|
+
}
|
|
121
|
+
}
|