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
|
@@ -0,0 +1,1200 @@
|
|
|
1
|
+
// Dynamic quantizers that produce tiled activations
|
|
2
|
+
|
|
3
|
+
static inline void quantize_block_f32_q8_1_tiled(float * restrict x, uint8_t * restrict y_block) {
|
|
4
|
+
assert((unsigned long) x % 128 == 0);
|
|
5
|
+
assert((unsigned long) y_block % 128 == 0);
|
|
6
|
+
|
|
7
|
+
HVX_Vector * vx = (HVX_Vector *) x;
|
|
8
|
+
HVX_Vector zero = Q6_V_vzero();
|
|
9
|
+
|
|
10
|
+
HVX_Vector vmax0_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[0]));
|
|
11
|
+
HVX_Vector vmax1_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[1]));
|
|
12
|
+
HVX_Vector vmax2_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[2]));
|
|
13
|
+
HVX_Vector vmax3_sf = hvx_vec_reduce_max_f32(hvx_vec_abs_f32(vx[3]));
|
|
14
|
+
|
|
15
|
+
HVX_Vector vx0_qf = Q6_Vqf32_vsub_VsfVsf(vx[0], zero);
|
|
16
|
+
HVX_Vector vx1_qf = Q6_Vqf32_vsub_VsfVsf(vx[1], zero);
|
|
17
|
+
HVX_Vector vx2_qf = Q6_Vqf32_vsub_VsfVsf(vx[2], zero);
|
|
18
|
+
HVX_Vector vx3_qf = Q6_Vqf32_vsub_VsfVsf(vx[3], zero);
|
|
19
|
+
|
|
20
|
+
HVX_Vector vmax0_qf = Q6_Vqf32_vsub_VsfVsf(vmax0_sf, zero);
|
|
21
|
+
HVX_Vector vmax1_qf = Q6_Vqf32_vsub_VsfVsf(vmax1_sf, zero);
|
|
22
|
+
HVX_Vector vmax2_qf = Q6_Vqf32_vsub_VsfVsf(vmax2_sf, zero);
|
|
23
|
+
HVX_Vector vmax3_qf = Q6_Vqf32_vsub_VsfVsf(vmax3_sf, zero);
|
|
24
|
+
|
|
25
|
+
HVX_Vector vmax01_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vmax1_qf, vmax0_qf)));
|
|
26
|
+
HVX_Vector vmax23_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vmax3_qf, vmax2_qf)));
|
|
27
|
+
|
|
28
|
+
HVX_Vector vx01_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vx1_qf, vx0_qf)));
|
|
29
|
+
HVX_Vector vx23_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vx3_qf, vx2_qf)));
|
|
30
|
+
|
|
31
|
+
HVX_Vector vd01_qf16 = Q6_Vqf16_vmpy_VhfVhf(vmax01_hf, Q6_Vh_vsplat_R(0x2008)); // 1.0 / 127.0
|
|
32
|
+
HVX_Vector vd23_qf16 = Q6_Vqf16_vmpy_VhfVhf(vmax23_hf, Q6_Vh_vsplat_R(0x2008)); // 1.0 / 127.0
|
|
33
|
+
HVX_Vector vd01_hf = Q6_Vhf_equals_Vqf16(vd01_qf16);
|
|
34
|
+
HVX_Vector vd23_hf = Q6_Vhf_equals_Vqf16(vd23_qf16);
|
|
35
|
+
|
|
36
|
+
HVX_Vector vd01_inv_hf = hvx_vec_inverse_f16(vd01_hf);
|
|
37
|
+
HVX_Vector vd23_inv_hf = hvx_vec_inverse_f16(vd23_hf);
|
|
38
|
+
vx01_hf = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(vx01_hf, vd01_inv_hf));
|
|
39
|
+
vx23_hf = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(vx23_hf, vd23_inv_hf));
|
|
40
|
+
|
|
41
|
+
HVX_Vector vx01_i16 = hvx_vec_i16_from_hf_rnd_sat(vx01_hf);
|
|
42
|
+
HVX_Vector vx23_i16 = hvx_vec_i16_from_hf_rnd_sat(vx23_hf);
|
|
43
|
+
HVX_Vector vx_i8 = Q6_Vb_vpack_VhVh_sat(vx23_i16, vx01_i16);
|
|
44
|
+
|
|
45
|
+
const HVX_Vector ones = Q6_Vb_vsplat_R(1);
|
|
46
|
+
HVX_Vector v_sums = Q6_Vw_vrmpy_VbVb(vx_i8, ones);
|
|
47
|
+
v_sums = Q6_Vw_vadd_VwVw(v_sums, Q6_V_vror_VR(v_sums, 4));
|
|
48
|
+
v_sums = Q6_Vw_vadd_VwVw(v_sums, Q6_V_vror_VR(v_sums, 8));
|
|
49
|
+
v_sums = Q6_Vw_vadd_VwVw(v_sums, Q6_V_vror_VR(v_sums, 16));
|
|
50
|
+
|
|
51
|
+
float vmax0[32] __attribute__((aligned(128)));
|
|
52
|
+
float vmax1[32] __attribute__((aligned(128)));
|
|
53
|
+
float vmax2[32] __attribute__((aligned(128)));
|
|
54
|
+
float vmax3[32] __attribute__((aligned(128)));
|
|
55
|
+
int32_t sums[32] __attribute__((aligned(128)));
|
|
56
|
+
|
|
57
|
+
hvx_vec_store_u(vmax0, 128, vmax0_sf);
|
|
58
|
+
hvx_vec_store_u(vmax1, 128, vmax1_sf);
|
|
59
|
+
hvx_vec_store_u(vmax2, 128, vmax2_sf);
|
|
60
|
+
hvx_vec_store_u(vmax3, 128, vmax3_sf);
|
|
61
|
+
hvx_vec_store_u(sums, 128, v_sums);
|
|
62
|
+
|
|
63
|
+
float d0 = vmax0[0] / 127.0f;
|
|
64
|
+
float d1 = vmax1[0] / 127.0f;
|
|
65
|
+
float d2 = vmax2[0] / 127.0f;
|
|
66
|
+
float d3 = vmax3[0] / 127.0f;
|
|
67
|
+
|
|
68
|
+
static const uint8_t __attribute__((aligned(128))) repl[128] = {
|
|
69
|
+
0x00, 0x00, 0x00, 0x00, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
|
|
70
|
+
0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
|
|
71
|
+
0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
|
|
72
|
+
0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
|
|
73
|
+
0x40, 0x40, 0x40, 0x40, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
|
|
74
|
+
0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
|
|
75
|
+
0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
|
|
76
|
+
0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
|
|
77
|
+
};
|
|
78
|
+
HVX_Vector v_repl_ctrl = * (const HVX_Vector *) repl;
|
|
79
|
+
|
|
80
|
+
for (int b = 0; b < 4; b++) {
|
|
81
|
+
HVX_Vector v_act = Q6_V_vror_VR(vx_i8, b * 32);
|
|
82
|
+
|
|
83
|
+
HVX_Vector r0 = Q6_V_vdelta_VV(v_act, v_repl_ctrl);
|
|
84
|
+
HVX_Vector r1 = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act, 4), v_repl_ctrl);
|
|
85
|
+
HVX_Vector r2 = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act, 8), v_repl_ctrl);
|
|
86
|
+
HVX_Vector r3 = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act, 12), v_repl_ctrl);
|
|
87
|
+
HVX_Vector r4 = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act, 16), v_repl_ctrl);
|
|
88
|
+
HVX_Vector r5 = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act, 20), v_repl_ctrl);
|
|
89
|
+
HVX_Vector r6 = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act, 24), v_repl_ctrl);
|
|
90
|
+
HVX_Vector r7 = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act, 28), v_repl_ctrl);
|
|
91
|
+
|
|
92
|
+
__fp16 scale_h, offset_h;
|
|
93
|
+
if (b == 0) {
|
|
94
|
+
scale_h = (__fp16) d0;
|
|
95
|
+
offset_h = (__fp16) (sums[0] * d0);
|
|
96
|
+
} else if (b == 1) {
|
|
97
|
+
scale_h = (__fp16) d1;
|
|
98
|
+
offset_h = (__fp16) (sums[8] * d1);
|
|
99
|
+
} else if (b == 2) {
|
|
100
|
+
scale_h = (__fp16) d2;
|
|
101
|
+
offset_h = (__fp16) (sums[16] * d2);
|
|
102
|
+
} else {
|
|
103
|
+
scale_h = (__fp16) d3;
|
|
104
|
+
offset_h = (__fp16) (sums[24] * d3);
|
|
105
|
+
}
|
|
106
|
+
|
|
107
|
+
HVX_Vector r_scale = Q6_Vh_vsplat_R(*(int16_t *)&scale_h);
|
|
108
|
+
HVX_Vector r_offset = Q6_Vh_vsplat_R(*(int16_t *)&offset_h);
|
|
109
|
+
|
|
110
|
+
HVX_Vector * restrict dst = (HVX_Vector *) (y_block + b * 1280);
|
|
111
|
+
dst[0] = r0;
|
|
112
|
+
dst[1] = r1;
|
|
113
|
+
dst[2] = r2;
|
|
114
|
+
dst[3] = r3;
|
|
115
|
+
dst[4] = r4;
|
|
116
|
+
dst[5] = r5;
|
|
117
|
+
dst[6] = r6;
|
|
118
|
+
dst[7] = r7;
|
|
119
|
+
dst[8] = r_scale;
|
|
120
|
+
dst[9] = r_offset;
|
|
121
|
+
}
|
|
122
|
+
}
|
|
123
|
+
|
|
124
|
+
static inline void quantize_block_f32_q8_0_tiled(float * restrict x, uint8_t * restrict y_block) {
|
|
125
|
+
assert((unsigned long) x % 128 == 0);
|
|
126
|
+
assert((unsigned long) y_block % 128 == 0);
|
|
127
|
+
|
|
128
|
+
HVX_Vector * vx = (HVX_Vector *) x;
|
|
129
|
+
HVX_Vector zero = Q6_V_vzero();
|
|
130
|
+
|
|
131
|
+
HVX_Vector vx0_qf = Q6_Vqf32_vsub_VsfVsf(vx[0], zero);
|
|
132
|
+
HVX_Vector vx1_qf = Q6_Vqf32_vsub_VsfVsf(vx[1], zero);
|
|
133
|
+
HVX_Vector vx2_qf = Q6_Vqf32_vsub_VsfVsf(vx[2], zero);
|
|
134
|
+
HVX_Vector vx3_qf = Q6_Vqf32_vsub_VsfVsf(vx[3], zero);
|
|
135
|
+
|
|
136
|
+
HVX_Vector vx01_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vx1_qf, vx0_qf)));
|
|
137
|
+
HVX_Vector vx23_hf = Q6_Vh_vdeal_Vh(Q6_Vhf_equals_Wqf32(Q6_W_vcombine_VV(vx3_qf, vx2_qf)));
|
|
138
|
+
|
|
139
|
+
HVX_Vector vmax_hf = hvx_vec_reduce_max_f16(hvx_vec_abs_f16(vx01_hf));
|
|
140
|
+
vmax_hf = hvx_vec_reduce_max2_f16(hvx_vec_abs_f16(vx23_hf), vmax_hf);
|
|
141
|
+
|
|
142
|
+
HVX_Vector vd_qf16 = Q6_Vqf16_vmpy_VhfVhf(vmax_hf, Q6_Vh_vsplat_R(0x2008));
|
|
143
|
+
HVX_Vector vd_hf = Q6_Vhf_equals_Vqf16(vd_qf16);
|
|
144
|
+
|
|
145
|
+
HVX_Vector vd_inv_hf = hvx_vec_inverse_f16(vd_hf);
|
|
146
|
+
vx01_hf = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(vx01_hf, vd_inv_hf));
|
|
147
|
+
vx23_hf = Q6_Vhf_equals_Vqf16(Q6_Vqf16_vmpy_VhfVhf(vx23_hf, vd_inv_hf));
|
|
148
|
+
|
|
149
|
+
HVX_Vector vx01_i16 = hvx_vec_i16_from_hf_rnd_sat(vx01_hf);
|
|
150
|
+
HVX_Vector vx23_i16 = hvx_vec_i16_from_hf_rnd_sat(vx23_hf);
|
|
151
|
+
HVX_Vector vx_i8 = Q6_Vb_vpack_VhVh_sat(vx23_i16, vx01_i16);
|
|
152
|
+
|
|
153
|
+
HVX_Vector r_scale = hvx_vec_repl_f16(vd_hf);
|
|
154
|
+
|
|
155
|
+
static const uint8_t __attribute__((aligned(128))) repl[128] = {
|
|
156
|
+
0x00, 0x00, 0x00, 0x00, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
|
|
157
|
+
0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
|
|
158
|
+
0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
|
|
159
|
+
0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
|
|
160
|
+
0x40, 0x40, 0x40, 0x40, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
|
|
161
|
+
0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
|
|
162
|
+
0x20, 0x20, 0x20, 0x20, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
|
|
163
|
+
0x10, 0x10, 0x10, 0x10, 0x04, 0x04, 0x04, 0x04, 0x08, 0x08, 0x08, 0x08, 0x04, 0x04, 0x04, 0x04,
|
|
164
|
+
};
|
|
165
|
+
HVX_Vector v_repl_ctrl = * (const HVX_Vector *) repl;
|
|
166
|
+
|
|
167
|
+
for (int b = 0; b < 4; b++) {
|
|
168
|
+
HVX_Vector v_act = Q6_V_vror_VR(vx_i8, b * 32);
|
|
169
|
+
|
|
170
|
+
HVX_Vector r0 = Q6_V_vdelta_VV(v_act, v_repl_ctrl);
|
|
171
|
+
HVX_Vector r1 = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act, 4), v_repl_ctrl);
|
|
172
|
+
HVX_Vector r2 = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act, 8), v_repl_ctrl);
|
|
173
|
+
HVX_Vector r3 = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act, 12), v_repl_ctrl);
|
|
174
|
+
HVX_Vector r4 = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act, 16), v_repl_ctrl);
|
|
175
|
+
HVX_Vector r5 = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act, 20), v_repl_ctrl);
|
|
176
|
+
HVX_Vector r6 = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act, 24), v_repl_ctrl);
|
|
177
|
+
HVX_Vector r7 = Q6_V_vdelta_VV(Q6_V_vror_VR(v_act, 28), v_repl_ctrl);
|
|
178
|
+
|
|
179
|
+
HVX_Vector * restrict dst = (HVX_Vector *) (y_block + b * 1152);
|
|
180
|
+
dst[0] = r0;
|
|
181
|
+
dst[1] = r1;
|
|
182
|
+
dst[2] = r2;
|
|
183
|
+
dst[3] = r3;
|
|
184
|
+
dst[4] = r4;
|
|
185
|
+
dst[5] = r5;
|
|
186
|
+
dst[6] = r6;
|
|
187
|
+
dst[7] = r7;
|
|
188
|
+
dst[8] = r_scale;
|
|
189
|
+
}
|
|
190
|
+
}
|
|
191
|
+
|
|
192
|
+
static void quantize_row_f32_q8_0_tiled(float * restrict x, uint8_t * restrict y, uint32_t k) {
|
|
193
|
+
assert(k % 32 == 0);
|
|
194
|
+
const uint32_t qk = QK_Q8_0_TILED;
|
|
195
|
+
const uint32_t nb = (k + qk - 1) / qk;
|
|
196
|
+
|
|
197
|
+
for (uint32_t i = 0; i < nb; i++) {
|
|
198
|
+
uint8_t * restrict y_block = y + i * 4 * 1152;
|
|
199
|
+
quantize_block_f32_q8_0_tiled(x + i * qk, y_block);
|
|
200
|
+
}
|
|
201
|
+
}
|
|
202
|
+
|
|
203
|
+
static void quantize_row_f32_q8_1_tiled(float * restrict x, uint8_t * restrict y, uint32_t k) {
|
|
204
|
+
assert(k % 32 == 0);
|
|
205
|
+
const uint32_t qk = QK_Q8_0_TILED;
|
|
206
|
+
const uint32_t nb = (k + qk - 1) / qk;
|
|
207
|
+
|
|
208
|
+
for (uint32_t i = 0; i < nb; i++) {
|
|
209
|
+
uint8_t * restrict y_block = y + i * 4 * 1280;
|
|
210
|
+
quantize_block_f32_q8_1_tiled(x + i * qk, y_block);
|
|
211
|
+
}
|
|
212
|
+
}
|
|
213
|
+
|
|
214
|
+
// Dot kernels & helpers that consume tiled activations
|
|
215
|
+
|
|
216
|
+
static inline HVX_Vector hvx_vec_mul_f16_f16_to_f32_lower32(HVX_Vector v1, HVX_Vector v2) {
|
|
217
|
+
#if __HVX_ARCH__ >= 79
|
|
218
|
+
HVX_VectorPair p = Q6_Wsf_vmpy_VhfVhf(v1, v2);
|
|
219
|
+
return Q6_V_lo_W(Q6_W_vshuff_VVR(Q6_V_hi_W(p), Q6_V_lo_W(p), -4));
|
|
220
|
+
#else
|
|
221
|
+
HVX_VectorPair p = Q6_Wqf32_vmpy_VhfVhf(v1, v2);
|
|
222
|
+
HVX_Vector hi = Q6_Vsf_equals_Vqf32(Q6_V_hi_W(p));
|
|
223
|
+
HVX_Vector lo = Q6_Vsf_equals_Vqf32(Q6_V_lo_W(p));
|
|
224
|
+
return Q6_V_lo_W(Q6_W_vshuff_VVR(hi, lo, -4));
|
|
225
|
+
#endif
|
|
226
|
+
}
|
|
227
|
+
|
|
228
|
+
static inline HVX_Vector unpack_and_interleave_4bit(HVX_Vector v_a, HVX_Vector v_b, HVX_Vector mask_h4) {
|
|
229
|
+
HVX_Vector v_W0 = Q6_V_vand_VV(v_a, mask_h4);
|
|
230
|
+
HVX_Vector v_W1 = Q6_Vub_vlsr_VubR(v_a, 4);
|
|
231
|
+
HVX_Vector v_W2 = Q6_V_vand_VV(v_b, mask_h4);
|
|
232
|
+
HVX_Vector v_W3 = Q6_Vub_vlsr_VubR(v_b, 4);
|
|
233
|
+
|
|
234
|
+
HVX_VectorPair v01_pair = Q6_W_vshuff_VVR(v_W1, v_W0, -1);
|
|
235
|
+
HVX_VectorPair v23_pair = Q6_W_vshuff_VVR(v_W3, v_W2, -1);
|
|
236
|
+
HVX_VectorPair v0123_pair = Q6_W_vshuff_VVR(Q6_V_lo_W(v23_pair), Q6_V_lo_W(v01_pair), -2);
|
|
237
|
+
return Q6_V_lo_W(v0123_pair);
|
|
238
|
+
}
|
|
239
|
+
|
|
240
|
+
static inline HVX_VectorPair unpack_and_interleave_4bit_x2(HVX_Vector v_src, HVX_Vector mask_h4) {
|
|
241
|
+
HVX_Vector v_lo = Q6_V_vand_VV(v_src, mask_h4);
|
|
242
|
+
HVX_Vector v_hi = Q6_Vub_vlsr_VubR(v_src, 4);
|
|
243
|
+
HVX_VectorPair v01_pair = Q6_W_vshuff_VVR(v_hi, v_lo, -1);
|
|
244
|
+
HVX_Vector v01_lo = Q6_V_lo_W(v01_pair);
|
|
245
|
+
HVX_Vector v01_hi = Q6_V_hi_W(v01_pair);
|
|
246
|
+
|
|
247
|
+
HVX_Vector v23_lo = Q6_V_valign_VVR(v01_hi, v01_lo, 64);
|
|
248
|
+
HVX_Vector v_W0 = Q6_V_lo_W(Q6_W_vshuff_VVR(v23_lo, v01_lo, -2));
|
|
249
|
+
|
|
250
|
+
HVX_Vector v67_lo = Q6_V_valign_VVR(v01_lo, v01_hi, 64);
|
|
251
|
+
HVX_Vector v_W1 = Q6_V_lo_W(Q6_W_vshuff_VVR(v67_lo, v01_hi, -2));
|
|
252
|
+
|
|
253
|
+
return Q6_W_vcombine_VV(v_W1, v_W0);
|
|
254
|
+
}
|
|
255
|
+
|
|
256
|
+
static inline HVX_Vector accum_4bit_32x1(
|
|
257
|
+
const HVX_Vector * restrict vptr,
|
|
258
|
+
const HVX_Vector * restrict v_act,
|
|
259
|
+
HVX_Vector i8
|
|
260
|
+
) {
|
|
261
|
+
HVX_Vector v_sum0 = Q6_V_vzero();
|
|
262
|
+
HVX_Vector v_sum1 = Q6_V_vzero();
|
|
263
|
+
HVX_Vector mask_h4 = Q6_Vb_vsplat_R(0x0F);
|
|
264
|
+
|
|
265
|
+
#pragma unroll
|
|
266
|
+
for (int i = 0; i < 4; i++) {
|
|
267
|
+
HVX_VectorPair v_W_pair = unpack_and_interleave_4bit_x2(vptr[i], mask_h4);
|
|
268
|
+
HVX_Vector v_W0 = Q6_Vb_vsub_VbVb(Q6_V_lo_W(v_W_pair), i8);
|
|
269
|
+
HVX_Vector v_W1 = Q6_Vb_vsub_VbVb(Q6_V_hi_W(v_W_pair), i8);
|
|
270
|
+
v_sum0 = Q6_Vw_vrmpyacc_VwVbVb(v_sum0, v_W0, v_act[i * 2 + 0]);
|
|
271
|
+
v_sum1 = Q6_Vw_vrmpyacc_VwVbVb(v_sum1, v_W1, v_act[i * 2 + 1]);
|
|
272
|
+
}
|
|
273
|
+
|
|
274
|
+
return Q6_Vw_vadd_VwVw(v_sum0, v_sum1);
|
|
275
|
+
}
|
|
276
|
+
|
|
277
|
+
static inline HVX_Vector accum_4bit_32x1_lut(
|
|
278
|
+
const HVX_Vector * restrict vptr,
|
|
279
|
+
const HVX_Vector * restrict v_act,
|
|
280
|
+
HVX_Vector mask_h4,
|
|
281
|
+
HVX_Vector lut
|
|
282
|
+
) {
|
|
283
|
+
HVX_Vector v_sum0 = Q6_V_vzero();
|
|
284
|
+
HVX_Vector v_sum1 = Q6_V_vzero();
|
|
285
|
+
|
|
286
|
+
#pragma unroll
|
|
287
|
+
for (int i = 0; i < 4; i++) {
|
|
288
|
+
HVX_VectorPair v_W_pair = unpack_and_interleave_4bit_x2(vptr[i], mask_h4);
|
|
289
|
+
HVX_Vector v_W0 = Q6_Vb_vlut32_VbVbI(Q6_V_lo_W(v_W_pair), lut, 0);
|
|
290
|
+
HVX_Vector v_W1 = Q6_Vb_vlut32_VbVbI(Q6_V_hi_W(v_W_pair), lut, 0);
|
|
291
|
+
v_sum0 = Q6_Vw_vrmpyacc_VwVbVb(v_sum0, v_W0, v_act[i * 2 + 0]);
|
|
292
|
+
v_sum1 = Q6_Vw_vrmpyacc_VwVbVb(v_sum1, v_W1, v_act[i * 2 + 1]);
|
|
293
|
+
}
|
|
294
|
+
|
|
295
|
+
return Q6_Vw_vadd_VwVw(v_sum0, v_sum1);
|
|
296
|
+
}
|
|
297
|
+
|
|
298
|
+
static inline HVX_VectorPair accum_4bit_32x2(
|
|
299
|
+
const HVX_Vector * restrict vptr,
|
|
300
|
+
const HVX_Vector * restrict v_act0,
|
|
301
|
+
const HVX_Vector * restrict v_act1,
|
|
302
|
+
HVX_Vector i8
|
|
303
|
+
) {
|
|
304
|
+
HVX_Vector v_sum0 = Q6_V_vzero();
|
|
305
|
+
HVX_Vector v_sum1 = Q6_V_vzero();
|
|
306
|
+
HVX_Vector mask_h4 = Q6_Vb_vsplat_R(0x0F);
|
|
307
|
+
|
|
308
|
+
#pragma unroll
|
|
309
|
+
for (int i = 0; i < 4; i++) {
|
|
310
|
+
HVX_VectorPair v_W_pair = unpack_and_interleave_4bit_x2(vptr[i], mask_h4);
|
|
311
|
+
HVX_Vector v_W0 = Q6_Vb_vsub_VbVb(Q6_V_lo_W(v_W_pair), i8);
|
|
312
|
+
HVX_Vector v_W1 = Q6_Vb_vsub_VbVb(Q6_V_hi_W(v_W_pair), i8);
|
|
313
|
+
|
|
314
|
+
v_sum0 = Q6_Vw_vrmpyacc_VwVbVb(v_sum0, v_W0, v_act0[i * 2 + 0]);
|
|
315
|
+
v_sum0 = Q6_Vw_vrmpyacc_VwVbVb(v_sum0, v_W1, v_act0[i * 2 + 1]);
|
|
316
|
+
|
|
317
|
+
v_sum1 = Q6_Vw_vrmpyacc_VwVbVb(v_sum1, v_W0, v_act1[i * 2 + 0]);
|
|
318
|
+
v_sum1 = Q6_Vw_vrmpyacc_VwVbVb(v_sum1, v_W1, v_act1[i * 2 + 1]);
|
|
319
|
+
}
|
|
320
|
+
|
|
321
|
+
return Q6_W_vcombine_VV(v_sum1, v_sum0);
|
|
322
|
+
}
|
|
323
|
+
|
|
324
|
+
static inline HVX_VectorPair accum_4bit_32x2_lut(
|
|
325
|
+
const HVX_Vector * restrict vptr,
|
|
326
|
+
const HVX_Vector * restrict v_act0,
|
|
327
|
+
const HVX_Vector * restrict v_act1,
|
|
328
|
+
HVX_Vector mask_h4,
|
|
329
|
+
HVX_Vector lut
|
|
330
|
+
) {
|
|
331
|
+
HVX_Vector v_sum0 = Q6_V_vzero();
|
|
332
|
+
HVX_Vector v_sum1 = Q6_V_vzero();
|
|
333
|
+
|
|
334
|
+
#pragma unroll
|
|
335
|
+
for (int i = 0; i < 4; i++) {
|
|
336
|
+
HVX_VectorPair v_W_pair = unpack_and_interleave_4bit_x2(vptr[i], mask_h4);
|
|
337
|
+
HVX_Vector v_W0 = Q6_Vb_vlut32_VbVbI(Q6_V_lo_W(v_W_pair), lut, 0);
|
|
338
|
+
HVX_Vector v_W1 = Q6_Vb_vlut32_VbVbI(Q6_V_hi_W(v_W_pair), lut, 0);
|
|
339
|
+
|
|
340
|
+
v_sum0 = Q6_Vw_vrmpyacc_VwVbVb(v_sum0, v_W0, v_act0[i * 2 + 0]);
|
|
341
|
+
v_sum0 = Q6_Vw_vrmpyacc_VwVbVb(v_sum0, v_W1, v_act0[i * 2 + 1]);
|
|
342
|
+
|
|
343
|
+
v_sum1 = Q6_Vw_vrmpyacc_VwVbVb(v_sum1, v_W0, v_act1[i * 2 + 0]);
|
|
344
|
+
v_sum1 = Q6_Vw_vrmpyacc_VwVbVb(v_sum1, v_W1, v_act1[i * 2 + 1]);
|
|
345
|
+
}
|
|
346
|
+
|
|
347
|
+
return Q6_W_vcombine_VV(v_sum1, v_sum0);
|
|
348
|
+
}
|
|
349
|
+
|
|
350
|
+
static inline HVX_Vector accum_q8_0_32x1(
|
|
351
|
+
const HVX_Vector * restrict vptr,
|
|
352
|
+
const HVX_Vector * restrict v_act
|
|
353
|
+
) {
|
|
354
|
+
HVX_Vector v_sum = Q6_V_vzero();
|
|
355
|
+
#pragma unroll
|
|
356
|
+
for (int g = 0; g < 8; g++) {
|
|
357
|
+
HVX_Vector v_rot = Q6_V_vror_VR(vptr[g], 64);
|
|
358
|
+
HVX_Vector v_W = Q6_V_lo_W(Q6_W_vshuff_VVR(v_rot, vptr[g], -2));
|
|
359
|
+
v_sum = Q6_Vw_vrmpyacc_VwVbVb(v_sum, v_W, v_act[g]);
|
|
360
|
+
}
|
|
361
|
+
return v_sum;
|
|
362
|
+
}
|
|
363
|
+
|
|
364
|
+
static inline HVX_VectorPair accum_q8_0_32x2(
|
|
365
|
+
const HVX_Vector * restrict vptr,
|
|
366
|
+
const HVX_Vector * restrict v_act0,
|
|
367
|
+
const HVX_Vector * restrict v_act1
|
|
368
|
+
) {
|
|
369
|
+
HVX_Vector v_sum0 = Q6_V_vzero();
|
|
370
|
+
HVX_Vector v_sum1 = Q6_V_vzero();
|
|
371
|
+
#pragma unroll
|
|
372
|
+
for (int g = 0; g < 8; g++) {
|
|
373
|
+
HVX_Vector v_rot = Q6_V_vror_VR(vptr[g], 64);
|
|
374
|
+
HVX_Vector v_W = Q6_V_lo_W(Q6_W_vshuff_VVR(v_rot, vptr[g], -2));
|
|
375
|
+
v_sum0 = Q6_Vw_vrmpyacc_VwVbVb(v_sum0, v_W, v_act0[g]);
|
|
376
|
+
v_sum1 = Q6_Vw_vrmpyacc_VwVbVb(v_sum1, v_W, v_act1[g]);
|
|
377
|
+
}
|
|
378
|
+
return Q6_W_vcombine_VV(v_sum1, v_sum0);
|
|
379
|
+
}
|
|
380
|
+
|
|
381
|
+
static void tiled_vec_dot_q4_0_32x1(const uint32_t n, float * restrict s, const void * restrict vx, const void * restrict vy, uint32_t valid_rows, const float * restrict sz) {
|
|
382
|
+
const uint8_t * restrict tile_ptr = vx;
|
|
383
|
+
const uint8_t * restrict y_q = vy;
|
|
384
|
+
|
|
385
|
+
HVX_Vector v_sum_float = Q6_V_vzero();
|
|
386
|
+
HVX_Vector i8 = Q6_Vb_vsplat_R(8);
|
|
387
|
+
|
|
388
|
+
uint32_t n_k_tiles = n / 32;
|
|
389
|
+
for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
|
|
390
|
+
const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 640);
|
|
391
|
+
const HVX_Vector * restrict v_act = (const HVX_Vector *) (y_q + kt * 1152);
|
|
392
|
+
|
|
393
|
+
HVX_Vector v_sum = accum_4bit_32x1(vptr, v_act, i8);
|
|
394
|
+
HVX_Vector v_sum_sf = Q6_Vsf_equals_Vw(v_sum);
|
|
395
|
+
|
|
396
|
+
HVX_Vector v_scale_w = vptr[4];
|
|
397
|
+
HVX_Vector v_scale_a = v_act[8];
|
|
398
|
+
HVX_Vector v_scale_comb = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w, v_scale_a);
|
|
399
|
+
HVX_Vector v_sum_scaled = hvx_vec_mul_f32_f32(v_sum_sf, v_scale_comb);
|
|
400
|
+
|
|
401
|
+
v_sum_float = hvx_vec_add_f32_f32(v_sum_float, v_sum_scaled);
|
|
402
|
+
}
|
|
403
|
+
|
|
404
|
+
if (sz) {
|
|
405
|
+
hvx_vec_store_u(s, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float, hvx_vmemu(sz)));
|
|
406
|
+
} else {
|
|
407
|
+
hvx_vec_store_u(s, valid_rows * sizeof(float), v_sum_float);
|
|
408
|
+
}
|
|
409
|
+
}
|
|
410
|
+
|
|
411
|
+
static void tiled_vec_dot_q4_0_32x2(const uint32_t n, float * restrict s0, float * restrict s1, const void * restrict vx, const void * restrict vy0, const void * restrict vy1, uint32_t valid_rows, const float * restrict sz0, const float * restrict sz1) {
|
|
412
|
+
const uint8_t * restrict tile_ptr = vx;
|
|
413
|
+
const uint8_t * restrict y0_q = vy0;
|
|
414
|
+
const uint8_t * restrict y1_q = vy1;
|
|
415
|
+
|
|
416
|
+
HVX_Vector v_sum_float_c0 = Q6_V_vzero();
|
|
417
|
+
HVX_Vector v_sum_float_c1 = Q6_V_vzero();
|
|
418
|
+
HVX_Vector i8 = Q6_Vb_vsplat_R(8);
|
|
419
|
+
|
|
420
|
+
uint32_t n_k_tiles = n / 32;
|
|
421
|
+
uint32_t kt = 0;
|
|
422
|
+
for (; kt + 1 < n_k_tiles; kt += 2) {
|
|
423
|
+
const HVX_Vector * restrict vptr0 = (const HVX_Vector *) (tile_ptr + (kt + 0) * 640);
|
|
424
|
+
const HVX_Vector * restrict v_act0_0 = (const HVX_Vector *) (y0_q + (kt + 0) * 1152);
|
|
425
|
+
const HVX_Vector * restrict v_act1_0 = (const HVX_Vector *) (y1_q + (kt + 0) * 1152);
|
|
426
|
+
|
|
427
|
+
const HVX_Vector * restrict vptr1 = (const HVX_Vector *) (tile_ptr + (kt + 1) * 640);
|
|
428
|
+
const HVX_Vector * restrict v_act0_1 = (const HVX_Vector *) (y0_q + (kt + 1) * 1152);
|
|
429
|
+
const HVX_Vector * restrict v_act1_1 = (const HVX_Vector *) (y1_q + (kt + 1) * 1152);
|
|
430
|
+
|
|
431
|
+
HVX_VectorPair v_sums0 = accum_4bit_32x2(vptr0, v_act0_0, v_act1_0, i8);
|
|
432
|
+
HVX_VectorPair v_sums1 = accum_4bit_32x2(vptr1, v_act0_1, v_act1_1, i8);
|
|
433
|
+
|
|
434
|
+
HVX_Vector v_sum_c0_0 = Q6_V_lo_W(v_sums0);
|
|
435
|
+
HVX_Vector v_sum_c1_0 = Q6_V_hi_W(v_sums0);
|
|
436
|
+
HVX_Vector v_sum_c0_1 = Q6_V_lo_W(v_sums1);
|
|
437
|
+
HVX_Vector v_sum_c1_1 = Q6_V_hi_W(v_sums1);
|
|
438
|
+
|
|
439
|
+
HVX_Vector v_sum_sf_c0_0 = Q6_Vsf_equals_Vw(v_sum_c0_0);
|
|
440
|
+
HVX_Vector v_sum_sf_c1_0 = Q6_Vsf_equals_Vw(v_sum_c1_0);
|
|
441
|
+
HVX_Vector v_sum_sf_c0_1 = Q6_Vsf_equals_Vw(v_sum_c0_1);
|
|
442
|
+
HVX_Vector v_sum_sf_c1_1 = Q6_Vsf_equals_Vw(v_sum_c1_1);
|
|
443
|
+
|
|
444
|
+
HVX_Vector v_scale_w0 = vptr0[4];
|
|
445
|
+
HVX_Vector v_scale_w1 = vptr1[4];
|
|
446
|
+
HVX_Vector v_scale_a_c0_0 = v_act0_0[8];
|
|
447
|
+
HVX_Vector v_scale_a_c1_0 = v_act1_0[8];
|
|
448
|
+
HVX_Vector v_scale_a_c0_1 = v_act0_1[8];
|
|
449
|
+
HVX_Vector v_scale_a_c1_1 = v_act1_1[8];
|
|
450
|
+
|
|
451
|
+
HVX_Vector v_scale_comb_c0_0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w0, v_scale_a_c0_0);
|
|
452
|
+
HVX_Vector v_scale_comb_c1_0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w0, v_scale_a_c1_0);
|
|
453
|
+
HVX_Vector v_scale_comb_c0_1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w1, v_scale_a_c0_1);
|
|
454
|
+
HVX_Vector v_scale_comb_c1_1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w1, v_scale_a_c1_1);
|
|
455
|
+
|
|
456
|
+
HVX_Vector v_sum_scaled_c0_0 = hvx_vec_mul_f32_f32(v_sum_sf_c0_0, v_scale_comb_c0_0);
|
|
457
|
+
HVX_Vector v_sum_scaled_c1_0 = hvx_vec_mul_f32_f32(v_sum_sf_c1_0, v_scale_comb_c1_0);
|
|
458
|
+
HVX_Vector v_sum_scaled_c0_1 = hvx_vec_mul_f32_f32(v_sum_sf_c0_1, v_scale_comb_c0_1);
|
|
459
|
+
HVX_Vector v_sum_scaled_c1_1 = hvx_vec_mul_f32_f32(v_sum_sf_c1_1, v_scale_comb_c1_1);
|
|
460
|
+
|
|
461
|
+
v_sum_float_c0 = hvx_vec_add_f32_f32(v_sum_float_c0, hvx_vec_add_f32_f32(v_sum_scaled_c0_0, v_sum_scaled_c0_1));
|
|
462
|
+
v_sum_float_c1 = hvx_vec_add_f32_f32(v_sum_float_c1, hvx_vec_add_f32_f32(v_sum_scaled_c1_0, v_sum_scaled_c1_1));
|
|
463
|
+
}
|
|
464
|
+
|
|
465
|
+
for (; kt < n_k_tiles; kt++) {
|
|
466
|
+
const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 640);
|
|
467
|
+
const HVX_Vector * restrict v_act0 = (const HVX_Vector *) (y0_q + kt * 1152);
|
|
468
|
+
const HVX_Vector * restrict v_act1 = (const HVX_Vector *) (y1_q + kt * 1152);
|
|
469
|
+
|
|
470
|
+
HVX_VectorPair v_sums = accum_4bit_32x2(vptr, v_act0, v_act1, i8);
|
|
471
|
+
HVX_Vector v_sum_c0 = Q6_V_lo_W(v_sums);
|
|
472
|
+
HVX_Vector v_sum_c1 = Q6_V_hi_W(v_sums);
|
|
473
|
+
|
|
474
|
+
HVX_Vector v_sum_sf_c0 = Q6_Vsf_equals_Vw(v_sum_c0);
|
|
475
|
+
HVX_Vector v_sum_sf_c1 = Q6_Vsf_equals_Vw(v_sum_c1);
|
|
476
|
+
|
|
477
|
+
HVX_Vector v_scale_w = vptr[4];
|
|
478
|
+
HVX_Vector v_scale_a_c0 = v_act0[8];
|
|
479
|
+
HVX_Vector v_scale_a_c1 = v_act1[8];
|
|
480
|
+
|
|
481
|
+
HVX_Vector v_scale_comb_c0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w, v_scale_a_c0);
|
|
482
|
+
HVX_Vector v_scale_comb_c1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w, v_scale_a_c1);
|
|
483
|
+
|
|
484
|
+
HVX_Vector v_sum_scaled_c0 = hvx_vec_mul_f32_f32(v_sum_sf_c0, v_scale_comb_c0);
|
|
485
|
+
HVX_Vector v_sum_scaled_c1 = hvx_vec_mul_f32_f32(v_sum_sf_c1, v_scale_comb_c1);
|
|
486
|
+
|
|
487
|
+
v_sum_float_c0 = hvx_vec_add_f32_f32(v_sum_float_c0, v_sum_scaled_c0);
|
|
488
|
+
v_sum_float_c1 = hvx_vec_add_f32_f32(v_sum_float_c1, v_sum_scaled_c1);
|
|
489
|
+
}
|
|
490
|
+
|
|
491
|
+
if (sz0) {
|
|
492
|
+
hvx_vec_store_u(s0, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c0, hvx_vmemu(sz0)));
|
|
493
|
+
} else {
|
|
494
|
+
hvx_vec_store_u(s0, valid_rows * sizeof(float), v_sum_float_c0);
|
|
495
|
+
}
|
|
496
|
+
if (sz1) {
|
|
497
|
+
hvx_vec_store_u(s1, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c1, hvx_vmemu(sz1)));
|
|
498
|
+
} else {
|
|
499
|
+
hvx_vec_store_u(s1, valid_rows * sizeof(float), v_sum_float_c1);
|
|
500
|
+
}
|
|
501
|
+
}
|
|
502
|
+
|
|
503
|
+
static void tiled_vec_dot_q4_1_32x1(const uint32_t n, float * restrict s, const void * restrict vx, const void * restrict vy, uint32_t valid_rows, const float * restrict sz) {
|
|
504
|
+
const uint8_t * restrict tile_ptr = vx;
|
|
505
|
+
const uint8_t * restrict y_q = vy;
|
|
506
|
+
|
|
507
|
+
HVX_Vector v_sum_float = Q6_V_vzero();
|
|
508
|
+
|
|
509
|
+
uint32_t n_k_tiles = n / 32;
|
|
510
|
+
for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
|
|
511
|
+
const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 640);
|
|
512
|
+
const HVX_Vector * restrict v_act = (const HVX_Vector *) (y_q + kt * 1280);
|
|
513
|
+
|
|
514
|
+
HVX_Vector v_sum = accum_4bit_32x1(vptr, v_act, Q6_V_vzero());
|
|
515
|
+
HVX_Vector v_sum_sf = Q6_Vsf_equals_Vw(v_sum);
|
|
516
|
+
|
|
517
|
+
HVX_Vector v_scale_offset = vptr[4];
|
|
518
|
+
HVX_VectorPair p_deal = Q6_W_vdeal_VVR(v_scale_offset, v_scale_offset, -2);
|
|
519
|
+
HVX_Vector v_scale = Q6_V_lo_W(p_deal);
|
|
520
|
+
HVX_Vector v_offset = Q6_V_hi_W(p_deal);
|
|
521
|
+
|
|
522
|
+
HVX_Vector v_scale_a = v_act[8];
|
|
523
|
+
HVX_Vector v_sum_a = v_act[9];
|
|
524
|
+
|
|
525
|
+
HVX_Vector v_scale_comb = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale, v_scale_a);
|
|
526
|
+
HVX_Vector v_offset_comb = hvx_vec_mul_f16_f16_to_f32_lower32(v_offset, v_sum_a);
|
|
527
|
+
|
|
528
|
+
HVX_Vector v_scaled_dot = hvx_vec_mul_f32_f32(v_sum_sf, v_scale_comb);
|
|
529
|
+
HVX_Vector v_sum_scaled = hvx_vec_add_f32_f32(v_scaled_dot, v_offset_comb);
|
|
530
|
+
|
|
531
|
+
v_sum_float = hvx_vec_add_f32_f32(v_sum_float, v_sum_scaled);
|
|
532
|
+
}
|
|
533
|
+
|
|
534
|
+
if (sz) {
|
|
535
|
+
hvx_vec_store_u(s, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float, hvx_vmemu(sz)));
|
|
536
|
+
} else {
|
|
537
|
+
hvx_vec_store_u(s, valid_rows * sizeof(float), v_sum_float);
|
|
538
|
+
}
|
|
539
|
+
}
|
|
540
|
+
|
|
541
|
+
static void tiled_vec_dot_q4_1_32x2(const uint32_t n, float * restrict s0, float * restrict s1, const void * restrict vx, const void * restrict vy0, const void * restrict vy1, uint32_t valid_rows, const float * restrict sz0, const float * restrict sz1) {
|
|
542
|
+
const uint8_t * restrict tile_ptr = vx;
|
|
543
|
+
const uint8_t * restrict y0_q = vy0;
|
|
544
|
+
const uint8_t * restrict y1_q = vy1;
|
|
545
|
+
|
|
546
|
+
HVX_Vector v_sum_float_c0 = Q6_V_vzero();
|
|
547
|
+
HVX_Vector v_sum_float_c1 = Q6_V_vzero();
|
|
548
|
+
|
|
549
|
+
uint32_t n_k_tiles = n / 32;
|
|
550
|
+
uint32_t kt = 0;
|
|
551
|
+
for (; kt + 1 < n_k_tiles; kt += 2) {
|
|
552
|
+
const HVX_Vector * restrict vptr0 = (const HVX_Vector *) (tile_ptr + (kt + 0) * 640);
|
|
553
|
+
const HVX_Vector * restrict v_act0_0 = (const HVX_Vector *) (y0_q + (kt + 0) * 1280);
|
|
554
|
+
const HVX_Vector * restrict v_act1_0 = (const HVX_Vector *) (y1_q + (kt + 0) * 1280);
|
|
555
|
+
|
|
556
|
+
const HVX_Vector * restrict vptr1 = (const HVX_Vector *) (tile_ptr + (kt + 1) * 640);
|
|
557
|
+
const HVX_Vector * restrict v_act0_1 = (const HVX_Vector *) (y0_q + (kt + 1) * 1280);
|
|
558
|
+
const HVX_Vector * restrict v_act1_1 = (const HVX_Vector *) (y1_q + (kt + 1) * 1280);
|
|
559
|
+
|
|
560
|
+
HVX_VectorPair v_sums0 = accum_4bit_32x2(vptr0, v_act0_0, v_act1_0, Q6_V_vzero());
|
|
561
|
+
HVX_VectorPair v_sums1 = accum_4bit_32x2(vptr1, v_act0_1, v_act1_1, Q6_V_vzero());
|
|
562
|
+
|
|
563
|
+
HVX_Vector v_sum_c0_0 = Q6_V_lo_W(v_sums0);
|
|
564
|
+
HVX_Vector v_sum_c1_0 = Q6_V_hi_W(v_sums0);
|
|
565
|
+
HVX_Vector v_sum_c0_1 = Q6_V_lo_W(v_sums1);
|
|
566
|
+
HVX_Vector v_sum_c1_1 = Q6_V_hi_W(v_sums1);
|
|
567
|
+
|
|
568
|
+
HVX_Vector v_sum_sf_c0_0 = Q6_Vsf_equals_Vw(v_sum_c0_0);
|
|
569
|
+
HVX_Vector v_sum_sf_c1_0 = Q6_Vsf_equals_Vw(v_sum_c1_0);
|
|
570
|
+
HVX_Vector v_sum_sf_c0_1 = Q6_Vsf_equals_Vw(v_sum_c0_1);
|
|
571
|
+
HVX_Vector v_sum_sf_c1_1 = Q6_Vsf_equals_Vw(v_sum_c1_1);
|
|
572
|
+
|
|
573
|
+
HVX_Vector v_scale_offset0 = vptr0[4];
|
|
574
|
+
HVX_VectorPair p_deal0 = Q6_W_vdeal_VVR(v_scale_offset0, v_scale_offset0, -2);
|
|
575
|
+
HVX_Vector v_scale0 = Q6_V_lo_W(p_deal0);
|
|
576
|
+
HVX_Vector v_offset0 = Q6_V_hi_W(p_deal0);
|
|
577
|
+
|
|
578
|
+
HVX_Vector v_scale_offset1 = vptr1[4];
|
|
579
|
+
HVX_VectorPair p_deal1 = Q6_W_vdeal_VVR(v_scale_offset1, v_scale_offset1, -2);
|
|
580
|
+
HVX_Vector v_scale1 = Q6_V_lo_W(p_deal1);
|
|
581
|
+
HVX_Vector v_offset1 = Q6_V_hi_W(p_deal1);
|
|
582
|
+
|
|
583
|
+
HVX_Vector v_scale_a_c0_0 = v_act0_0[8];
|
|
584
|
+
HVX_Vector v_sum_a_c0_0 = v_act0_0[9];
|
|
585
|
+
HVX_Vector v_scale_a_c1_0 = v_act1_0[8];
|
|
586
|
+
HVX_Vector v_sum_a_c1_0 = v_act1_0[9];
|
|
587
|
+
|
|
588
|
+
HVX_Vector v_scale_a_c0_1 = v_act0_1[8];
|
|
589
|
+
HVX_Vector v_sum_a_c0_1 = v_act0_1[9];
|
|
590
|
+
HVX_Vector v_scale_a_c1_1 = v_act1_1[8];
|
|
591
|
+
HVX_Vector v_sum_a_c1_1 = v_act1_1[9];
|
|
592
|
+
|
|
593
|
+
HVX_Vector v_scale_comb_c0_0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale0, v_scale_a_c0_0);
|
|
594
|
+
HVX_Vector v_offset_comb_c0_0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_offset0, v_sum_a_c0_0);
|
|
595
|
+
HVX_Vector v_scale_comb_c1_0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale0, v_scale_a_c1_0);
|
|
596
|
+
HVX_Vector v_offset_comb_c1_0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_offset0, v_sum_a_c1_0);
|
|
597
|
+
|
|
598
|
+
HVX_Vector v_scale_comb_c0_1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale1, v_scale_a_c0_1);
|
|
599
|
+
HVX_Vector v_offset_comb_c0_1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_offset1, v_sum_a_c0_1);
|
|
600
|
+
HVX_Vector v_scale_comb_c1_1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale1, v_scale_a_c1_1);
|
|
601
|
+
HVX_Vector v_offset_comb_c1_1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_offset1, v_sum_a_c1_1);
|
|
602
|
+
|
|
603
|
+
HVX_Vector v_scaled_dot_c0_0 = hvx_vec_mul_f32_f32(v_sum_sf_c0_0, v_scale_comb_c0_0);
|
|
604
|
+
HVX_Vector v_sum_scaled_c0_0 = hvx_vec_add_f32_f32(v_scaled_dot_c0_0, v_offset_comb_c0_0);
|
|
605
|
+
|
|
606
|
+
HVX_Vector v_scaled_dot_c1_0 = hvx_vec_mul_f32_f32(v_sum_sf_c1_0, v_scale_comb_c1_0);
|
|
607
|
+
HVX_Vector v_sum_scaled_c1_0 = hvx_vec_add_f32_f32(v_scaled_dot_c1_0, v_offset_comb_c1_0);
|
|
608
|
+
|
|
609
|
+
HVX_Vector v_scaled_dot_c0_1 = hvx_vec_mul_f32_f32(v_sum_sf_c0_1, v_scale_comb_c0_1);
|
|
610
|
+
HVX_Vector v_sum_scaled_c0_1 = hvx_vec_add_f32_f32(v_scaled_dot_c0_1, v_offset_comb_c0_1);
|
|
611
|
+
|
|
612
|
+
HVX_Vector v_scaled_dot_c1_1 = hvx_vec_mul_f32_f32(v_sum_sf_c1_1, v_scale_comb_c1_1);
|
|
613
|
+
HVX_Vector v_sum_scaled_c1_1 = hvx_vec_add_f32_f32(v_scaled_dot_c1_1, v_offset_comb_c1_1);
|
|
614
|
+
|
|
615
|
+
v_sum_float_c0 = hvx_vec_add_f32_f32(v_sum_float_c0, hvx_vec_add_f32_f32(v_sum_scaled_c0_0, v_sum_scaled_c0_1));
|
|
616
|
+
v_sum_float_c1 = hvx_vec_add_f32_f32(v_sum_float_c1, hvx_vec_add_f32_f32(v_sum_scaled_c1_0, v_sum_scaled_c1_1));
|
|
617
|
+
}
|
|
618
|
+
|
|
619
|
+
for (; kt < n_k_tiles; kt++) {
|
|
620
|
+
const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 640);
|
|
621
|
+
const HVX_Vector * restrict v_act0 = (const HVX_Vector *) (y0_q + kt * 1280);
|
|
622
|
+
const HVX_Vector * restrict v_act1 = (const HVX_Vector *) (y1_q + kt * 1280);
|
|
623
|
+
|
|
624
|
+
HVX_VectorPair v_sums = accum_4bit_32x2(vptr, v_act0, v_act1, Q6_V_vzero());
|
|
625
|
+
HVX_Vector v_sum_c0 = Q6_V_lo_W(v_sums);
|
|
626
|
+
HVX_Vector v_sum_c1 = Q6_V_hi_W(v_sums);
|
|
627
|
+
|
|
628
|
+
HVX_Vector v_sum_sf_c0 = Q6_Vsf_equals_Vw(v_sum_c0);
|
|
629
|
+
HVX_Vector v_sum_sf_c1 = Q6_Vsf_equals_Vw(v_sum_c1);
|
|
630
|
+
|
|
631
|
+
HVX_Vector v_scale_offset = vptr[4];
|
|
632
|
+
HVX_VectorPair p_deal = Q6_W_vdeal_VVR(v_scale_offset, v_scale_offset, -2);
|
|
633
|
+
HVX_Vector v_scale = Q6_V_lo_W(p_deal);
|
|
634
|
+
HVX_Vector v_offset = Q6_V_hi_W(p_deal);
|
|
635
|
+
|
|
636
|
+
HVX_Vector v_scale_a_c0 = v_act0[8];
|
|
637
|
+
HVX_Vector v_sum_a_c0 = v_act0[9];
|
|
638
|
+
HVX_Vector v_scale_a_c1 = v_act1[8];
|
|
639
|
+
HVX_Vector v_sum_a_c1 = v_act1[9];
|
|
640
|
+
|
|
641
|
+
HVX_Vector v_scale_comb_c0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale, v_scale_a_c0);
|
|
642
|
+
HVX_Vector v_offset_comb_c0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_offset, v_sum_a_c0);
|
|
643
|
+
HVX_Vector v_scale_comb_c1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale, v_scale_a_c1);
|
|
644
|
+
HVX_Vector v_offset_comb_c1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_offset, v_sum_a_c1);
|
|
645
|
+
|
|
646
|
+
HVX_Vector v_scaled_dot_c0 = hvx_vec_mul_f32_f32(v_sum_sf_c0, v_scale_comb_c0);
|
|
647
|
+
HVX_Vector v_sum_scaled_c0 = hvx_vec_add_f32_f32(v_scaled_dot_c0, v_offset_comb_c0);
|
|
648
|
+
|
|
649
|
+
HVX_Vector v_scaled_dot_c1 = hvx_vec_mul_f32_f32(v_sum_sf_c1, v_scale_comb_c1);
|
|
650
|
+
HVX_Vector v_sum_scaled_c1 = hvx_vec_add_f32_f32(v_scaled_dot_c1, v_offset_comb_c1);
|
|
651
|
+
|
|
652
|
+
v_sum_float_c0 = hvx_vec_add_f32_f32(v_sum_float_c0, v_sum_scaled_c0);
|
|
653
|
+
v_sum_float_c1 = hvx_vec_add_f32_f32(v_sum_float_c1, v_sum_scaled_c1);
|
|
654
|
+
}
|
|
655
|
+
|
|
656
|
+
if (sz0) {
|
|
657
|
+
hvx_vec_store_u(s0, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c0, hvx_vmemu(sz0)));
|
|
658
|
+
} else {
|
|
659
|
+
hvx_vec_store_u(s0, valid_rows * sizeof(float), v_sum_float_c0);
|
|
660
|
+
}
|
|
661
|
+
if (sz1) {
|
|
662
|
+
hvx_vec_store_u(s1, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c1, hvx_vmemu(sz1)));
|
|
663
|
+
} else {
|
|
664
|
+
hvx_vec_store_u(s1, valid_rows * sizeof(float), v_sum_float_c1);
|
|
665
|
+
}
|
|
666
|
+
}
|
|
667
|
+
|
|
668
|
+
static void tiled_vec_dot_q8_0_32x1(const uint32_t n, float * restrict s, const void * restrict vx, const void * restrict vy, uint32_t valid_rows, const float * restrict sz) {
|
|
669
|
+
const uint8_t * restrict tile_ptr = vx;
|
|
670
|
+
const uint8_t * restrict y_q = vy;
|
|
671
|
+
|
|
672
|
+
HVX_Vector v_sum_float = Q6_V_vzero();
|
|
673
|
+
|
|
674
|
+
uint32_t n_k_tiles = n / 32;
|
|
675
|
+
for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
|
|
676
|
+
const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 1152);
|
|
677
|
+
const HVX_Vector * restrict v_act = (const HVX_Vector *) (y_q + kt * 1152);
|
|
678
|
+
|
|
679
|
+
HVX_Vector v_sum = accum_q8_0_32x1(vptr, v_act);
|
|
680
|
+
HVX_Vector v_sum_sf = Q6_Vsf_equals_Vw(v_sum);
|
|
681
|
+
|
|
682
|
+
HVX_Vector v_scale_w = vptr[8];
|
|
683
|
+
HVX_Vector v_scale_a = v_act[8];
|
|
684
|
+
HVX_Vector v_scale_comb = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w, v_scale_a);
|
|
685
|
+
HVX_Vector v_sum_scaled = hvx_vec_mul_f32_f32(v_sum_sf, v_scale_comb);
|
|
686
|
+
|
|
687
|
+
v_sum_float = hvx_vec_add_f32_f32(v_sum_float, v_sum_scaled);
|
|
688
|
+
}
|
|
689
|
+
|
|
690
|
+
if (sz) {
|
|
691
|
+
hvx_vec_store_u(s, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float, hvx_vmemu(sz)));
|
|
692
|
+
} else {
|
|
693
|
+
hvx_vec_store_u(s, valid_rows * sizeof(float), v_sum_float);
|
|
694
|
+
}
|
|
695
|
+
}
|
|
696
|
+
|
|
697
|
+
static void tiled_vec_dot_q8_0_32x2(const uint32_t n, float * restrict s0, float * restrict s1, const void * restrict vx, const void * restrict vy0, const void * restrict vy1, uint32_t valid_rows, const float * restrict sz0, const float * restrict sz1) {
|
|
698
|
+
const uint8_t * restrict tile_ptr = vx;
|
|
699
|
+
const uint8_t * restrict y0_q = vy0;
|
|
700
|
+
const uint8_t * restrict y1_q = vy1;
|
|
701
|
+
|
|
702
|
+
HVX_Vector v_sum_float_c0 = Q6_V_vzero();
|
|
703
|
+
HVX_Vector v_sum_float_c1 = Q6_V_vzero();
|
|
704
|
+
|
|
705
|
+
uint32_t n_k_tiles = n / 32;
|
|
706
|
+
uint32_t kt = 0;
|
|
707
|
+
for (; kt + 1 < n_k_tiles; kt += 2) {
|
|
708
|
+
const HVX_Vector * restrict vptr0 = (const HVX_Vector *) (tile_ptr + (kt + 0) * 1152);
|
|
709
|
+
const HVX_Vector * restrict v_act0_0 = (const HVX_Vector *) (y0_q + (kt + 0) * 1152);
|
|
710
|
+
const HVX_Vector * restrict v_act1_0 = (const HVX_Vector *) (y1_q + (kt + 0) * 1152);
|
|
711
|
+
|
|
712
|
+
const HVX_Vector * restrict vptr1 = (const HVX_Vector *) (tile_ptr + (kt + 1) * 1152);
|
|
713
|
+
const HVX_Vector * restrict v_act0_1 = (const HVX_Vector *) (y0_q + (kt + 1) * 1152);
|
|
714
|
+
const HVX_Vector * restrict v_act1_1 = (const HVX_Vector *) (y1_q + (kt + 1) * 1152);
|
|
715
|
+
|
|
716
|
+
HVX_VectorPair v_sums0 = accum_q8_0_32x2(vptr0, v_act0_0, v_act1_0);
|
|
717
|
+
HVX_VectorPair v_sums1 = accum_q8_0_32x2(vptr1, v_act0_1, v_act1_1);
|
|
718
|
+
|
|
719
|
+
HVX_Vector v_sum_c0_0 = Q6_V_lo_W(v_sums0);
|
|
720
|
+
HVX_Vector v_sum_c1_0 = Q6_V_hi_W(v_sums0);
|
|
721
|
+
HVX_Vector v_sum_c0_1 = Q6_V_lo_W(v_sums1);
|
|
722
|
+
HVX_Vector v_sum_c1_1 = Q6_V_hi_W(v_sums1);
|
|
723
|
+
|
|
724
|
+
HVX_Vector v_sum_sf_c0_0 = Q6_Vsf_equals_Vw(v_sum_c0_0);
|
|
725
|
+
HVX_Vector v_sum_sf_c1_0 = Q6_Vsf_equals_Vw(v_sum_c1_0);
|
|
726
|
+
HVX_Vector v_sum_sf_c0_1 = Q6_Vsf_equals_Vw(v_sum_c0_1);
|
|
727
|
+
HVX_Vector v_sum_sf_c1_1 = Q6_Vsf_equals_Vw(v_sum_c1_1);
|
|
728
|
+
|
|
729
|
+
HVX_Vector v_scale_w0 = vptr0[8];
|
|
730
|
+
HVX_Vector v_scale_w1 = vptr1[8];
|
|
731
|
+
HVX_Vector v_scale_a_c0_0 = v_act0_0[8];
|
|
732
|
+
HVX_Vector v_scale_a_c1_0 = v_act1_0[8];
|
|
733
|
+
HVX_Vector v_scale_a_c0_1 = v_act0_1[8];
|
|
734
|
+
HVX_Vector v_scale_a_c1_1 = v_act1_1[8];
|
|
735
|
+
|
|
736
|
+
HVX_Vector v_scale_comb_c0_0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w0, v_scale_a_c0_0);
|
|
737
|
+
HVX_Vector v_scale_comb_c1_0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w0, v_scale_a_c1_0);
|
|
738
|
+
HVX_Vector v_scale_comb_c0_1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w1, v_scale_a_c0_1);
|
|
739
|
+
HVX_Vector v_scale_comb_c1_1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w1, v_scale_a_c1_1);
|
|
740
|
+
|
|
741
|
+
HVX_Vector v_sum_scaled_c0_0 = hvx_vec_mul_f32_f32(v_sum_sf_c0_0, v_scale_comb_c0_0);
|
|
742
|
+
HVX_Vector v_sum_scaled_c1_0 = hvx_vec_mul_f32_f32(v_sum_sf_c1_0, v_scale_comb_c1_0);
|
|
743
|
+
HVX_Vector v_sum_scaled_c0_1 = hvx_vec_mul_f32_f32(v_sum_sf_c0_1, v_scale_comb_c0_1);
|
|
744
|
+
HVX_Vector v_sum_scaled_c1_1 = hvx_vec_mul_f32_f32(v_sum_sf_c1_1, v_scale_comb_c1_1);
|
|
745
|
+
|
|
746
|
+
v_sum_float_c0 = hvx_vec_add_f32_f32(v_sum_float_c0, hvx_vec_add_f32_f32(v_sum_scaled_c0_0, v_sum_scaled_c0_1));
|
|
747
|
+
v_sum_float_c1 = hvx_vec_add_f32_f32(v_sum_float_c1, hvx_vec_add_f32_f32(v_sum_scaled_c1_0, v_sum_scaled_c1_1));
|
|
748
|
+
}
|
|
749
|
+
|
|
750
|
+
for (; kt < n_k_tiles; kt++) {
|
|
751
|
+
const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 1152);
|
|
752
|
+
const HVX_Vector * restrict v_act0 = (const HVX_Vector *) (y0_q + kt * 1152);
|
|
753
|
+
const HVX_Vector * restrict v_act1 = (const HVX_Vector *) (y1_q + kt * 1152);
|
|
754
|
+
|
|
755
|
+
HVX_VectorPair v_sums = accum_q8_0_32x2(vptr, v_act0, v_act1);
|
|
756
|
+
HVX_Vector v_sum_c0 = Q6_V_lo_W(v_sums);
|
|
757
|
+
HVX_Vector v_sum_c1 = Q6_V_hi_W(v_sums);
|
|
758
|
+
|
|
759
|
+
HVX_Vector v_sum_sf_c0 = Q6_Vsf_equals_Vw(v_sum_c0);
|
|
760
|
+
HVX_Vector v_sum_sf_c1 = Q6_Vsf_equals_Vw(v_sum_c1);
|
|
761
|
+
|
|
762
|
+
HVX_Vector v_scale_w = vptr[8];
|
|
763
|
+
HVX_Vector v_scale_a_c0 = v_act0[8];
|
|
764
|
+
HVX_Vector v_scale_a_c1 = v_act1[8];
|
|
765
|
+
|
|
766
|
+
HVX_Vector v_scale_comb_c0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w, v_scale_a_c0);
|
|
767
|
+
HVX_Vector v_scale_comb_c1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w, v_scale_a_c1);
|
|
768
|
+
|
|
769
|
+
HVX_Vector v_sum_scaled_c0 = hvx_vec_mul_f32_f32(v_sum_sf_c0, v_scale_comb_c0);
|
|
770
|
+
HVX_Vector v_sum_scaled_c1 = hvx_vec_mul_f32_f32(v_sum_sf_c1, v_scale_comb_c1);
|
|
771
|
+
|
|
772
|
+
v_sum_float_c0 = hvx_vec_add_f32_f32(v_sum_float_c0, v_sum_scaled_c0);
|
|
773
|
+
v_sum_float_c1 = hvx_vec_add_f32_f32(v_sum_float_c1, v_sum_scaled_c1);
|
|
774
|
+
}
|
|
775
|
+
|
|
776
|
+
if (sz0) {
|
|
777
|
+
hvx_vec_store_u(s0, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c0, hvx_vmemu(sz0)));
|
|
778
|
+
} else {
|
|
779
|
+
hvx_vec_store_u(s0, valid_rows * sizeof(float), v_sum_float_c0);
|
|
780
|
+
}
|
|
781
|
+
if (sz1) {
|
|
782
|
+
hvx_vec_store_u(s1, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c1, hvx_vmemu(sz1)));
|
|
783
|
+
} else {
|
|
784
|
+
hvx_vec_store_u(s1, valid_rows * sizeof(float), v_sum_float_c1);
|
|
785
|
+
}
|
|
786
|
+
}
|
|
787
|
+
|
|
788
|
+
static void tiled_vec_dot_iq4nl_32x1(const uint32_t n, float * restrict s, const void * restrict vx, const void * restrict vy, uint32_t valid_rows, const float * restrict sz) {
|
|
789
|
+
const uint8_t * restrict tile_ptr = vx;
|
|
790
|
+
const uint8_t * restrict y_q = vy;
|
|
791
|
+
|
|
792
|
+
HVX_Vector v_sum_float = Q6_V_vzero();
|
|
793
|
+
HVX_Vector mask_h4 = Q6_Vb_vsplat_R(0x0F);
|
|
794
|
+
HVX_Vector lut = *(const HVX_Vector *) kvalues_iq4nl_lut;
|
|
795
|
+
|
|
796
|
+
uint32_t n_k_tiles = n / 32;
|
|
797
|
+
for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
|
|
798
|
+
const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 640);
|
|
799
|
+
const HVX_Vector * restrict v_act = (const HVX_Vector *) (y_q + kt * 1152);
|
|
800
|
+
|
|
801
|
+
HVX_Vector v_sum = accum_4bit_32x1_lut(vptr, v_act, mask_h4, lut);
|
|
802
|
+
HVX_Vector v_sum_sf = Q6_Vsf_equals_Vw(v_sum);
|
|
803
|
+
|
|
804
|
+
HVX_Vector v_scale_w = vptr[4];
|
|
805
|
+
HVX_Vector v_scale_a = v_act[8];
|
|
806
|
+
HVX_Vector v_scale_comb = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w, v_scale_a);
|
|
807
|
+
HVX_Vector v_sum_scaled = hvx_vec_mul_f32_f32(v_sum_sf, v_scale_comb);
|
|
808
|
+
|
|
809
|
+
v_sum_float = hvx_vec_add_f32_f32(v_sum_float, v_sum_scaled);
|
|
810
|
+
}
|
|
811
|
+
|
|
812
|
+
if (sz) {
|
|
813
|
+
hvx_vec_store_u(s, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float, hvx_vmemu(sz)));
|
|
814
|
+
} else {
|
|
815
|
+
hvx_vec_store_u(s, valid_rows * sizeof(float), v_sum_float);
|
|
816
|
+
}
|
|
817
|
+
}
|
|
818
|
+
|
|
819
|
+
static void tiled_vec_dot_iq4nl_32x2(const uint32_t n, float * restrict s0, float * restrict s1, const void * restrict vx, const void * restrict vy0, const void * restrict vy1, uint32_t valid_rows, const float * restrict sz0, const float * restrict sz1) {
|
|
820
|
+
const uint8_t * restrict tile_ptr = vx;
|
|
821
|
+
const uint8_t * restrict y0_q = vy0;
|
|
822
|
+
const uint8_t * restrict y1_q = vy1;
|
|
823
|
+
|
|
824
|
+
HVX_Vector v_sum_float_c0 = Q6_V_vzero();
|
|
825
|
+
HVX_Vector v_sum_float_c1 = Q6_V_vzero();
|
|
826
|
+
HVX_Vector mask_h4 = Q6_Vb_vsplat_R(0x0F);
|
|
827
|
+
HVX_Vector lut = *(const HVX_Vector *) kvalues_iq4nl_lut;
|
|
828
|
+
|
|
829
|
+
uint32_t n_k_tiles = n / 32;
|
|
830
|
+
uint32_t kt = 0;
|
|
831
|
+
for (; kt + 1 < n_k_tiles; kt += 2) {
|
|
832
|
+
const HVX_Vector * restrict vptr0 = (const HVX_Vector *) (tile_ptr + (kt + 0) * 640);
|
|
833
|
+
const HVX_Vector * restrict v_act0_0 = (const HVX_Vector *) (y0_q + (kt + 0) * 1152);
|
|
834
|
+
const HVX_Vector * restrict v_act1_0 = (const HVX_Vector *) (y1_q + (kt + 0) * 1152);
|
|
835
|
+
|
|
836
|
+
const HVX_Vector * restrict vptr1 = (const HVX_Vector *) (tile_ptr + (kt + 1) * 640);
|
|
837
|
+
const HVX_Vector * restrict v_act0_1 = (const HVX_Vector *) (y0_q + (kt + 1) * 1152);
|
|
838
|
+
const HVX_Vector * restrict v_act1_1 = (const HVX_Vector *) (y1_q + (kt + 1) * 1152);
|
|
839
|
+
|
|
840
|
+
HVX_VectorPair v_sums0 = accum_4bit_32x2_lut(vptr0, v_act0_0, v_act1_0, mask_h4, lut);
|
|
841
|
+
HVX_VectorPair v_sums1 = accum_4bit_32x2_lut(vptr1, v_act0_1, v_act1_1, mask_h4, lut);
|
|
842
|
+
|
|
843
|
+
HVX_Vector v_sum_c0_0 = Q6_V_lo_W(v_sums0);
|
|
844
|
+
HVX_Vector v_sum_c1_0 = Q6_V_hi_W(v_sums0);
|
|
845
|
+
HVX_Vector v_sum_c0_1 = Q6_V_lo_W(v_sums1);
|
|
846
|
+
HVX_Vector v_sum_c1_1 = Q6_V_hi_W(v_sums1);
|
|
847
|
+
|
|
848
|
+
HVX_Vector v_sum_sf_c0_0 = Q6_Vsf_equals_Vw(v_sum_c0_0);
|
|
849
|
+
HVX_Vector v_sum_sf_c1_0 = Q6_Vsf_equals_Vw(v_sum_c1_0);
|
|
850
|
+
HVX_Vector v_sum_sf_c0_1 = Q6_Vsf_equals_Vw(v_sum_c0_1);
|
|
851
|
+
HVX_Vector v_sum_sf_c1_1 = Q6_Vsf_equals_Vw(v_sum_c1_1);
|
|
852
|
+
|
|
853
|
+
HVX_Vector v_scale_w0 = vptr0[4];
|
|
854
|
+
HVX_Vector v_scale_w1 = vptr1[4];
|
|
855
|
+
HVX_Vector v_scale_a_c0_0 = v_act0_0[8];
|
|
856
|
+
HVX_Vector v_scale_a_c1_0 = v_act1_0[8];
|
|
857
|
+
HVX_Vector v_scale_a_c0_1 = v_act0_1[8];
|
|
858
|
+
HVX_Vector v_scale_a_c1_1 = v_act1_1[8];
|
|
859
|
+
|
|
860
|
+
HVX_Vector v_scale_comb_c0_0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w0, v_scale_a_c0_0);
|
|
861
|
+
HVX_Vector v_scale_comb_c1_0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w0, v_scale_a_c1_0);
|
|
862
|
+
HVX_Vector v_scale_comb_c0_1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w1, v_scale_a_c0_1);
|
|
863
|
+
HVX_Vector v_scale_comb_c1_1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w1, v_scale_a_c1_1);
|
|
864
|
+
|
|
865
|
+
HVX_Vector v_sum_scaled_c0_0 = hvx_vec_mul_f32_f32(v_sum_sf_c0_0, v_scale_comb_c0_0);
|
|
866
|
+
HVX_Vector v_sum_scaled_c1_0 = hvx_vec_mul_f32_f32(v_sum_sf_c1_0, v_scale_comb_c1_0);
|
|
867
|
+
HVX_Vector v_sum_scaled_c0_1 = hvx_vec_mul_f32_f32(v_sum_sf_c0_1, v_scale_comb_c0_1);
|
|
868
|
+
HVX_Vector v_sum_scaled_c1_1 = hvx_vec_mul_f32_f32(v_sum_sf_c1_1, v_scale_comb_c1_1);
|
|
869
|
+
|
|
870
|
+
v_sum_float_c0 = hvx_vec_add_f32_f32(v_sum_float_c0, hvx_vec_add_f32_f32(v_sum_scaled_c0_0, v_sum_scaled_c0_1));
|
|
871
|
+
v_sum_float_c1 = hvx_vec_add_f32_f32(v_sum_float_c1, hvx_vec_add_f32_f32(v_sum_scaled_c1_0, v_sum_scaled_c1_1));
|
|
872
|
+
}
|
|
873
|
+
|
|
874
|
+
for (; kt < n_k_tiles; kt++) {
|
|
875
|
+
const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 640);
|
|
876
|
+
const HVX_Vector * restrict v_act0 = (const HVX_Vector *) (y0_q + kt * 1152);
|
|
877
|
+
const HVX_Vector * restrict v_act1 = (const HVX_Vector *) (y1_q + kt * 1152);
|
|
878
|
+
|
|
879
|
+
HVX_VectorPair v_sums = accum_4bit_32x2_lut(vptr, v_act0, v_act1, mask_h4, lut);
|
|
880
|
+
HVX_Vector v_sum_c0 = Q6_V_lo_W(v_sums);
|
|
881
|
+
HVX_Vector v_sum_c1 = Q6_V_hi_W(v_sums);
|
|
882
|
+
|
|
883
|
+
HVX_Vector v_sum_sf_c0 = Q6_Vsf_equals_Vw(v_sum_c0);
|
|
884
|
+
HVX_Vector v_sum_sf_c1 = Q6_Vsf_equals_Vw(v_sum_c1);
|
|
885
|
+
|
|
886
|
+
HVX_Vector v_scale_w = vptr[4];
|
|
887
|
+
HVX_Vector v_scale_a_c0 = v_act0[8];
|
|
888
|
+
HVX_Vector v_scale_a_c1 = v_act1[8];
|
|
889
|
+
|
|
890
|
+
HVX_Vector v_scale_comb_c0 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w, v_scale_a_c0);
|
|
891
|
+
HVX_Vector v_scale_comb_c1 = hvx_vec_mul_f16_f16_to_f32_lower32(v_scale_w, v_scale_a_c1);
|
|
892
|
+
|
|
893
|
+
HVX_Vector v_sum_scaled_c0 = hvx_vec_mul_f32_f32(v_sum_sf_c0, v_scale_comb_c0);
|
|
894
|
+
HVX_Vector v_sum_scaled_c1 = hvx_vec_mul_f32_f32(v_sum_sf_c1, v_scale_comb_c1);
|
|
895
|
+
|
|
896
|
+
v_sum_float_c0 = hvx_vec_add_f32_f32(v_sum_float_c0, v_sum_scaled_c0);
|
|
897
|
+
v_sum_float_c1 = hvx_vec_add_f32_f32(v_sum_float_c1, v_sum_scaled_c1);
|
|
898
|
+
}
|
|
899
|
+
|
|
900
|
+
if (sz0) {
|
|
901
|
+
hvx_vec_store_u(s0, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c0, hvx_vmemu(sz0)));
|
|
902
|
+
} else {
|
|
903
|
+
hvx_vec_store_u(s0, valid_rows * sizeof(float), v_sum_float_c0);
|
|
904
|
+
}
|
|
905
|
+
if (sz1) {
|
|
906
|
+
hvx_vec_store_u(s1, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c1, hvx_vmemu(sz1)));
|
|
907
|
+
} else {
|
|
908
|
+
hvx_vec_store_u(s1, valid_rows * sizeof(float), v_sum_float_c1);
|
|
909
|
+
}
|
|
910
|
+
}
|
|
911
|
+
|
|
912
|
+
static void tiled_vec_dot_mxfp4_32x1(const uint32_t n, float * restrict s, const void * restrict vx, const void * restrict vy, uint32_t valid_rows, const float * restrict sz) {
|
|
913
|
+
const uint8_t * restrict tile_ptr = vx;
|
|
914
|
+
const uint8_t * restrict y_q = vy;
|
|
915
|
+
|
|
916
|
+
HVX_Vector v_sum_float = Q6_V_vzero();
|
|
917
|
+
HVX_Vector mask_h4 = Q6_Vb_vsplat_R(0x0F);
|
|
918
|
+
HVX_Vector lut = *(const HVX_Vector *) kvalues_mxfp4_lut;
|
|
919
|
+
HVX_Vector expand = *(const HVX_Vector *) expand_x32_e8m0;
|
|
920
|
+
HVX_Vector e8m0_mask = Q6_V_vsplat_R(0x000000ff);
|
|
921
|
+
|
|
922
|
+
uint32_t n_k_tiles = n / 32;
|
|
923
|
+
for (uint32_t kt = 0; kt < n_k_tiles; kt++) {
|
|
924
|
+
const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 640);
|
|
925
|
+
const HVX_Vector * restrict v_act = (const HVX_Vector *) (y_q + kt * 1152);
|
|
926
|
+
|
|
927
|
+
HVX_Vector v_sum = accum_4bit_32x1_lut(vptr, v_act, mask_h4, lut);
|
|
928
|
+
HVX_Vector v_sum_sf = Q6_Vsf_equals_Vw(v_sum);
|
|
929
|
+
|
|
930
|
+
HVX_Vector v_scale_w = hvx_vmem(tile_ptr + kt * 640 + 512);
|
|
931
|
+
HVX_Vector r0_d = Q6_V_vdelta_VV(v_scale_w, expand);
|
|
932
|
+
r0_d = Q6_V_vand_VV(r0_d, e8m0_mask);
|
|
933
|
+
HVX_Vector v_scale_w_f32 = Q6_Vw_vasl_VwR(r0_d, 23);
|
|
934
|
+
|
|
935
|
+
HVX_Vector v_scale_a_f16 = v_act[8];
|
|
936
|
+
HVX_VectorPair p_scale_a_f32 = hvx_vec_f16_to_f32_shuff(v_scale_a_f16);
|
|
937
|
+
HVX_Vector v_scale_a = Q6_V_lo_W(p_scale_a_f32);
|
|
938
|
+
|
|
939
|
+
HVX_Vector v_scale_comb = hvx_vec_mul_f32_f32(v_scale_w_f32, v_scale_a);
|
|
940
|
+
HVX_Vector v_sum_scaled = hvx_vec_mul_f32_f32(v_sum_sf, v_scale_comb);
|
|
941
|
+
|
|
942
|
+
v_sum_float = hvx_vec_add_f32_f32(v_sum_float, v_sum_scaled);
|
|
943
|
+
}
|
|
944
|
+
|
|
945
|
+
v_sum_float = hvx_vec_mul_f32_f32(v_sum_float, hvx_vec_splat_f32(0.5f));
|
|
946
|
+
|
|
947
|
+
if (sz) {
|
|
948
|
+
hvx_vec_store_u(s, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float, hvx_vmemu(sz)));
|
|
949
|
+
} else {
|
|
950
|
+
hvx_vec_store_u(s, valid_rows * sizeof(float), v_sum_float);
|
|
951
|
+
}
|
|
952
|
+
}
|
|
953
|
+
|
|
954
|
+
static void tiled_vec_dot_mxfp4_32x2(const uint32_t n, float * restrict s0, float * restrict s1, const void * restrict vx, const void * restrict vy0, const void * restrict vy1, uint32_t valid_rows, const float * restrict sz0, const float * restrict sz1) {
|
|
955
|
+
const uint8_t * restrict tile_ptr = vx;
|
|
956
|
+
const uint8_t * restrict y0_q = vy0;
|
|
957
|
+
const uint8_t * restrict y1_q = vy1;
|
|
958
|
+
|
|
959
|
+
HVX_Vector v_sum_float_c0 = Q6_V_vzero();
|
|
960
|
+
HVX_Vector v_sum_float_c1 = Q6_V_vzero();
|
|
961
|
+
HVX_Vector mask_h4 = Q6_Vb_vsplat_R(0x0F);
|
|
962
|
+
HVX_Vector lut = *(const HVX_Vector *) kvalues_mxfp4_lut;
|
|
963
|
+
HVX_Vector expand = *(const HVX_Vector *) expand_x32_e8m0;
|
|
964
|
+
HVX_Vector e8m0_mask = Q6_V_vsplat_R(0x000000ff);
|
|
965
|
+
|
|
966
|
+
uint32_t n_k_tiles = n / 32;
|
|
967
|
+
uint32_t kt = 0;
|
|
968
|
+
for (; kt + 1 < n_k_tiles; kt += 2) {
|
|
969
|
+
const HVX_Vector * restrict vptr0 = (const HVX_Vector *) (tile_ptr + (kt + 0) * 640);
|
|
970
|
+
const HVX_Vector * restrict v_act0_0 = (const HVX_Vector *) (y0_q + (kt + 0) * 1152);
|
|
971
|
+
const HVX_Vector * restrict v_act1_0 = (const HVX_Vector *) (y1_q + (kt + 0) * 1152);
|
|
972
|
+
|
|
973
|
+
const HVX_Vector * restrict vptr1 = (const HVX_Vector *) (tile_ptr + (kt + 1) * 640);
|
|
974
|
+
const HVX_Vector * restrict v_act0_1 = (const HVX_Vector *) (y0_q + (kt + 1) * 1152);
|
|
975
|
+
const HVX_Vector * restrict v_act1_1 = (const HVX_Vector *) (y1_q + (kt + 1) * 1152);
|
|
976
|
+
|
|
977
|
+
HVX_VectorPair v_sums0 = accum_4bit_32x2_lut(vptr0, v_act0_0, v_act1_0, mask_h4, lut);
|
|
978
|
+
HVX_VectorPair v_sums1 = accum_4bit_32x2_lut(vptr1, v_act0_1, v_act1_1, mask_h4, lut);
|
|
979
|
+
|
|
980
|
+
HVX_Vector v_sum_c0_0 = Q6_V_lo_W(v_sums0);
|
|
981
|
+
HVX_Vector v_sum_c1_0 = Q6_V_hi_W(v_sums0);
|
|
982
|
+
HVX_Vector v_sum_c0_1 = Q6_V_lo_W(v_sums1);
|
|
983
|
+
HVX_Vector v_sum_c1_1 = Q6_V_hi_W(v_sums1);
|
|
984
|
+
|
|
985
|
+
HVX_Vector v_sum_sf_c0_0 = Q6_Vsf_equals_Vw(v_sum_c0_0);
|
|
986
|
+
HVX_Vector v_sum_sf_c1_0 = Q6_Vsf_equals_Vw(v_sum_c1_0);
|
|
987
|
+
HVX_Vector v_sum_sf_c0_1 = Q6_Vsf_equals_Vw(v_sum_c0_1);
|
|
988
|
+
HVX_Vector v_sum_sf_c1_1 = Q6_Vsf_equals_Vw(v_sum_c1_1);
|
|
989
|
+
|
|
990
|
+
HVX_Vector v_scale_w0 = hvx_vmem(tile_ptr + (kt + 0) * 640 + 512);
|
|
991
|
+
HVX_Vector r0_d0 = Q6_V_vdelta_VV(v_scale_w0, expand);
|
|
992
|
+
r0_d0 = Q6_V_vand_VV(r0_d0, e8m0_mask);
|
|
993
|
+
HVX_Vector v_scale_w_f32_0 = Q6_Vw_vasl_VwR(r0_d0, 23);
|
|
994
|
+
|
|
995
|
+
HVX_Vector v_scale_w1 = hvx_vmem(tile_ptr + (kt + 1) * 640 + 512);
|
|
996
|
+
HVX_Vector r0_d1 = Q6_V_vdelta_VV(v_scale_w1, expand);
|
|
997
|
+
r0_d1 = Q6_V_vand_VV(r0_d1, e8m0_mask);
|
|
998
|
+
HVX_Vector v_scale_w_f32_1 = Q6_Vw_vasl_VwR(r0_d1, 23);
|
|
999
|
+
|
|
1000
|
+
HVX_Vector v_scale_a_c0_f16_0 = v_act0_0[8];
|
|
1001
|
+
HVX_Vector v_scale_a_c1_f16_0 = v_act1_0[8];
|
|
1002
|
+
HVX_Vector v_scale_a_c0_f16_1 = v_act0_1[8];
|
|
1003
|
+
HVX_Vector v_scale_a_c1_f16_1 = v_act1_1[8];
|
|
1004
|
+
|
|
1005
|
+
HVX_VectorPair p_scale_a_c0_f32_0 = hvx_vec_f16_to_f32_shuff(v_scale_a_c0_f16_0);
|
|
1006
|
+
HVX_VectorPair p_scale_a_c1_f32_0 = hvx_vec_f16_to_f32_shuff(v_scale_a_c1_f16_0);
|
|
1007
|
+
HVX_VectorPair p_scale_a_c0_f32_1 = hvx_vec_f16_to_f32_shuff(v_scale_a_c0_f16_1);
|
|
1008
|
+
HVX_VectorPair p_scale_a_c1_f32_1 = hvx_vec_f16_to_f32_shuff(v_scale_a_c1_f16_1);
|
|
1009
|
+
|
|
1010
|
+
HVX_Vector v_scale_a_c0_0 = Q6_V_lo_W(p_scale_a_c0_f32_0);
|
|
1011
|
+
HVX_Vector v_scale_a_c1_0 = Q6_V_lo_W(p_scale_a_c1_f32_0);
|
|
1012
|
+
HVX_Vector v_scale_a_c0_1 = Q6_V_lo_W(p_scale_a_c0_f32_1);
|
|
1013
|
+
HVX_Vector v_scale_a_c1_1 = Q6_V_lo_W(p_scale_a_c1_f32_1);
|
|
1014
|
+
|
|
1015
|
+
HVX_Vector v_scale_comb_c0_0 = hvx_vec_mul_f32_f32(v_scale_w_f32_0, v_scale_a_c0_0);
|
|
1016
|
+
HVX_Vector v_scale_comb_c1_0 = hvx_vec_mul_f32_f32(v_scale_w_f32_0, v_scale_a_c1_0);
|
|
1017
|
+
HVX_Vector v_scale_comb_c0_1 = hvx_vec_mul_f32_f32(v_scale_w_f32_1, v_scale_a_c0_1);
|
|
1018
|
+
HVX_Vector v_scale_comb_c1_1 = hvx_vec_mul_f32_f32(v_scale_w_f32_1, v_scale_a_c1_1);
|
|
1019
|
+
|
|
1020
|
+
HVX_Vector v_sum_scaled_c0_0 = hvx_vec_mul_f32_f32(v_sum_sf_c0_0, v_scale_comb_c0_0);
|
|
1021
|
+
HVX_Vector v_sum_scaled_c1_0 = hvx_vec_mul_f32_f32(v_sum_sf_c1_0, v_scale_comb_c1_0);
|
|
1022
|
+
HVX_Vector v_sum_scaled_c0_1 = hvx_vec_mul_f32_f32(v_sum_sf_c0_1, v_scale_comb_c0_1);
|
|
1023
|
+
HVX_Vector v_sum_scaled_c1_1 = hvx_vec_mul_f32_f32(v_sum_sf_c1_1, v_scale_comb_c1_1);
|
|
1024
|
+
|
|
1025
|
+
v_sum_float_c0 = hvx_vec_add_f32_f32(v_sum_float_c0, hvx_vec_add_f32_f32(v_sum_scaled_c0_0, v_sum_scaled_c0_1));
|
|
1026
|
+
v_sum_float_c1 = hvx_vec_add_f32_f32(v_sum_float_c1, hvx_vec_add_f32_f32(v_sum_scaled_c1_0, v_sum_scaled_c1_1));
|
|
1027
|
+
}
|
|
1028
|
+
|
|
1029
|
+
for (; kt < n_k_tiles; kt++) {
|
|
1030
|
+
const HVX_Vector * restrict vptr = (const HVX_Vector *) (tile_ptr + kt * 640);
|
|
1031
|
+
const HVX_Vector * restrict v_act0 = (const HVX_Vector *) (y0_q + kt * 1152);
|
|
1032
|
+
const HVX_Vector * restrict v_act1 = (const HVX_Vector *) (y1_q + kt * 1152);
|
|
1033
|
+
|
|
1034
|
+
HVX_VectorPair v_sums = accum_4bit_32x2_lut(vptr, v_act0, v_act1, mask_h4, lut);
|
|
1035
|
+
HVX_Vector v_sum_c0 = Q6_V_lo_W(v_sums);
|
|
1036
|
+
HVX_Vector v_sum_c1 = Q6_V_hi_W(v_sums);
|
|
1037
|
+
|
|
1038
|
+
HVX_Vector v_sum_sf_c0 = Q6_Vsf_equals_Vw(v_sum_c0);
|
|
1039
|
+
HVX_Vector v_sum_sf_c1 = Q6_Vsf_equals_Vw(v_sum_c1);
|
|
1040
|
+
|
|
1041
|
+
HVX_Vector v_scale_w = hvx_vmem(tile_ptr + kt * 640 + 512);
|
|
1042
|
+
HVX_Vector r0_d = Q6_V_vdelta_VV(v_scale_w, expand);
|
|
1043
|
+
r0_d = Q6_V_vand_VV(r0_d, e8m0_mask);
|
|
1044
|
+
HVX_Vector v_scale_w_f32 = Q6_Vw_vasl_VwR(r0_d, 23);
|
|
1045
|
+
|
|
1046
|
+
HVX_Vector v_scale_a_c0_f16 = v_act0[8];
|
|
1047
|
+
HVX_Vector v_scale_a_c1_f16 = v_act1[8];
|
|
1048
|
+
|
|
1049
|
+
HVX_VectorPair p_scale_a_c0_f32 = hvx_vec_f16_to_f32_shuff(v_scale_a_c0_f16);
|
|
1050
|
+
HVX_VectorPair p_scale_a_c1_f32 = hvx_vec_f16_to_f32_shuff(v_scale_a_c1_f16);
|
|
1051
|
+
|
|
1052
|
+
HVX_Vector v_scale_a_c0 = Q6_V_lo_W(p_scale_a_c0_f32);
|
|
1053
|
+
HVX_Vector v_scale_a_c1 = Q6_V_lo_W(p_scale_a_c1_f32);
|
|
1054
|
+
|
|
1055
|
+
HVX_Vector v_scale_comb_c0 = hvx_vec_mul_f32_f32(v_scale_w_f32, v_scale_a_c0);
|
|
1056
|
+
HVX_Vector v_scale_comb_c1 = hvx_vec_mul_f32_f32(v_scale_w_f32, v_scale_a_c1);
|
|
1057
|
+
|
|
1058
|
+
HVX_Vector v_sum_scaled_c0 = hvx_vec_mul_f32_f32(v_sum_sf_c0, v_scale_comb_c0);
|
|
1059
|
+
HVX_Vector v_sum_scaled_c1 = hvx_vec_mul_f32_f32(v_sum_sf_c1, v_scale_comb_c1);
|
|
1060
|
+
|
|
1061
|
+
v_sum_float_c0 = hvx_vec_add_f32_f32(v_sum_float_c0, v_sum_scaled_c0);
|
|
1062
|
+
v_sum_float_c1 = hvx_vec_add_f32_f32(v_sum_float_c1, v_sum_scaled_c1);
|
|
1063
|
+
}
|
|
1064
|
+
|
|
1065
|
+
v_sum_float_c0 = hvx_vec_mul_f32_f32(v_sum_float_c0, hvx_vec_splat_f32(0.5f));
|
|
1066
|
+
v_sum_float_c1 = hvx_vec_mul_f32_f32(v_sum_float_c1, hvx_vec_splat_f32(0.5f));
|
|
1067
|
+
|
|
1068
|
+
if (sz0) {
|
|
1069
|
+
hvx_vec_store_u(s0, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c0, hvx_vmemu(sz0)));
|
|
1070
|
+
} else {
|
|
1071
|
+
hvx_vec_store_u(s0, valid_rows * sizeof(float), v_sum_float_c0);
|
|
1072
|
+
}
|
|
1073
|
+
if (sz1) {
|
|
1074
|
+
hvx_vec_store_u(s1, valid_rows * sizeof(float), hvx_vec_add_f32_f32(v_sum_float_c1, hvx_vmemu(sz1)));
|
|
1075
|
+
} else {
|
|
1076
|
+
hvx_vec_store_u(s1, valid_rows * sizeof(float), v_sum_float_c1);
|
|
1077
|
+
}
|
|
1078
|
+
}
|
|
1079
|
+
|
|
1080
|
+
static inline void quantize_f32_q8_0_tiled_kernel(
|
|
1081
|
+
const uint8_t * restrict src_data,
|
|
1082
|
+
uint8_t * restrict dst_data,
|
|
1083
|
+
uint8_t * restrict tmp_data,
|
|
1084
|
+
uint32_t ne0,
|
|
1085
|
+
uint32_t nrows,
|
|
1086
|
+
size_t src_row_size,
|
|
1087
|
+
size_t dst_row_size
|
|
1088
|
+
) {
|
|
1089
|
+
const size_t src_row_size_padded = hex_round_up(src_row_size, QK_Q8_0_TILED * sizeof(float));
|
|
1090
|
+
hvx_splat_f32_a(tmp_data, 0.0f, src_row_size_padded / sizeof(float));
|
|
1091
|
+
|
|
1092
|
+
for (uint32_t i = 0; i < nrows; ++i) {
|
|
1093
|
+
hex_l2fetch(src_data, src_row_size, src_row_size, 2);
|
|
1094
|
+
hvx_copy_f32_aa(tmp_data, src_data, ne0);
|
|
1095
|
+
|
|
1096
|
+
quantize_row_f32_q8_0_tiled((float *) tmp_data, dst_data, ne0);
|
|
1097
|
+
dst_data += dst_row_size;
|
|
1098
|
+
src_data += src_row_size;
|
|
1099
|
+
}
|
|
1100
|
+
}
|
|
1101
|
+
|
|
1102
|
+
static inline void quantize_f32_q8_1_tiled_kernel(
|
|
1103
|
+
const uint8_t * restrict src_data,
|
|
1104
|
+
uint8_t * restrict dst_data,
|
|
1105
|
+
uint8_t * restrict tmp_data,
|
|
1106
|
+
uint32_t ne0,
|
|
1107
|
+
uint32_t nrows,
|
|
1108
|
+
size_t src_row_size,
|
|
1109
|
+
size_t dst_row_size
|
|
1110
|
+
) {
|
|
1111
|
+
const size_t src_row_size_padded = hex_round_up(src_row_size, QK_Q8_0_TILED * sizeof(float));
|
|
1112
|
+
hvx_splat_f32_a(tmp_data, 0.0f, src_row_size_padded / sizeof(float));
|
|
1113
|
+
|
|
1114
|
+
for (uint32_t i = 0; i < nrows; ++i) {
|
|
1115
|
+
hex_l2fetch(src_data, src_row_size, src_row_size, 2);
|
|
1116
|
+
hvx_copy_f32_aa(tmp_data, src_data, ne0);
|
|
1117
|
+
|
|
1118
|
+
quantize_row_f32_q8_1_tiled((float *) tmp_data, dst_data, ne0);
|
|
1119
|
+
dst_data += dst_row_size;
|
|
1120
|
+
src_data += src_row_size;
|
|
1121
|
+
}
|
|
1122
|
+
}
|
|
1123
|
+
|
|
1124
|
+
static inline void quantize_f32_q8_0_tiled_block_kernel(
|
|
1125
|
+
const float * restrict src,
|
|
1126
|
+
uint8_t * restrict dst,
|
|
1127
|
+
uint8_t * restrict tmp_data,
|
|
1128
|
+
uint32_t ne0,
|
|
1129
|
+
uint32_t ib_first,
|
|
1130
|
+
uint32_t ib_last,
|
|
1131
|
+
size_t src_row_size,
|
|
1132
|
+
size_t dst_row_size,
|
|
1133
|
+
uint32_t r,
|
|
1134
|
+
uint32_t c
|
|
1135
|
+
) {
|
|
1136
|
+
const uint32_t qk = QK_Q8_0_TILED;
|
|
1137
|
+
const uint32_t nb = (ne0 + qk - 1) / qk;
|
|
1138
|
+
|
|
1139
|
+
for (uint32_t ib = ib_first; ib < ib_last; ++ib) {
|
|
1140
|
+
const uint8_t * restrict src_ptr = (const uint8_t *) src + r * src_row_size + c * qk * sizeof(float);
|
|
1141
|
+
uint8_t * restrict dst_ptr = dst + r * dst_row_size + c * 4 * 1152;
|
|
1142
|
+
|
|
1143
|
+
hex_l2fetch(src_ptr, qk * sizeof(float), qk * sizeof(float), 1);
|
|
1144
|
+
|
|
1145
|
+
if (c == nb - 1) {
|
|
1146
|
+
uint32_t active_elements = ne0 - c * qk;
|
|
1147
|
+
hvx_splat_f32_a(tmp_data, 0.0f, qk);
|
|
1148
|
+
hvx_copy_f32_aa(tmp_data, src_ptr, active_elements);
|
|
1149
|
+
} else {
|
|
1150
|
+
hvx_copy_f32_aa(tmp_data, src_ptr, qk);
|
|
1151
|
+
}
|
|
1152
|
+
|
|
1153
|
+
quantize_block_f32_q8_0_tiled((float *) tmp_data, dst_ptr);
|
|
1154
|
+
|
|
1155
|
+
c++;
|
|
1156
|
+
if (c == nb) {
|
|
1157
|
+
c = 0;
|
|
1158
|
+
r++;
|
|
1159
|
+
}
|
|
1160
|
+
}
|
|
1161
|
+
}
|
|
1162
|
+
|
|
1163
|
+
static inline void quantize_f32_q8_1_tiled_block_kernel(
|
|
1164
|
+
const float * restrict src,
|
|
1165
|
+
uint8_t * restrict dst,
|
|
1166
|
+
uint8_t * restrict tmp_data,
|
|
1167
|
+
uint32_t ne0,
|
|
1168
|
+
uint32_t ib_first,
|
|
1169
|
+
uint32_t ib_last,
|
|
1170
|
+
size_t src_row_size,
|
|
1171
|
+
size_t dst_row_size,
|
|
1172
|
+
uint32_t r,
|
|
1173
|
+
uint32_t c
|
|
1174
|
+
) {
|
|
1175
|
+
const uint32_t qk = QK_Q8_0_TILED;
|
|
1176
|
+
const uint32_t nb = (ne0 + qk - 1) / qk;
|
|
1177
|
+
|
|
1178
|
+
for (uint32_t ib = ib_first; ib < ib_last; ++ib) {
|
|
1179
|
+
const uint8_t * restrict src_ptr = (const uint8_t *) src + r * src_row_size + c * qk * sizeof(float);
|
|
1180
|
+
uint8_t * restrict dst_ptr = dst + r * dst_row_size + c * 4 * 1280;
|
|
1181
|
+
|
|
1182
|
+
hex_l2fetch(src_ptr, qk * sizeof(float), qk * sizeof(float), 1);
|
|
1183
|
+
|
|
1184
|
+
if (c == nb - 1) {
|
|
1185
|
+
uint32_t active_elements = ne0 - c * qk;
|
|
1186
|
+
hvx_splat_f32_a(tmp_data, 0.0f, qk);
|
|
1187
|
+
hvx_copy_f32_aa(tmp_data, src_ptr, active_elements);
|
|
1188
|
+
} else {
|
|
1189
|
+
hvx_copy_f32_aa(tmp_data, src_ptr, qk);
|
|
1190
|
+
}
|
|
1191
|
+
|
|
1192
|
+
quantize_block_f32_q8_1_tiled((float *) tmp_data, dst_ptr);
|
|
1193
|
+
|
|
1194
|
+
c++;
|
|
1195
|
+
if (c == nb) {
|
|
1196
|
+
c = 0;
|
|
1197
|
+
r++;
|
|
1198
|
+
}
|
|
1199
|
+
}
|
|
1200
|
+
}
|